Support running host processes alongside VMs (#482)

* Regenerate .pb with protoc-gen-go v1.36.11

* Support running host processes alongside VMs

* Replace existing host processes when starting a set
This commit is contained in:
edi-oai
2026-09-03 00:24:09 +01:00
committed by GitHub
parent 3fe0a284c5
commit a3e08ffcbf
33 changed files with 1454 additions and 309 deletions
+17 -3
View File
@@ -5,14 +5,16 @@ import (
"encoding/json"
"errors"
"fmt"
"time"
"github.com/cirruslabs/orchard/internal/responder"
v1 "github.com/cirruslabs/orchard/pkg/resource/v1"
"github.com/cirruslabs/orchard/rpc"
"github.com/coder/websocket"
"github.com/gin-gonic/gin"
"time"
)
//nolint:protogetter // Preserve the original host-process wire conversion.
func (controller *Controller) rpcWatch(ctx *gin.Context) responder.Responder {
if responder := controller.authorize(ctx, v1.ServiceAccountRoleComputeRead); responder != nil {
return responder
@@ -61,8 +63,20 @@ func (controller *Controller) rpcWatch(ctx *gin.Context) responder.Responder {
case *rpc.WatchInstruction_PortForwardAction:
watchInstruction.PortForwardAction = &v1.PortForwardAction{
Session: typedAction.PortForwardAction.Session,
VMUID: typedAction.PortForwardAction.VmUid,
Port: uint16(typedAction.PortForwardAction.Port),
}
if target := typedAction.PortForwardAction.GetTarget(); target != nil {
watchInstruction.PortForwardAction.Target = &v1.PortForwardTarget{}
if hostProcess := target.GetHostProcess(); hostProcess != nil {
watchInstruction.PortForwardAction.Target.HostProcess = &v1.PortForwardTargetHostProcess{
VMUID: hostProcess.VmUid,
Name: hostProcess.Name,
}
}
} else {
watchInstruction.PortForwardAction.VMUID = typedAction.PortForwardAction.VmUid
watchInstruction.PortForwardAction.Port = uint16(typedAction.PortForwardAction.Port)
}
case *rpc.WatchInstruction_SyncVmsAction:
watchInstruction.SyncVMsAction = &v1.SyncVMsAction{}
+25
View File
@@ -34,6 +34,13 @@ func (controller *Controller) createVM(ctx *gin.Context) responder.Responder {
return responder.JSON(http.StatusBadRequest, NewErrorResponse("invalid JSON was provided"))
}
// Host processes require an additional role
if len(vm.HostProcesses) != 0 {
if responder := controller.authorize(ctx, v1.ServiceAccountRoleHostProcessWrite); responder != nil {
return responder
}
}
if vm.Name == "" {
return responder.JSON(http.StatusPreconditionFailed, NewErrorResponse("VM name is empty"))
} else if err := simplename.Validate(vm.Name); err != nil {
@@ -122,6 +129,11 @@ func (controller *Controller) createVM(ctx *gin.Context) responder.Responder {
vm.RestartPolicy = v1.RestartPolicyNever
}
// Validate hostProcesses
if err := v1.ValidateHostProcesses(vm.HostProcesses); err != nil {
return responder.JSON(http.StatusBadRequest, NewErrorResponse("invalid host processes: %v", err))
}
// Validate hostDirs
if responder := controller.validateHostDirs(vm.HostDirs); responder != nil {
return responder
@@ -182,6 +194,14 @@ func (controller *Controller) updateVMSpec(ctx *gin.Context) responder.Responder
return responder.Error(err)
}
// Changes to host processes require an additional role
//nolint:staticcheck // Preserve the original explicit VMSpec comparison.
if !dbVM.VMSpec.HostProcessesEqual(userVM.VMSpec) {
if responder := controller.authorize(ctx, v1.ServiceAccountRoleHostProcessWrite); responder != nil {
return responder
}
}
if dbVM.TerminalState() {
return responder.JSON(http.StatusPreconditionFailed,
NewErrorResponse("cannot update VM in a terminal state"))
@@ -197,6 +217,11 @@ func (controller *Controller) updateVMSpec(ctx *gin.Context) responder.Responder
return responder.JSON(http.StatusPreconditionFailed, NewErrorResponse("%v", err))
}
// Validate hostProcesses
if err := v1.ValidateHostProcesses(userVM.HostProcesses); err != nil {
return responder.JSON(http.StatusBadRequest, NewErrorResponse("invalid host processes: %v", err))
}
// Softnet-specific logic: automatically enable Softnet when NetSoftnetAllow or NetSoftnetBlock are set
// and propagate deprecated and non-deprecated boolean fields into each other
if userVM.NetSoftnetDeprecated || userVM.NetSoftnet || len(userVM.NetSoftnetAllow) != 0 || len(userVM.NetSoftnetBlock) != 0 {
+1
View File
@@ -211,6 +211,7 @@ func (controller *Controller) newSSHExecSession(
vm.Worker,
vm.UID,
22,
"",
)
if err != nil {
return nil, err
+52 -14
View File
@@ -19,6 +19,7 @@ import (
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"github.com/pkg/errors"
"github.com/samber/lo"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
)
@@ -33,14 +34,29 @@ func (controller *Controller) portForwardVM(ctx *gin.Context) responder.Responde
// Retrieve and parse path and query parameters
name := ctx.Param("name")
portRaw := ctx.Query("port")
port, err := strconv.ParseUint(portRaw, 10, 16)
if err != nil {
return responder.Code(http.StatusBadRequest)
}
if port < 1 || port > 65535 {
return responder.Code(http.StatusBadRequest)
hostProcess := ctx.Query("hostProcess")
// Host process connections require an additional role
var port uint64
var err error
if hostProcess != "" {
if responder := controller.authorizeAny(ctx, v1.ServiceAccountRoleHostProcessWrite,
v1.ServiceAccountRoleHostProcessConnect); responder != nil {
return responder
}
// Host process forwarding cannot also target a VM port
if portRaw != "" {
return responder.Code(http.StatusBadRequest)
}
} else {
// VM port forwarding requires a valid non-zero TCP port
port, err = strconv.ParseUint(portRaw, 10, 16)
if err != nil || port < 1 || port > 65535 {
return responder.Code(http.StatusBadRequest)
}
}
waitRaw := ctx.DefaultQuery("wait", "10")
@@ -58,8 +74,16 @@ func (controller *Controller) portForwardVM(ctx *gin.Context) responder.Responde
return responderImpl
}
// Verify that the requested host process is declared on the VM.
if hostProcess != "" && !lo.SomeBy(vm.HostProcesses, func(process v1.HostProcess) bool {
return process.Name == hostProcess
}) {
return responder.JSON(http.StatusNotFound,
NewErrorResponse("host process %q is not declared on VM %q", hostProcess, vm.Name))
}
// Commence port forwarding
return controller.portForward(ctx, waitContext, vm.Worker, vm.UID, uint32(port))
return controller.portForward(ctx, waitContext, vm.Worker, vm.UID, uint32(port), hostProcess)
}
func (controller *Controller) portForward(
@@ -68,6 +92,7 @@ func (controller *Controller) portForward(
workerName string,
vmUID string,
port uint32,
hostProcess string,
) responder.Responder {
// Request and wait for a connection with a worker
rendezvousConn, err := retry.NewWithData[net.Conn](
@@ -77,7 +102,7 @@ func (controller *Controller) portForward(
retry.Attempts(0),
retry.LastErrorOnly(true),
).Do(func() (net.Conn, error) {
return controller.portForwardConnection(ctx, notifyContext, workerName, vmUID, port)
return controller.portForwardConnection(ctx, notifyContext, workerName, vmUID, port, hostProcess)
})
if err != nil {
if errors.Is(err, errPortForwardRequest) {
@@ -200,6 +225,7 @@ func (controller *Controller) portForwardConnection(
workerName string,
vmUID string,
port uint32,
hostProcess string,
) (net.Conn, error) {
// Create a rendezvous connection point
rendezvousCtx, rendezvousCtxCancel := context.WithCancel(ctx)
@@ -213,13 +239,25 @@ func (controller *Controller) portForwardConnection(
}
// Send request to a worker to initiate a port forwarding connection back to us
portForwardAction := &rpc.WatchInstruction_PortForward{
Session: session,
}
if hostProcess != "" {
portForwardAction.Target = &rpc.WatchInstruction_PortForward_Target{
Value: &rpc.WatchInstruction_PortForward_Target_HostProcess_{
HostProcess: &rpc.WatchInstruction_PortForward_Target_HostProcess{
VmUid: vmUID,
Name: hostProcess,
},
},
}
} else {
portForwardAction.VmUid = vmUID
portForwardAction.Port = port
}
err := controller.workerNotifier.Notify(waitContext, workerName, &rpc.WatchInstruction{
Action: &rpc.WatchInstruction_PortForwardAction{
PortForwardAction: &rpc.WatchInstruction_PortForward{
Session: session,
VmUid: vmUID,
Port: port,
},
PortForwardAction: portForwardAction,
},
})
if err != nil {
@@ -50,5 +50,5 @@ func (controller *Controller) portForwardWorker(ctx *gin.Context) responder.Resp
}
// Commence port-forwarding
return controller.portForward(ctx, waitContext, worker.Name, "", uint32(port))
return controller.portForward(ctx, waitContext, worker.Name, "", uint32(port), "")
}