diff --git a/internal/server/control/service.go b/internal/server/control/service.go index 14e1093..510d65e 100644 --- a/internal/server/control/service.go +++ b/internal/server/control/service.go @@ -555,6 +555,9 @@ func (service *Service) SignalCommand(ctx context.Context, request *rvboxv1.Cont if err != nil { return nil, mapStoreError(err) } + if service.wakeClient != nil { + _ = service.wakeClient(request.GetClientId()) + } return &rvboxv1.ControlSignalCommandResponse{CommandRevision: result.CommandRevision}, nil } diff --git a/internal/server/session/agent_server.go b/internal/server/session/agent_server.go index 0311dd4..49709a6 100644 --- a/internal/server/session/agent_server.go +++ b/internal/server/session/agent_server.go @@ -142,13 +142,14 @@ func (server *AgentServer) serveConnection(parent context.Context, connection *w var capacityMu sync.Mutex reservations := make(map[string]DispatchLane) stdinSent := make(map[string]struct{}) + signalSent := make(map[string]struct{}) reconciled := make(chan struct{}) var reconcileOnce sync.Once if !capacity.UpdateAdvertised(0, 0, hello.GetMaxRunningCommands(), hello.GetMaxQueuedCommands()) { server.close(connection, websocket.StatusPolicyViolation, "invalid initial client capacity") return } - go server.dispatchLoop(sessionContext, queue, handle.DispatchWake(), reconciled, &capacityMu, capacity, reservations, stdinSent, hello.GetClientId(), hello.GetPlatform(), encodeSessionID(sessionID), registration.Generation, cancel) + go server.dispatchLoop(sessionContext, queue, handle.DispatchWake(), reconciled, &capacityMu, capacity, reservations, stdinSent, signalSent, hello.GetClientId(), hello.GetPlatform(), encodeSessionID(sessionID), registration.Generation, cancel) encodedSessionID := encodeSessionID(sessionID) welcome, err := proto.Marshal(&rvboxv1.AgentEnvelope{ SessionId: encodedSessionID, SessionGeneration: registration.Generation, @@ -273,6 +274,13 @@ func (server *AgentServer) serveConnection(parent context.Context, connection *w capacityMu.Unlock() handle.SignalDispatch() } + if signalResult := event.GetSignalResult(); signalResult != nil { + issue, _ := domain.ParseUUIDv7(event.GetIssueUuid()) + capacityMu.Lock() + delete(signalSent, signalIntentKey(issue, signalResult.GetCommandRevision(), signalResult.GetSignal())) + capacityMu.Unlock() + handle.SignalDispatch() + } } } } @@ -422,10 +430,49 @@ func stdinIntentKey(issue domain.UUID, writeSeq uint64) string { return issue.String() + ":" + fmt.Sprint(writeSeq) } +func signalIntentKey(issue domain.UUID, revision uint64, signal rvboxv1.SignalKind) string { + return issue.String() + ":" + fmt.Sprint(revision) + ":" + fmt.Sprint(signal) +} + +func (server *AgentServer) enqueueNextSignal(ctx context.Context, queue *WriterQueue, clientID, sessionID string, generation uint64, sent map[string]struct{}, sentMu *sync.Mutex) (bool, error) { + intents, err := server.Store.PendingSignals(ctx, clientID, generation) + if err != nil { + return false, err + } + for _, intent := range intents { + key := signalIntentKey(intent.IssueUUID, intent.CommandRevision, intent.Signal) + sentMu.Lock() + _, alreadySent := sent[key] + if !alreadySent { + sent[key] = struct{}{} + } + sentMu.Unlock() + if alreadySent { + continue + } + envelope := &rvboxv1.AgentEnvelope{SessionId: sessionID, SessionGeneration: generation, Payload: &rvboxv1.AgentEnvelope_SignalCommand{SignalCommand: &rvboxv1.SignalCommand{IssueUuid: intent.IssueUUID.String(), CommandRevision: intent.CommandRevision, Signal: intent.Signal}}} + encoded, marshalErr := proto.Marshal(envelope) + if marshalErr != nil { + sentMu.Lock() + delete(sent, key) + sentMu.Unlock() + return false, marshalErr + } + if enqueueErr := queue.EnqueueControl(Frame{Kind: FrameControl, Payload: encoded}); enqueueErr != nil { + sentMu.Lock() + delete(sent, key) + sentMu.Unlock() + return false, enqueueErr + } + return true, nil + } + return false, nil +} + // dispatchLoop is the per-session serialized dispatcher. It waits for a // complete reconciliation result before consuming queued work, then coalesces // wakeups from local control RPCs, capacity advertisements, and acceptances. -func (server *AgentServer) dispatchLoop(ctx context.Context, queue *WriterQueue, wake <-chan struct{}, reconciled <-chan struct{}, capacityMu *sync.Mutex, capacity *CapacityShadow, reservations map[string]DispatchLane, stdinSent map[string]struct{}, clientID string, platform rvboxv1.Platform, sessionID string, generation uint64, cancel context.CancelFunc) { +func (server *AgentServer) dispatchLoop(ctx context.Context, queue *WriterQueue, wake <-chan struct{}, reconciled <-chan struct{}, capacityMu *sync.Mutex, capacity *CapacityShadow, reservations map[string]DispatchLane, stdinSent, signalSent map[string]struct{}, clientID string, platform rvboxv1.Platform, sessionID string, generation uint64, cancel context.CancelFunc) { select { case <-reconciled: case <-ctx.Done(): @@ -441,6 +488,14 @@ func (server *AgentServer) dispatchLoop(ctx context.Context, queue *WriterQueue, if stdinQueued { continue } + signalQueued, signalErr := server.enqueueNextSignal(ctx, queue, clientID, sessionID, generation, signalSent, capacityMu) + if signalErr != nil { + server.closeForDispatchFailure(cancel) + return + } + if signalQueued { + continue + } capacityMu.Lock() lane := capacity.Reserve() capacityMu.Unlock() diff --git a/internal/server/store/event.go b/internal/server/store/event.go index 994d0bb..e33eaf1 100644 --- a/internal/server/store/event.go +++ b/internal/server/store/event.go @@ -245,7 +245,32 @@ retained_compressed_bytes = retained_compressed_bytes + ?, output_charged_bytes if unmarshalErr := proto.Unmarshal(event.Payload, &wire); unmarshalErr != nil || wire.GetStdinAck() == nil || wire.GetStdinAck().GetWriteSeq() == 0 { err = ErrInvalidSegmentRecord } else { - _, err = tx.ExecContext(ctx, `UPDATE stdin_writes SET acknowledged = 1 WHERE issue_uuid = ? AND write_seq <= ?`, event.IssueUUID[:], wire.GetStdinAck().GetWriteSeq()) + var acknowledged sql.Result + acknowledged, err = tx.ExecContext(ctx, `UPDATE stdin_writes SET acknowledged = 1 WHERE issue_uuid = ? AND write_seq <= ?`, event.IssueUUID[:], wire.GetStdinAck().GetWriteSeq()) + if err == nil { + var affected int64 + affected, err = acknowledged.RowsAffected() + if err == nil && affected == 0 { + err = ErrInvalidSegmentRecord + } + } + } + } + if err == nil && event.EventType == 8 { + var wire rvboxv1.CommandEvent + if unmarshalErr := proto.Unmarshal(event.Payload, &wire); unmarshalErr != nil || wire.GetSignalResult() == nil || wire.GetSignalResult().GetCommandRevision() == 0 || wire.GetSignalResult().GetSignal() < rvboxv1.SignalKind_SIGNAL_HUP || wire.GetSignalResult().GetSignal() > rvboxv1.SignalKind_SIGNAL_USR2 { + err = ErrInvalidSegmentRecord + } else { + var acknowledged sql.Result + result := wire.GetSignalResult() + acknowledged, err = tx.ExecContext(ctx, `UPDATE signal_intents SET acknowledged = 1 WHERE issue_uuid = ? AND command_revision = ? AND signal = ?`, event.IssueUUID[:], result.GetCommandRevision(), result.GetSignal()) + if err == nil { + var affected int64 + affected, err = acknowledged.RowsAffected() + if err == nil && affected == 0 { + err = ErrInvalidSegmentRecord + } + } } } if err == nil { diff --git a/internal/server/store/migrations.go b/internal/server/store/migrations.go index 4e658b9..ab43b5e 100644 --- a/internal/server/store/migrations.go +++ b/internal/server/store/migrations.go @@ -13,7 +13,10 @@ type migration struct { sql string } -var migrations = []migration{{version: 1, sql: schemaV1}} +var migrations = []migration{ + {version: 1, sql: schemaV1}, + {version: 2, sql: schemaV2}, +} func applyMigrations(ctx context.Context, db *sql.DB) error { if _, err := db.ExecContext(ctx, `CREATE TABLE IF NOT EXISTS schema_migrations ( @@ -181,3 +184,15 @@ CREATE TABLE storage_counters ( ) STRICT; INSERT INTO storage_counters(singleton, command_charged_bytes, charge_version) VALUES (1, 0, 1); ` + +const schemaV2 = ` +CREATE TABLE signal_intents ( + request_uuid BLOB PRIMARY KEY CHECK(length(request_uuid) = 16), + issue_uuid BLOB NOT NULL REFERENCES commands(issue_uuid) ON DELETE CASCADE CHECK(length(issue_uuid) = 16), + signal INTEGER NOT NULL CHECK(signal BETWEEN 1 AND 6), + command_revision INTEGER NOT NULL CHECK(command_revision > 0), + acknowledged INTEGER NOT NULL DEFAULT 0 CHECK(acknowledged IN (0,1)), + created_at INTEGER NOT NULL +) STRICT; +CREATE INDEX signal_intents_pending ON signal_intents(issue_uuid, command_revision, signal, acknowledged, created_at, request_uuid); +` diff --git a/internal/server/store/signal.go b/internal/server/store/signal.go index 2009a5d..e7548dd 100644 --- a/internal/server/store/signal.go +++ b/internal/server/store/signal.go @@ -27,12 +27,19 @@ type SignalResult struct { CommandRevision uint64 Duplicate bool Cancelled bool + Pending bool } -// SignalCommand applies the only locally-completable signal operation: a -// queued command can be cancelled before dispatch. For dispatched/accepted/ -// running work the revisioned remote delivery path is intentionally not -// claimed until a live session control queue is available. +type SignalIntent struct { + IssueUUID domain.UUID + CommandRevision uint64 + Signal rvboxv1.SignalKind +} + +// SignalCommand cancels queued work locally or records a revision-bound, +// replayable signal intent for dispatched/accepted/running work. The session +// layer later exposes that intent and the client acknowledges it as a durable +// command event. func (store *Store) SignalCommand(ctx context.Context, input SignalInput) (SignalResult, error) { if input.IssueUUID == (domain.UUID{}) || input.ClientID == "" || input.RequestUUID == (domain.UUID{}) || input.Signal < rvboxv1.SignalKind_SIGNAL_HUP || input.Signal > rvboxv1.SignalKind_SIGNAL_USR2 || input.Hash == [32]byte{} || input.OccurredAt.IsZero() { return SignalResult{}, errors.New("invalid signal request") @@ -60,16 +67,63 @@ func (store *Store) SignalCommand(ctx context.Context, input SignalInput) (Signa defer tx.Rollback() var lifecycle uint32 var revision uint64 - err = tx.QueryRowContext(ctx, `SELECT lifecycle, revision FROM commands WHERE issue_uuid = ? AND client_id = ?`, input.IssueUUID[:], input.ClientID).Scan(&lifecycle, &revision) + var targetGeneration sql.NullInt64 + err = tx.QueryRowContext(ctx, `SELECT lifecycle, revision, target_session_generation FROM commands WHERE issue_uuid = ? AND client_id = ?`, input.IssueUUID[:], input.ClientID).Scan(&lifecycle, &revision, &targetGeneration) if errors.Is(err, sql.ErrNoRows) { return SignalResult{}, ErrCommandNotFound } if err != nil { return SignalResult{}, err } - if lifecycle != uint32(rvboxv1.CommandLifecycle_COMMAND_QUEUED) { + if lifecycle != uint32(rvboxv1.CommandLifecycle_COMMAND_QUEUED) && (lifecycle < uint32(rvboxv1.CommandLifecycle_COMMAND_DISPATCHED) || lifecycle > uint32(rvboxv1.CommandLifecycle_COMMAND_RUNNING) || !targetGeneration.Valid || targetGeneration.Int64 <= 0) { return SignalResult{}, ErrSignalDeliveryUnavailable } + if lifecycle != uint32(rvboxv1.CommandLifecycle_COMMAND_QUEUED) { + charge, chargeErr := EstimateCharge(ChargeInput{SQLiteRows: 1, IndexEntries: 1}) + if chargeErr != nil { + return SignalResult{}, chargeErr + } + var commandCharged, closeout, clientCharged, serverCharged uint64 + if err := tx.QueryRowContext(ctx, `SELECT commands.charged_bytes, commands.closeout_remaining_bytes, clients.charged_bytes, storage_counters.command_charged_bytes +FROM commands JOIN clients ON clients.client_id = commands.client_id JOIN storage_counters ON storage_counters.singleton = 1 +WHERE commands.issue_uuid = ? AND commands.client_id = ?`, input.IssueUUID[:], input.ClientID).Scan(&commandCharged, &closeout, &clientCharged, &serverCharged); err != nil { + return SignalResult{}, err + } + freeBytes, freeErr := store.freeSpaceProbe.AvailableBytes(store.dataDir) + if freeErr != nil { + return SignalResult{}, freeErr + } + reservation, reserveErr := CheckReservation(store.quotaLimits, ReservationState{CommandTotalCharged: commandCharged, CloseoutRemaining: closeout, ClientTotalCharged: clientCharged, ServerTotalCharged: serverCharged, FilesystemFreeBytes: freeBytes}, ReservationRequest{ChargedBytes: charge}) + if reserveErr != nil { + return SignalResult{}, reserveErr + } + if _, err := tx.ExecContext(ctx, `INSERT INTO signal_intents (request_uuid, issue_uuid, signal, command_revision, acknowledged, created_at) VALUES (?, ?, ?, ?, 0, ?)`, input.RequestUUID[:], input.IssueUUID[:], input.Signal, revision, input.OccurredAt.UTC().UnixNano()); err != nil { + return SignalResult{}, err + } + if _, err := tx.ExecContext(ctx, `UPDATE commands SET charged_bytes = ?, closeout_remaining_bytes = ? WHERE issue_uuid = ? AND charged_bytes = ?`, reservation.CommandTotalCharged, reservation.CloseoutRemaining, input.IssueUUID[:], commandCharged); err != nil { + return SignalResult{}, err + } + if _, err := tx.ExecContext(ctx, `UPDATE clients SET charged_bytes = ? WHERE client_id = ? AND charged_bytes = ?`, reservation.ClientTotalCharged, input.ClientID, clientCharged); err != nil { + return SignalResult{}, err + } + if _, err := tx.ExecContext(ctx, `UPDATE storage_counters SET command_charged_bytes = ? WHERE singleton = 1 AND command_charged_bytes = ?`, reservation.ServerTotalCharged, serverCharged); err != nil { + return SignalResult{}, err + } + result := make([]byte, 8) + binary.BigEndian.PutUint64(result, revision) + if _, err := tx.ExecContext(ctx, `INSERT INTO control_mutations (request_uuid, method, owner_kind, owner_id, target, immutable_sha256, assigned_revision, result, created_at) VALUES (?, ?, 'command', ?, ?, ?, ?, ?, ?)`, input.RequestUUID[:], method, input.IssueUUID.String(), target, input.Hash[:], revision, result, input.OccurredAt.UTC().UnixNano()); err != nil { + return SignalResult{}, err + } + auditPayload := []byte{byte(input.Signal)} + auditHash := sha256.Sum256(auditPayload) + if _, err := tx.ExecContext(ctx, `INSERT INTO audit_events (occurred_at, source, action, outcome, compression, payload, raw_bytes, stored_bytes, sha256) VALUES (?, 'control', ?, 'success', 1, ?, ?, ?, ?)`, input.OccurredAt.UTC().UnixNano(), method, auditPayload, len(auditPayload), len(auditPayload), auditHash[:]); err != nil { + return SignalResult{}, err + } + if err := tx.Commit(); err != nil { + return SignalResult{}, err + } + return SignalResult{CommandRevision: revision, Pending: true}, nil + } if revision == ^uint64(0) { return SignalResult{}, errors.New("command revision exhausted") } @@ -96,3 +150,35 @@ func (store *Store) SignalCommand(ctx context.Context, input SignalInput) (Signa } return SignalResult{CommandRevision: nextRevision, Cancelled: true}, nil } + +func (store *Store) PendingSignals(ctx context.Context, clientID string, generation uint64) ([]SignalIntent, error) { + if clientID == "" || generation == 0 { + return nil, errors.New("client ID and generation are required") + } + database, err := store.openDatabase() + if err != nil { + return nil, err + } + rows, err := database.QueryContext(ctx, `SELECT i.issue_uuid, i.command_revision, i.signal +FROM signal_intents i JOIN commands c ON c.issue_uuid = i.issue_uuid +WHERE c.client_id = ? AND c.target_session_generation = ? AND c.lifecycle BETWEEN 2 AND 4 AND i.acknowledged = 0 +ORDER BY i.created_at, i.request_uuid`, clientID, generation) + if err != nil { + return nil, err + } + defer rows.Close() + var result []SignalIntent + for rows.Next() { + var encoded []byte + var intent SignalIntent + if err := rows.Scan(&encoded, &intent.CommandRevision, &intent.Signal); err != nil { + return nil, err + } + if len(encoded) != len(domain.UUID{}) || intent.CommandRevision == 0 || intent.Signal < rvboxv1.SignalKind_SIGNAL_HUP || intent.Signal > rvboxv1.SignalKind_SIGNAL_USR2 { + return nil, ErrInvalidSegmentRecord + } + copy(intent.IssueUUID[:], encoded) + result = append(result, intent) + } + return result, rows.Err() +} diff --git a/test/coverage.toml b/test/coverage.toml index 2478e3c..8c59988 100644 --- a/test/coverage.toml +++ b/test/coverage.toml @@ -338,6 +338,12 @@ layer = "integration" status = "implemented" tests = ["test/integration/clientagent/clientagent_integration_test.go:TestControlStdinIntentReplaysAndAcknowledges_HP_DISPATCH_08"] +[[requirements]] +id = "HP-SIGNAL-01" +layer = "integration" +status = "implemented" +tests = ["test/integration/clientagent/clientagent_integration_test.go:TestControlStdinIntentReplaysAndAcknowledges_HP_DISPATCH_08"] + [[requirements]] id = "HP-EVENT-01" layer = "unit" diff --git a/test/integration/clientagent/clientagent_integration_test.go b/test/integration/clientagent/clientagent_integration_test.go index 9b33143..30c06b3 100644 --- a/test/integration/clientagent/clientagent_integration_test.go +++ b/test/integration/clientagent/clientagent_integration_test.go @@ -300,6 +300,30 @@ func TestControlStdinIntentReplaysAndAcknowledges_HP_DISPATCH_08(t *testing.T) { if err := persistence.DB().QueryRow(`SELECT acknowledged FROM stdin_writes WHERE issue_uuid = ? AND write_seq = 2`, parsedIssue[:]).Scan(&closeAcknowledged); err != nil || closeAcknowledged != 1 { t.Fatalf("stdin close acknowledgement = %d, %v", closeAcknowledged, err) } + signalRequestID := fixedIssueID(0x7d) + signalResponse, err := controlService.SignalCommand(ctx, &rvboxv1.ControlSignalCommandRequest{ClientId: hello.GetClientId(), IssueUuid: issue, RequestId: signalRequestID, Signal: rvboxv1.SignalKind_SIGNAL_TERM}) + if err != nil || signalResponse.GetCommandRevision() != 1 { + t.Fatalf("active signal request = %#v, %v", signalResponse, err) + } + encoded, err = transport.Read(readContext) + if err != nil { + t.Fatal(err) + } + signalEnvelope, err := agentproto.DecodeEnvelope(encoded, agentproto.DefaultLimits(), rvboxv1.Platform_PLATFORM_WINDOWS) + if err != nil || signalEnvelope.GetSignalCommand() == nil || signalEnvelope.GetSignalCommand().GetIssueUuid() != issue || signalEnvelope.GetSignalCommand().GetCommandRevision() != 1 || signalEnvelope.GetSignalCommand().GetSignal() != rvboxv1.SignalKind_SIGNAL_TERM { + t.Fatalf("signal delivery = %#v, %v", signalEnvelope, err) + } + signalAck := &rvboxv1.CommandEvent{IssueUuid: issue, EventSeq: 3, ObservedAt: timestamppb.Now(), Payload: &rvboxv1.CommandEvent_SignalResult{SignalResult: &rvboxv1.SignalResult{Signal: rvboxv1.SignalKind_SIGNAL_TERM, Accepted: true, GracefulDeliveryAttempted: true, CommandRevision: 1, Detail: "delivered"}}} + if err := agent.SendCommandEvent(ctx, transport, accepted, signalAck, agentproto.DefaultLimits()); err != nil { + t.Fatal(err) + } + if _, err := transport.Read(readContext); err != nil { + t.Fatal(err) + } + var signalAcknowledged int + if err := persistence.DB().QueryRow(`SELECT acknowledged FROM signal_intents WHERE issue_uuid = ? AND command_revision = 1 AND signal = ?`, parsedIssue[:], rvboxv1.SignalKind_SIGNAL_TERM).Scan(&signalAcknowledged); err != nil || signalAcknowledged != 1 { + t.Fatalf("signal acknowledgement = %d, %v", signalAcknowledged, err) + } } func fixedIssueID(last byte) string { diff --git a/test/integration/store/store_integration_test.go b/test/integration/store/store_integration_test.go index 2a6640c..280a2c6 100644 --- a/test/integration/store/store_integration_test.go +++ b/test/integration/store/store_integration_test.go @@ -67,13 +67,13 @@ func TestRealSQLiteInitializationAndRestart_HP_STORE_01(t *testing.T) { if err := rows.Close(); err != nil { t.Fatal(err) } - want := []string{"audit_events", "clients", "command_events", "command_payloads", "command_tombstones", "commands", "control_mutations", "output_segments", "output_truncations", "schema_migrations", "sessions", "stdin_writes", "storage_counters", "storage_incidents", "takeover_authorizations"} + want := []string{"audit_events", "clients", "command_events", "command_payloads", "command_tombstones", "commands", "control_mutations", "output_segments", "output_truncations", "schema_migrations", "sessions", "signal_intents", "stdin_writes", "storage_counters", "storage_incidents", "takeover_authorizations"} sort.Strings(want) if strings.Join(names, ",") != strings.Join(want, ",") { t.Fatalf("tables = %v, want %v", names, want) } var migrationCount int - if err := opened.DB().QueryRow(`SELECT count(*) FROM schema_migrations`).Scan(&migrationCount); err != nil || migrationCount != 1 { + if err := opened.DB().QueryRow(`SELECT count(*) FROM schema_migrations`).Scan(&migrationCount); err != nil || migrationCount != 2 { t.Fatalf("migration count = %d, err = %v", migrationCount, err) } if err := opened.Close(); err != nil { @@ -81,7 +81,7 @@ func TestRealSQLiteInitializationAndRestart_HP_STORE_01(t *testing.T) { } reopened := openStore(t, dataDir) - if err := reopened.DB().QueryRow(`SELECT count(*) FROM schema_migrations`).Scan(&migrationCount); err != nil || migrationCount != 1 { + if err := reopened.DB().QueryRow(`SELECT count(*) FROM schema_migrations`).Scan(&migrationCount); err != nil || migrationCount != 2 { t.Fatalf("reopened migration count = %d, err = %v", migrationCount, err) } if err := reopened.Close(); err != nil {