Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions apps/daemon/internal/dispatch/steering.go
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,11 @@ func (r *Router) queueSteering(ctx context.Context, env proto.Envelope, input pr
if state.steering == nil {
state.steering = make(map[string]steeringReceipt)
}
if err := proto.ValidateSelection(state.declaration, proto.Selection{Messages: input.Input}); err != nil {
ack.ErrorCode, ack.Error = "unsupported", err.Error()
state.steering[input.InputID] = steeringReceipt{fingerprint: fingerprint, ack: ack}
return &ack
}
ack.ErrorCode, ack.Error = "not_ready", "The run is still starting."
// Bind input identity before any retryable state so changed text cannot
// slip through a startup or in-flight retry.
Expand Down
102 changes: 102 additions & 0 deletions apps/daemon/internal/dispatch/steering_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import (
"context"
"errors"
"fmt"
"sync/atomic"
"testing"
"time"

Expand Down Expand Up @@ -248,3 +249,104 @@ func handleSteeringAndWait(t *testing.T, h *harness, env proto.Envelope) error {
waitFor(t, func() bool { return len(h.sender.snapshot()) > before }, "steering ack")
return nil
}

func TestSteeringRetainsAdmittedDeclaration(t *testing.T) {
for _, test := range []struct {
name string
admitted, current proto.CapabilitySupport
available bool
}{
{"narrowed", proto.CapabilitySupported, proto.CapabilityUnsupported, true},
{"widened", proto.CapabilityUnsupported, proto.CapabilitySupported, true},
{"unavailable", proto.CapabilitySupported, proto.CapabilitySupported, false},
} {
t.Run(test.name, func(t *testing.T) {
h := newHarness(t)
defer h.router.Shutdown(context.Background())
var calls atomic.Int32
factory := func(_ context.Context, _ fixtureRun, out chan<- proto.Envelope) (fixtureSession, error) {
return &steeringSession{fakeSession: &fakeSession{out: out, closeOutOnCancel: true}, steer: func(_ context.Context, input proto.PromptSteerPayload, _ func()) error {
calls.Add(1)
if !input.Input.HasImages() {
t.Error("native input lost its image")
}
return nil
}}, nil
}
info := proto.SupportedAgentKind{Kind: "fixture", Available: true, Capabilities: prototest.Capabilities(proto.AgentKindCapabilities{MessageImages: test.admitted})}
registerSession(h.reg, info, factory)
startRun(t, h.router, h.sender, "fixture", "active")
info.Capabilities.MessageImages, info.Available = test.current, test.available
registerSession(h.reg, info, factory)
image := "https://example.com/input.png"
input := proto.PromptSteerPayload{InputID: "image", Input: proto.MessageInput{{Content: []proto.InputContent{{Type: "input_image", ImageURL: &image}}}}}
env := scoped(t, "active", proto.TypePromptSteer, "active", input)
accepted := test.admitted.IsSupported()
code := "unsupported"
if accepted {
code = ""
}
for range 2 {
if err := handleSteeringAndWait(t, h, env); err != nil {
t.Fatal(err)
}
if ack := lastSteeringAck(t, h.sender, "active", "image"); ack.Accepted != accepted || ack.Written || ack.ErrorCode != code {
t.Fatalf("admitted declaration changed: %+v", ack)
}
}
wanted := int32(0)
if accepted {
wanted = 1
}
if calls.Load() != wanted {
t.Fatalf("native calls=%d, want %d", calls.Load(), wanted)
}
// Rejected and accepted receipts both bind the original input identity.
changed := input
changed.Input = proto.TextInput("changed input")
if err := handleSteeringAndWait(t, h, scoped(t, "active", proto.TypePromptSteer, "active", changed)); err != nil {
t.Fatal(err)
}
if ack := lastSteeringAck(t, h.sender, "active", "image"); ack.ErrorCode != "input_conflict" {
t.Fatalf("receipt identity lost: %+v", ack)
}
foreign := env
foreign.Assignment.AssignmentID = "other"
if err := handleSteeringAndWait(t, h, foreign); err != nil {
t.Fatal(err)
}
if ack := lastSteeringAck(t, h.sender, "active", "image"); ack.ErrorCode != proto.AssignmentConflict {
t.Fatalf("receipt escaped assignment: %+v", ack)
}
if test.available {
// Another admitted Turn uses the current declaration, without changing
// the active Turn's permissions or replaying its receipt.
startRun(t, h.router, h.sender, "fixture", "new")
if err := handleSteeringAndWait(t, h, scoped(t, "new", proto.TypePromptSteer, "new", input)); err != nil {
t.Fatal(err)
}
accepted = test.current.IsSupported()
code = "unsupported"
if accepted {
code = ""
wanted++
}
if ack := lastSteeringAck(t, h.sender, "new", "image"); ack.Accepted != accepted || ack.ErrorCode != code {
t.Fatalf("new Turn ignored current declaration: %+v", ack)
}
} else {
assign(t, h.router, "new", "")
prepare := scoped(t, "new", proto.TypeExecutionPrepare, "prepare-new", noEnvironmentPreparation("new", proto.PromptRequestPayload{AgentKind: "fixture"}))
if err := h.router.Handle(t.Context(), prepare); err == nil {
t.Fatal("new admission accepted an unavailable kind")
}
if status := waitPreparationStatus(t, h.sender, prepare.ID, "rejected", ""); status.ErrorCode != "resource_unavailable" {
t.Fatalf("new admission ignored unavailability: %+v", status)
}
}
if calls.Load() != wanted {
t.Fatalf("native calls=%d, want %d", calls.Load(), wanted)
}
})
}
}
8 changes: 4 additions & 4 deletions services/core/internal/execution/delivery.go
Original file line number Diff line number Diff line change
Expand Up @@ -53,7 +53,7 @@ func abort(peer *runtimegateway.Session, ref proto.AssignmentRef, runID string)
_ = send(context.Background(), peer, ref, proto.TypePromptCancel, runID, proto.PromptCancelPayload{})
}

func (d *Dispatcher) deliver(ctx context.Context, tenantID, sessionID string, peer *runtimegateway.Session, request proto.PromptRequestPayload, runID string, input proto.MessageInput, first int64, prepared *preparedStart) (result Result, status string) {
func (d *Dispatcher) deliver(ctx context.Context, tenantID, sessionID string, peer *runtimegateway.Session, request proto.PromptRequestPayload, declaration proto.Declaration, runID string, input proto.MessageInput, first int64, prepared *preparedStart) (result Result, status string) {
changed, unsubscribeChanges := d.notifications.subscribe(tenantID, sessionID)
defer unsubscribeChanges()
status = sessions.TurnFailed
Expand Down Expand Up @@ -102,7 +102,7 @@ func (d *Dispatcher) deliver(ctx context.Context, tenantID, sessionID string, pe
var pending *pendingInput
var cancelSent time.Time
var cancelReply <-chan cancellationResult
functions := &functionExchange{assignment: prepared.assignment, kind: request.AgentKind, turns: d.SessionsReader, sessions: d.sessionExecution, tenant: tenantID, session: sessionID, turn: runID, tools: request.FunctionTools}
functions := &functionExchange{assignment: prepared.assignment, turns: d.SessionsReader, sessions: d.sessionExecution, tenant: tenantID, session: sessionID, turn: runID, tools: request.FunctionTools}
done := false
cancelCtx, stopCancellation := context.WithCancel(ctx)
defer stopCancellation()
Expand Down Expand Up @@ -289,7 +289,7 @@ func (d *Dispatcher) deliver(ctx context.Context, tenantID, sessionID string, pe
cancelSent = time.Now()
continue
}
if err := functions.start(cancelCtx, peer); err != nil {
if err := functions.start(cancelCtx, peer, declaration); err != nil {
result.ErrorCode = "function_result_invalid"
return
}
Expand Down Expand Up @@ -328,7 +328,7 @@ func (d *Dispatcher) deliver(ctx context.Context, tenantID, sessionID string, pe
return
}
if !pending.waiting && !pending.written {
if validateDelivery(peer, request.AgentKind, proto.Selection{Messages: pending.input}) != nil {
if proto.ValidateSelection(declaration, proto.Selection{Messages: pending.input}) != nil {
result.ErrorCode = "message_input_unsupported"
return
}
Expand Down
2 changes: 1 addition & 1 deletion services/core/internal/execution/dispatcher.go
Original file line number Diff line number Diff line change
Expand Up @@ -120,7 +120,7 @@ func (d *Dispatcher) Run(ctx context.Context, tenantID, sessionID, turnID string
if _, err := d.sessionExecution.TransitionTurn(ctx, tenantID, sessionID, turnID, sessions.TurnTransition{ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnInProgress}); err != nil {
return sessions.Turn{}, err
}
result, status := d.deliver(ctx, tenantID, sessionID, peer, req, turnID, text, through, prepared)
result, status := d.deliver(ctx, tenantID, sessionID, peer, req, declaration, turnID, text, through, prepared)
return d.finishRun(tenantID, sessionID, turnID, snapshot.Agent.Model, result, status)
}

Expand Down
5 changes: 2 additions & 3 deletions services/core/internal/execution/functions.go
Original file line number Diff line number Diff line change
Expand Up @@ -53,7 +53,6 @@ type functionExchange struct {
sessions *sessions.ExecutionOperations
tenant, session, turn string
assignment proto.AssignmentRef
kind string
tools []proto.FunctionTool
callID string
reply <-chan functionReply
Expand All @@ -77,7 +76,7 @@ func (f *functionExchange) record(ctx context.Context, env proto.Envelope) error
return f.unlessCancelling(ctx, err)
}

func (f *functionExchange) start(ctx context.Context, peer *runtimegateway.Session) error {
func (f *functionExchange) start(ctx context.Context, peer *runtimegateway.Session, declaration proto.Declaration) error {
if f.reply != nil || len(f.tools) == 0 {
return nil
}
Expand All @@ -93,7 +92,7 @@ func (f *functionExchange) start(ctx context.Context, peer *runtimegateway.Sessi
if err != nil {
return err
}
if err := validateDelivery(peer, f.kind, proto.Selection{FunctionResult: &result}); err != nil {
if err := proto.ValidateSelection(declaration, proto.Selection{FunctionResult: &result}); err != nil {
return err
}
env, err := proto.NewEnvelope(proto.TypeFunctionResult, f.turn, result)
Expand Down
4 changes: 3 additions & 1 deletion services/core/internal/execution/prepared_dispatch.go
Original file line number Diff line number Diff line change
Expand Up @@ -88,6 +88,8 @@ func (d *Dispatcher) RunEnvironmentInput(ctx context.Context, lease Ownership, t
if err != nil || run.Reservation.State != sessions.EnvironmentInputPending {
return run, err
}
// Recheck before claiming without replacing the declaration that admitted
// preparation; later heartbeats cannot expand this Turn's operations.
if _, err := admitSession(peer, session.Engine, snapshot, messages); err != nil {
return run, err
}
Expand All @@ -108,7 +110,7 @@ func (d *Dispatcher) RunEnvironmentInput(ctx context.Context, lease Ownership, t
}
turnID := run.Reservation.Receipts[0].TurnID
through := run.Reservation.Receipts[len(run.Reservation.Receipts)-1].Sequence
result, status := d.deliver(owner, tenantID, sessionID, peer, req, turnID, messages, through, prepared)
result, status := d.deliver(owner, tenantID, sessionID, peer, req, declaration, turnID, messages, through, prepared)
result, status = d.captureCompletedArtifacts(owner, peer, session, environment, bound.Device, turnID, result, status)
run.Turn, err = d.finishRun(tenantID, sessionID, turnID, snapshot.Agent.Model, result, status)
return run, err
Expand Down
10 changes: 0 additions & 10 deletions services/core/internal/execution/support.go
Original file line number Diff line number Diff line change
Expand Up @@ -141,13 +141,3 @@ func admitSession(peer *runtimegateway.Session, engine string, snapshot Snapshot
selection.Messages = messages
return declaration, proto.ValidateSelection(declaration, selection)
}

// validateDelivery checks a message or function result delivered to a running
// Turn against the peer's declaration.
func validateDelivery(peer *runtimegateway.Session, engine string, selection proto.Selection) error {
declaration, err := runtimeDeclaration(peer, engine)
if err != nil {
return err
}
return proto.ValidateSelection(declaration, selection)
}
Loading
Loading