diff --git a/internal/server/session/agent_server.go b/internal/server/session/agent_server.go index 13d9ac1..16f539c 100644 --- a/internal/server/session/agent_server.go +++ b/internal/server/session/agent_server.go @@ -275,6 +275,10 @@ func eventAppendFromWire(event *rvboxv1.CommandEvent, clientID string, generatio var immutable [32]byte copy(immutable[:], event.GetImmutableEventSha256()) result := store.EventAppend{IssueUUID: [16]byte(issue), ClientID: clientID, SessionGeneration: generation, EventSeq: event.GetEventSeq(), ObservedUnixNano: event.GetObservedAt().AsTime().UnixNano(), ReceiptUnixNano: receipt.UnixNano(), EventType: eventType(event), Compression: 1, RawLength: uint64(len(payload)), Payload: payload, ImmutableSHA256: immutable} + if lifecycle := event.GetLifecycle(); lifecycle != nil { + value := lifecycle.GetLifecycle() + result.Lifecycle, result.LifecycleRevision = &value, lifecycle.GetCommandRevision() + } if output := event.GetOutput(); output != nil { result.Stream = uint16(output.GetStream()) result.Output = true diff --git a/internal/server/store/event.go b/internal/server/store/event.go index d5b2d89..a3473f6 100644 --- a/internal/server/store/event.go +++ b/internal/server/store/event.go @@ -9,6 +9,9 @@ import ( "io/fs" "math" "path/filepath" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/domain" ) const ( @@ -44,6 +47,8 @@ type EventAppend struct { RawLength uint64 Payload []byte ImmutableSHA256 [32]byte + Lifecycle *rvboxv1.CommandLifecycle + LifecycleRevision uint64 Output bool UseCloseout bool } @@ -81,10 +86,12 @@ func (store *Store) AppendCommandEvent(ctx context.Context, event EventAppend) ( return result, err } var lastSequence, commandOutputCharged, commandCharged, clientCharged, serverCharged, closeoutRemaining uint64 + var currentLifecycle uint32 + var currentRevision uint64 var clientID string query := `SELECT commands.last_event_seq, commands.output_charged_bytes, commands.charged_bytes, commands.closeout_remaining_bytes, commands.client_id, clients.charged_bytes, -storage_counters.command_charged_bytes +storage_counters.command_charged_bytes, commands.lifecycle, commands.revision FROM commands JOIN clients ON clients.client_id = commands.client_id JOIN storage_counters ON storage_counters.singleton = 1 WHERE commands.issue_uuid = ?` arguments := []any{event.IssueUUID[:]} @@ -97,7 +104,7 @@ JOIN storage_counters ON storage_counters.singleton = 1 WHERE commands.issue_uui arguments = append(arguments, event.SessionGeneration) } err = database.QueryRowContext(ctx, query, arguments...).Scan( - &lastSequence, &commandOutputCharged, &commandCharged, &closeoutRemaining, &clientID, &clientCharged, &serverCharged) + &lastSequence, &commandOutputCharged, &commandCharged, &closeoutRemaining, &clientID, &clientCharged, &serverCharged, ¤tLifecycle, ¤tRevision) if errors.Is(err, sql.ErrNoRows) { return result, ErrCommandNotFound } @@ -117,6 +124,14 @@ JOIN storage_counters ON storage_counters.singleton = 1 WHERE commands.issue_uui if event.EventSeq != lastSequence+1 { return result, ErrEventSequenceGap } + if event.Lifecycle != nil { + if event.LifecycleRevision == 0 || event.LifecycleRevision != currentRevision { + return result, errors.New("lifecycle event revision does not match command") + } + if !domain.CanTransition(rvboxv1.CommandLifecycle(currentLifecycle), *event.Lifecycle) { + return result, domain.ValidateTransition(rvboxv1.CommandLifecycle(currentLifecycle), *event.Lifecycle) + } + } rowsCharged, indexesCharged := uint64(1), uint64(1) if store.willCreateSegment(event.IssueUUID, uint64(len(encoded))) { rowsCharged++ @@ -197,11 +212,21 @@ payload, segment_ordinal, segment_record_offset, segment_record_length, immutabl } if err == nil { var update sql.Result - update, err = tx.ExecContext(ctx, `UPDATE commands SET last_event_seq = ?, + commandUpdate := `UPDATE commands SET last_event_seq = ?, retained_compressed_bytes = retained_compressed_bytes + ?, output_charged_bytes = ?, charged_bytes = ?, -closeout_remaining_bytes = ? WHERE issue_uuid = ? AND last_event_seq = ?`, event.EventSeq, len(event.Payload), - reservation.CommandOutputCharged, reservation.CommandTotalCharged, reservation.CloseoutRemaining, - event.IssueUUID[:], lastSequence) + closeout_remaining_bytes = ?` + commandArgs := []any{event.EventSeq, len(event.Payload), reservation.CommandOutputCharged, reservation.CommandTotalCharged, reservation.CloseoutRemaining} + if event.Lifecycle != nil { + terminal := 0 + if domain.IsTerminal(*event.Lifecycle) { + terminal = 1 + } + commandUpdate = `UPDATE commands SET lifecycle = ?, terminal_time = CASE WHEN ? = 1 THEN ? ELSE terminal_time END, ` + commandUpdate[len("UPDATE commands SET "):] + commandArgs = append([]any{uint32(*event.Lifecycle), terminal, event.ReceiptUnixNano}, commandArgs...) + } + commandUpdate += ` WHERE issue_uuid = ? AND last_event_seq = ?` + commandArgs = append(commandArgs, event.IssueUUID[:], lastSequence) + update, err = tx.ExecContext(ctx, commandUpdate, commandArgs...) if err == nil { var affected int64 affected, err = update.RowsAffected() diff --git a/test/integration/clientagent/clientagent_integration_test.go b/test/integration/clientagent/clientagent_integration_test.go index eb2e0e6..e5e9af3 100644 --- a/test/integration/clientagent/clientagent_integration_test.go +++ b/test/integration/clientagent/clientagent_integration_test.go @@ -109,6 +109,13 @@ func TestWebSocketDispatchAfterReconciliation_HP_DISPATCH_03(t *testing.T) { if dispatch == nil || dispatch.GetIssueUuid() != issue.String() || dispatch.GetCommandRevision() != 1 || dispatch.GetTargetSessionGeneration() != accepted.Generation || !proto.Equal(dispatch.GetSpec(), spec) || string(dispatch.GetImmutableRequestSha256()) != string(requestHash[:]) { t.Fatalf("dispatch = %#v", dispatch) } + acceptedEnvelope, err := proto.Marshal(&rvboxv1.AgentEnvelope{SessionId: accepted.ID, SessionGeneration: accepted.Generation, Payload: &rvboxv1.AgentEnvelope_CommandAccepted{CommandAccepted: &rvboxv1.CommandAccepted{IssueUuid: issue.String(), CommandRevision: 1, Accepted: true}}}) + if err != nil { + t.Fatal(err) + } + if err := transport.Write(ctx, acceptedEnvelope); err != nil { + t.Fatal(err) + } event := &rvboxv1.CommandEvent{IssueUuid: issue.String(), EventSeq: 1, ObservedAt: timestamppb.Now(), Payload: &rvboxv1.CommandEvent_Lifecycle{Lifecycle: &rvboxv1.LifecycleChange{Lifecycle: rvboxv1.CommandLifecycle_COMMAND_RUNNING, CommandRevision: 1}}} if err := agent.SendCommandEvent(ctx, transport, accepted, event, agentproto.DefaultLimits()); err != nil { t.Fatal(err) @@ -125,4 +132,8 @@ func TestWebSocketDispatchAfterReconciliation_HP_DISPATCH_03(t *testing.T) { if err := persistence.DB().QueryRow(`SELECT count(*) FROM command_events WHERE issue_uuid = ?`, issue[:]).Scan(&persistedEvents); err != nil || persistedEvents != 1 { t.Fatalf("persisted events = %d, %v", persistedEvents, err) } + var lifecycle uint32 + if err := persistence.DB().QueryRow(`SELECT lifecycle FROM commands WHERE issue_uuid = ?`, issue[:]).Scan(&lifecycle); err != nil || lifecycle != uint32(rvboxv1.CommandLifecycle_COMMAND_RUNNING) { + t.Fatalf("command lifecycle = %d, %v", lifecycle, err) + } }