diff --git a/.github/workflows/complexity.yml b/.github/workflows/complexity.yml index c37b4dc3..b1db4dac 100644 --- a/.github/workflows/complexity.yml +++ b/.github/workflows/complexity.yml @@ -8,6 +8,7 @@ on: - '.golangci-complexity.yml' - 'go.mod' - 'go.sum' + - 'scripts/generate-go.sh' - 'scripts/go-complexity.sh' - 'web/**' - '.github/workflows/complexity.yml' @@ -20,6 +21,7 @@ on: - '.golangci-complexity.yml' - 'go.mod' - 'go.sum' + - 'scripts/generate-go.sh' - 'scripts/go-complexity.sh' - 'web/**' - '.github/workflows/complexity.yml' @@ -47,7 +49,7 @@ jobs: cache: true - name: Generate Go sources - run: go generate ./... + run: ./scripts/generate-go.sh - name: Check Go complexity uses: golangci/golangci-lint-action@v9 diff --git a/.github/workflows/dead-code.yml b/.github/workflows/dead-code.yml index 14f74723..ee58b116 100644 --- a/.github/workflows/dead-code.yml +++ b/.github/workflows/dead-code.yml @@ -8,6 +8,7 @@ on: - '.golangci-dead-code.yml' - 'go.mod' - 'go.sum' + - 'scripts/generate-go.sh' - 'scripts/go-dead-code.sh' - '.github/workflows/dead-code.yml' push: @@ -19,6 +20,7 @@ on: - '.golangci-dead-code.yml' - 'go.mod' - 'go.sum' + - 'scripts/generate-go.sh' - 'scripts/go-dead-code.sh' - '.github/workflows/dead-code.yml' @@ -45,7 +47,7 @@ jobs: cache: true - name: Generate Go sources - run: go generate ./... + run: ./scripts/generate-go.sh - name: Detect unreachable Go declarations uses: golangci/golangci-lint-action@v9 diff --git a/.github/workflows/lint.yml b/.github/workflows/lint.yml index 1061d4c2..c80a1b9b 100644 --- a/.github/workflows/lint.yml +++ b/.github/workflows/lint.yml @@ -29,7 +29,7 @@ jobs: cache: true - name: Generate Go sources - run: go generate ./... + run: ./scripts/generate-go.sh - name: Run golangci-lint uses: golangci/golangci-lint-action@v9 diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index d91f427a..36ee7c95 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -21,9 +21,9 @@ repos: hooks: - id: go-generate name: Generate ignored Go sources - entry: go generate ./... - language: system - files: ^(?:.*\.go|internal/db/.*_mapper\.xml|go\.(?:mod|sum))$ + entry: scripts/generate-go.sh + language: script + files: ^(?:.*\.go|internal/db/.*_mapper\.xml|go\.(?:mod|sum)|scripts/generate-go\.sh)$ pass_filenames: false require_serial: true @@ -34,17 +34,17 @@ repos: types: [go] - id: golangci-lint - name: Lint packages containing staged Go files + name: Lint packages containing staged Go inputs entry: scripts/pre-commit-go-lint.sh language: script - types: [go] + files: ^(?:.*\.go|internal/db/.*_mapper\.xml|go\.(?:mod|sum)|scripts/generate-go\.sh)$ require_serial: true - id: go-dead-code name: Detect unreachable Go declarations entry: scripts/go-dead-code.sh language: script - files: ^(?:.*\.go|\.golangci-dead-code\.yml|go\.(?:mod|sum)|scripts/go-dead-code\.sh)$ + files: ^(?:.*\.go|internal/db/.*_mapper\.xml|\.golangci-dead-code\.yml|go\.(?:mod|sum)|scripts/(?:generate-go|go-dead-code)\.sh)$ pass_filenames: false require_serial: true @@ -52,7 +52,7 @@ repos: name: Enforce Go cyclomatic complexity entry: scripts/go-complexity.sh language: script - files: ^(?:.*\.go|\.golangci-complexity\.yml|go\.(?:mod|sum)|scripts/go-complexity\.sh)$ + files: ^(?:.*\.go|internal/db/.*_mapper\.xml|\.golangci-complexity\.yml|go\.(?:mod|sum)|scripts/(?:generate-go|go-complexity)\.sh)$ pass_filenames: false require_serial: true diff --git a/AGENTS.md b/AGENTS.md index 9e5faebd..1800ba5e 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -153,7 +153,7 @@ ## 测试要求 - 测试组织顺序应先写失败场景,再写成功场景。 -- `*.gen.go` 不纳入版本控制;干净 checkout 在直接运行 Go 编译、测试或静态分析前先执行 `go generate ./...`。仓库标准 `just` 命令会自动完成生成。 +- `*.gen.go` 不纳入版本控制;干净 checkout 在直接运行 Go 编译、测试或静态分析前先执行 `./scripts/generate-go.sh`(内部为 `go generate ./internal/db`)。仓库标准 `just` 命令会自动完成生成。 - 修改 `web/` 下的文件后,运行 `just web-format-check`,确保 Prettier 格式门禁通过。 - 修改 Go 代码后,运行 `just lint`;该命令使用仓库根目录的 `.golangci.yml` 执行与 CI 相同的静态分析和格式检查。 - 修改 schema 或 handler 后,运行 `just test`(等价于先生成 Go 源码,再运行 `go test ./... -count=1`)。 diff --git a/Dockerfile b/Dockerfile index 0dfbaa6b..404c6d91 100644 --- a/Dockerfile +++ b/Dockerfile @@ -21,7 +21,7 @@ WORKDIR /src COPY go.mod go.sum ./ RUN go mod download COPY . . -RUN go generate ./... +RUN ./scripts/generate-go.sh RUN CGO_ENABLED=0 go build -ldflags="-s -w" -o /oma-server . # ---- 前端构建 (Bun) --------------------------------------------------------- diff --git a/docs/design/be/ccrv2/ccr-v2-epoch-design.md b/docs/design/be/ccrv2/ccr-v2-epoch-design.md index b8df2aa2..45be33af 100644 --- a/docs/design/be/ccrv2/ccr-v2-epoch-design.md +++ b/docs/design/be/ccrv2/ccr-v2-epoch-design.md @@ -164,6 +164,14 @@ update code_sessions sequence_num commit() ``` +该事务由共享的 `yourbatis.DB.Transaction` 创建;回调中的同一个事务 `Executor` +同时构造 `CodeSessionMapper`、`CodeSessionInboundEventMapper` 和 +`CodeSessionOutboundEventMapper`。`CodeSessionMapper.LockCodeSessionByExternalID` +负责日常 append 的行锁,不附加 `initializing` 状态限制;启动激活仍使用独立的 +`LockInitializingCodeSession`。方向对应的 event mapper 在锁内完成幂等单条查询和 +`INSERT ... RETURNING`,随后由 `CodeSessionMapper` 更新对应方向的 sequence。 +激活使用的 inbound 批量 insert 保持独立,不经过单条 append。 + 这样线性化语义才成立: - 如果旧 worker 先拿到锁并写完,再发生 register,则这次写入发生在抢占之前,可以接受。 diff --git a/docs/design/be/session-startup-message-delivery.md b/docs/design/be/session-startup-message-delivery.md new file mode 100644 index 00000000..91997a42 --- /dev/null +++ b/docs/design/be/session-startup-message-delivery.md @@ -0,0 +1,86 @@ +# Session 启动期消息可靠投递 + +## 问题 + +Runner 过去在 prepare 阶段读取一次 `session_events` 快照,随后才创建 sandbox 和 Code +Session。快照之后、Code Session 创建之前发送的消息虽然已经写入 `session_events`,但不会 +进入 runtime 消费的 `code_session_inbound_events`。 + +## 设计 + +`session_events` 是启动输入的唯一事实源,不增加临时 queue、watermark 或公开状态。 + +Send Events 与 Code Session activation 都先锁同一条 Session 行: + +- Send 通过 `DB.AppendSessionEvents` 锁 Session 并提交公开事件; +- activation 锁 Session 后读取完整公开历史,在同一事务中写 inbound 并将 Code Session 从 + `initializing` 切为 `active`。 + +因此只可能有两种顺序: + +1. Send 先提交,activation 随后读取到该事件并写入 inbound; +2. activation 先提交,Send 随后由现有 active realtime 路径投递当前 batch。 + +```mermaid +sequenceDiagram + participant Client + participant Send as Send transaction + participant Session as sessions row + participant Events as session_events + participant Activate as Activation transaction + participant CS as code_sessions + participant Inbound as code_session_inbound_events + + alt Send 先锁 Session + Send->>Session: SELECT FOR UPDATE + Activate->>Session: wait + Send->>Events: INSERT current batch + Send->>Send: COMMIT + Activate->>Session: acquire lock + Activate->>Events: read complete history + Activate->>Inbound: append forwardable events + Activate->>CS: initializing → active + Activate->>Activate: COMMIT + else Activate 先锁 Session + Activate->>Session: SELECT FOR UPDATE + Send->>Session: wait + Activate->>Events: read complete history + Activate->>Inbound: append forwardable events + Activate->>CS: initializing → active + Activate->>Activate: COMMIT + Send->>Session: acquire lock + Send->>Events: INSERT current batch + Send->>Inbound: active realtime delivery + end +``` + +## 激活流程 + +`Service.CreateManagedAgentCodeSession`: + +1. 创建 `initializing` Code Session; +2. 写入 `initialize` inbound; +3. 调用 `ActivateManagedAgentCodeSession`; +4. activation 事务锁定 Session 和 initializing Code Session; +5. 按 `created_at ASC, id ASC` 读取完整 `session_events`; +6. 过滤并转换可转发事件,幂等写入 inbound; +7. 将 Code Session 切为 `active` 并提交。 + +任一历史转换、inbound 写入或状态更新失败时,activation 事务整体回滚,Code Session 保持 +`initializing`。Deployment initial events 已经属于 `session_events`,无需单独交接路径。 + +## Realtime cutover + +Send 提交公开事件后始终调用 `Service.QueuePublicSessionEvents`。该方法重新读取最新 Code +Session,只有 `status == active` 时才写 inbound;不存在或仍为 `initializing` 时直接返回。 + +如果 activation 恰好在公开事件提交后、realtime 检查前完成,同一事件可能同时出现在 +activation 历史和 realtime 尝试中;现有 inbound idempotency key 会保留一份,不会重复投递。 + +## 验收 + +- Runner prepare 后、Code Session 创建前接受的消息最终进入 inbound; +- 启动期接受多条用户消息,activation 按公开历史顺序全部重放; +- activation 失败时不留下部分 inbound,也不切换为 active; +- activation 后的新 batch 只通过 realtime 路径追加; +- Deployment initial user messages 在 `initialize` 后按输入顺序进入 inbound。 diff --git a/docs/design/be/yourbatis-admin-api-keys-prototype.md b/docs/design/be/yourbatis-admin-api-keys-prototype.md index bd1c7009..1b13907d 100644 --- a/docs/design/be/yourbatis-admin-api-keys-prototype.md +++ b/docs/design/be/yourbatis-admin-api-keys-prototype.md @@ -77,9 +77,10 @@ Yourbatis runtime 与生成器均由 `go.mod` 固定到已发布的 后续升级只需更新 `go.mod` 中的模块版本,运行时和生成器会保持一致。 生成的 `*.gen.go` 不纳入版本控制,由 `.gitignore` 排除。干净 checkout 必须先执行 -`go generate ./...`;`just server`、`just test`、Go lint/死代码/复杂度门禁、pre-commit、 -GitHub Actions 和 Docker 构建都在消费生成代码前自动执行该命令。XML 和生成器版本因此成为唯一 -受版本控制的 Mapper 代码来源,避免生成输出与声明发生漂移。 +`./scripts/generate-go.sh`(内部为 `go generate ./internal/db`);`just server`、`just test`、Go +lint/死代码/复杂度门禁、pre-commit、GitHub Actions 和 Docker 构建都在消费生成代码前自动 +执行该入口。XML 和生成器版本因此成为唯一受版本控制的 Mapper 代码来源,避免生成输出与声明 +发生漂移。 ## 验证 diff --git a/docs/design/development-quality-gates.md b/docs/design/development-quality-gates.md index 69f1fbb5..d765f34d 100644 --- a/docs/design/development-quality-gates.md +++ b/docs/design/development-quality-gates.md @@ -15,6 +15,7 @@ ## 配置与执行 +- yourbatis 根据 `internal/db/*.xml` 生成的 `*.sqlmap.gen.go` 不提交;`scripts/generate-go.sh` 是统一生成入口。Go lint、死代码、复杂度、测试、开发启动和 Docker 构建在类型检查或编译前调用该入口,避免本地残留的 ignored 文件掩盖干净 checkout 中的缺失生成代码。 - `.golangci.yml` 是常规 Go lint 规则来源;复杂度和死代码等需要不同扫描范围的专项门禁使用独立的固定配置。 - `just lint` 在本地对所有 Go package(包括测试)运行相同配置。 - `.golangci-dead-code.yml` 单独启用 golangci-lint 的 `unused` 分析器并覆盖测试代码;`just dead-code` 通过 `scripts/go-dead-code.sh` 枚举当前 Go module 的仓库 package,避免本地前端依赖中的第三方 Go 示例污染结果。 @@ -36,5 +37,5 @@ just large-files just lint just dead-code just duplicates -go test ./... -count=1 +just test ``` diff --git a/go.mod b/go.mod index 79ee9c6c..5f546b70 100644 --- a/go.mod +++ b/go.mod @@ -74,5 +74,3 @@ require ( google.golang.org/genproto/googleapis/rpc v0.0.0-20260420184626-e10c466a9529 // indirect google.golang.org/grpc v1.80.0 // indirect ) - -replace ybatis => /Users/arthur/GolandProjects/ybatis diff --git a/go.sum b/go.sum index 3d6a235a..052e4609 100644 --- a/go.sum +++ b/go.sum @@ -141,8 +141,6 @@ github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= github.com/superduck-ai/e2b-go-sdk v0.0.1 h1:yW23X/fEQvaPdJHhDnp8PgnsaIixn+UTVMt94X4Kr0s= github.com/superduck-ai/e2b-go-sdk v0.0.1/go.mod h1:dWlhv18vJamYp5gQ3ktmjjrX7tUaMe2OyTJXUo8EeXY= -github.com/superduck-ai/yourbatis v0.1.0 h1:U49je+0PhoR6PVJKfR3mbNo31rr+2SzkWwBi28wndwE= -github.com/superduck-ai/yourbatis v0.1.0/go.mod h1:BlCyyT1yfU2Zxya89rDf89keXqsdcwP6Q3PFscIjvig= github.com/superduck-ai/yourbatis v0.1.1 h1:iAEi8Hrx+p6MjS8ciI2kL8ueQ7FOrFDEphnN1NP3GSc= github.com/superduck-ai/yourbatis v0.1.1/go.mod h1:BlCyyT1yfU2Zxya89rDf89keXqsdcwP6Q3PFscIjvig= github.com/tidwall/gjson v1.14.2/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk= diff --git a/internal/codesessions/managed_agent_code_session.go b/internal/codesessions/managed_agent_code_session.go index c2e65c36..e9b6a3d0 100644 --- a/internal/codesessions/managed_agent_code_session.go +++ b/internal/codesessions/managed_agent_code_session.go @@ -23,7 +23,6 @@ type ManagedAgentCreateInput struct { PermissionMode string DangerouslySkipPermissions bool Config json.RawMessage - InitialEvents []json.RawMessage } // ManagedAgentCreateResult 只在创建链路内短暂携带两份明文凭证,调用方应立即交给 @@ -66,7 +65,7 @@ func (s *Service) CreateManagedAgentCodeSession(ctx context.Context, input Manag WorkDir: strings.TrimSpace(input.WorkDir), PermissionMode: strings.TrimSpace(input.PermissionMode), Model: strings.TrimSpace(input.Model), - Status: "active", + Status: "initializing", Metadata: metadata, // OAuth-compatible token 只落 SHA-256 hash;明文仅存在于当前返回值中。 OAuthAccessTokenHash: auth.HashAPIKey(oauthAccessToken), @@ -99,7 +98,7 @@ func (s *Service) CreateManagedAgentCodeSession(ctx context.Context, input Manag if err := s.queueInitialize(ctx, record, input.Config, now); err != nil { return ManagedAgentCreateResult{}, err } - if err := s.queueInitialPublicSessionEvents(ctx, record, input.InitialEvents, now); err != nil { + if err := s.ActivateManagedAgentCodeSession(ctx, record); err != nil { return ManagedAgentCreateResult{}, err } credentialContext, err := s.db.GetCodeSessionCredentialContextForIssue( @@ -126,6 +125,76 @@ func (s *Service) CreateManagedAgentCodeSession(ctx context.Context, input Manag }, nil } +// ActivateManagedAgentCodeSession locks the owning Session, replays complete +// public history in stable order, and activates the Code Session atomically. +func (s *Service) ActivateManagedAgentCodeSession( + ctx context.Context, + codeSession db.CodeSession, +) error { + if s == nil || s.db == nil { + return db.ErrNotFound + } + return s.db.WithManagedAgentActivationTx(ctx, func(tx db.ManagedAgentActivationTx) error { + // lock session by session external id + lockedSession, err := tx.LockSessionForEvents( + ctx, + codeSession.WorkspaceUUID, + codeSession.SessionExternalID, + ) + if err != nil { + return err + } + // lock code_session by code session id + lockedCodeSession, err := tx.LockInitializingCodeSession( + ctx, + codeSession.WorkspaceUUID, + codeSession.UUID, + ) + if err != nil { + return err + } + sessionEvents, err := tx.ListSessionEventsForActivation(ctx, lockedSession) + if err != nil { + return err + } + inboundInputs := make([]db.AppendCodeSessionEventInput, 0, len(sessionEvents)) + for _, event := range sessionEvents { + if !shouldForwardPublicEventToWorker(event.EventType) { + continue + } + inbound, err := s.convertSessionEventToInbound(lockedCodeSession.ExternalID, event) + if err != nil { + return err + } + inboundInputs = append(inboundInputs, inbound) + } + if err := tx.AppendCodeSessionInboundEvents(ctx, lockedCodeSession, inboundInputs); err != nil { + return err + } + activated, err := tx.ActivateCodeSession(ctx, lockedCodeSession.UUID, time.Now().UTC()) + if err != nil { + return err + } + if !activated { + return db.ErrInvalidState + } + return nil + }) +} + +// convertSessionEventToInbound maps one public session event payload into a +// Code Session inbound append input. +func (s *Service) convertSessionEventToInbound( + codeSessionID string, + event db.SessionEvent, +) (db.AppendCodeSessionEventInput, error) { + payload, err := workerPayloadForPublicEvent(codeSessionID, event.Payload, event.ProcessedAt) + if err != nil { + return db.AppendCodeSessionEventInput{}, err + } + return newInboundEventInput(codeSessionID, payload, "public-session") +} + // TerminateManagedAgentCodeSession revokes a Code Session created for a // sandbox launch that failed before the runtime became usable. func (s *Service) TerminateManagedAgentCodeSession( diff --git a/internal/codesessions/service.go b/internal/codesessions/service.go index 0dc1d9b2..0e4084be 100644 --- a/internal/codesessions/service.go +++ b/internal/codesessions/service.go @@ -43,30 +43,6 @@ func NewServiceWithCredentials(database *db.DB, credentials *SessionCredentials, return &Service{db: database, credentials: credentials, logger: logger} } -func (s *Service) queueInitialPublicSessionEvents(ctx context.Context, codeSession db.CodeSession, payloads []json.RawMessage, now time.Time) error { - if len(payloads) == 0 { - return nil - } - workerPayloads := make([]json.RawMessage, 0, len(payloads)) - for _, raw := range payloads { - object, err := decodeJSONObject(raw) - if err != nil { - s.logger.WarnContext(ctx, "skip initial code session event", "code_session_id", codeSession.ExternalID, "error", err) - continue - } - if !forwardPublicEventToWorker(stringField(object, "type")) { - continue - } - payload, err := workerPayloadForPublicEvent(codeSession.ExternalID, raw, now) - if err != nil { - s.logger.ErrorContext(ctx, "convert initial code session event", "code_session_id", codeSession.ExternalID, "error", err) - continue - } - workerPayloads = append(workerPayloads, payload) - } - return s.QueueRawPublicSessionEvents(ctx, codeSession, workerPayloads) -} - func (s *Service) QueuePublicSessionEvents(ctx context.Context, session db.Session, events []db.SessionEvent) error { if s == nil || len(events) == 0 { return nil @@ -78,9 +54,12 @@ func (s *Service) QueuePublicSessionEvents(ctx context.Context, session db.Sessi } return err } + if codeSession.Status != "active" { + return nil + } payloads := make([]json.RawMessage, 0, len(events)) for _, event := range events { - if !forwardPublicEventToWorker(event.EventType) { + if !shouldForwardPublicEventToWorker(event.EventType) { continue } if event.EventType == "user.tool_confirmation" { @@ -216,7 +195,7 @@ func (s *Service) AppendWorkerOutputEventsForEpoch(ctx context.Context, codeSess return err } } - if event.Ephemeral || !publicWorkerOutputEvent(event.EventType) { + if event.Ephemeral || !isPublicWorkerOutputEvent(event.EventType) { continue } publicPayloads, ok, err := publicPayloadsFromWorkerEvent(codeSessionID, event, publicObject) @@ -282,7 +261,7 @@ func (s *Service) appendWorkerEvent(ctx context.Context, codeSessionID string, r if meta.EventType == "control_request" && meta.EventSubtype == "can_use_tool" { return s.handleToolPermissionRequest(ctx, codeSessionID, object, meta) } - if hiddenWorkerEvent(meta.EventType) { + if isHiddenWorkerEvent(meta.EventType) { return nil } publicPayloads, ok, err := publicPayloadsFromWorkerEvent(codeSessionID, event, object) @@ -324,15 +303,23 @@ func (s *Service) queueInitialize(ctx context.Context, codeSession db.CodeSessio } func (s *Service) appendInboundPayload(ctx context.Context, codeSessionID string, payload json.RawMessage, source string) (db.CodeSessionEvent, bool, error) { - meta, err := BuildEventMetadata(codeSessionID, "inbound", payload) + input, err := newInboundEventInput(codeSessionID, payload, source) if err != nil { return db.CodeSessionEvent{}, false, err } + return s.db.AppendCodeSessionInboundEvent(ctx, codeSessionID, input) +} + +func newInboundEventInput(codeSessionID string, payload json.RawMessage, source string) (db.AppendCodeSessionEventInput, error) { + meta, err := BuildEventMetadata(codeSessionID, "inbound", payload) + if err != nil { + return db.AppendCodeSessionEventInput{}, err + } eventID, err := ids.New("csev_") if err != nil { - return db.CodeSessionEvent{}, false, err + return db.AppendCodeSessionEventInput{}, err } - return s.db.AppendCodeSessionInboundEvent(ctx, codeSessionID, db.AppendCodeSessionEventInput{ + return db.AppendCodeSessionEventInput{ ExternalID: eventID, EventType: meta.EventType, EventSubtype: meta.EventSubtype, @@ -344,7 +331,7 @@ func (s *Service) appendInboundPayload(ctx context.Context, codeSessionID string DeliveryStatus: "queued", Source: strings.TrimSpace(source), CreatedAt: time.Now().UTC(), - }) + }, nil } func (s *Service) publishPublicPayloads(ctx context.Context, codeSessionID string, payloads []json.RawMessage) error { @@ -460,7 +447,7 @@ func (s *Service) subagentThreadMappings(ctx context.Context, codeSession db.Cod return threadByAgent, nil } -func forwardPublicEventToWorker(eventType string) bool { +func shouldForwardPublicEventToWorker(eventType string) bool { switch eventType { case "user.message", "user.interrupt", "user.tool_confirmation", "user.tool_result", "user.custom_tool_result": return true @@ -469,7 +456,7 @@ func forwardPublicEventToWorker(eventType string) bool { } } -func hiddenWorkerEvent(eventType string) bool { +func isHiddenWorkerEvent(eventType string) bool { switch eventType { case "control_request", "control_response", "control_cancel_request": return true @@ -478,7 +465,7 @@ func hiddenWorkerEvent(eventType string) bool { } } -func publicWorkerOutputEvent(eventType string) bool { +func isPublicWorkerOutputEvent(eventType string) bool { return maevents.IsWorkerOutputEvent(eventType) || maevents.IsStreamDelta(eventType) } diff --git a/internal/db/code_session_inbound_event_mapper.go b/internal/db/code_session_inbound_event_mapper.go new file mode 100644 index 00000000..301f6319 --- /dev/null +++ b/internal/db/code_session_inbound_event_mapper.go @@ -0,0 +1,36 @@ +package db + +import ( + "context" + + "github.com/google/uuid" +) + +//go:generate go tool sqlmapgen -mapper CodeSessionInboundEventMapper -sql ./code_session_inbound_event_mapper.xml -dialect postgres + +// CodeSessionInboundEventMapper contains queries whose primary table is +// code_session_inbound_events. +type CodeSessionInboundEventMapper interface { + GetCodeSessionInboundEventByIdempotencyKey( + ctx context.Context, + workspaceUUID uuid.UUID, + idempotencyKey string, + ) (codeSessionEventRow, bool, error) + + InsertCodeSessionInboundEvent( + ctx context.Context, + row codeSessionInboundEventInsertRow, + ) (codeSessionEventRow, error) + + ListExistingActivationInboundEvents( + ctx context.Context, + organizationUUID uuid.UUID, + workspaceUUID uuid.UUID, + idempotencyKeys []string, + ) ([]codeSessionInboundEventIdentityRow, error) + + InsertCodeSessionInboundEvents( + ctx context.Context, + rows []codeSessionInboundEventInsertRow, + ) (int64, error) +} diff --git a/internal/db/code_session_inbound_event_mapper.xml b/internal/db/code_session_inbound_event_mapper.xml new file mode 100644 index 00000000..0317f9da --- /dev/null +++ b/internal/db/code_session_inbound_event_mapper.xml @@ -0,0 +1,152 @@ + + + + + + uuid, + external_id, + organization_uuid, + workspace_uuid, + code_session_uuid, + code_session_external_id, + sequence_num, + event_type, + event_subtype, + payload_uuid, + request_id, + payload, + payload_hash, + idempotency_key, + delivery_status, + source, + sent_at, + delivery_worker_epoch, + received_at, + processing_at, + processed_at, + last_delivery_attempt_at, + last_delivery_update_at, + delivery_attempts, + false AS ephemeral, + created_at, + updated_at, + deleted_at + + + + + + INSERT INTO code_session_inbound_events ( + external_id, + organization_uuid, + workspace_uuid, + code_session_uuid, + code_session_external_id, + sequence_num, + event_type, + event_subtype, + payload_uuid, + request_id, + payload, + payload_hash, + idempotency_key, + delivery_status, + source, + created_at, + updated_at + ) VALUES ( + #{row.ExternalID}, + #{row.OrganizationUUID}, + #{row.WorkspaceUUID}, + #{row.CodeSessionUUID}, + #{row.CodeSessionExternalID}, + #{row.SequenceNum}, + #{row.EventType}, + #{row.EventSubtype}, + #{row.PayloadUUID}, + #{row.RequestID}, + CAST(#{row.Payload} AS jsonb), + #{row.PayloadHash}, + #{row.IdempotencyKey}, + #{row.DeliveryStatus}, + #{row.Source}, + #{row.CreatedAt}, + #{row.CreatedAt} + ) + RETURNING + + + + + + + INSERT INTO code_session_inbound_events ( + external_id, + organization_uuid, + workspace_uuid, + code_session_uuid, + code_session_external_id, + sequence_num, + event_type, + event_subtype, + payload_uuid, + request_id, + payload, + payload_hash, + idempotency_key, + delivery_status, + source, + created_at, + updated_at + ) VALUES + + ( + #{row.ExternalID}, + #{row.OrganizationUUID}, + #{row.WorkspaceUUID}, + #{row.CodeSessionUUID}, + #{row.CodeSessionExternalID}, + #{row.SequenceNum}, + #{row.EventType}, + #{row.EventSubtype}, + #{row.PayloadUUID}, + #{row.RequestID}, + CAST(#{row.Payload} AS jsonb), + #{row.PayloadHash}, + #{row.IdempotencyKey}, + #{row.DeliveryStatus}, + #{row.Source}, + #{row.CreatedAt}, + #{row.CreatedAt} + ) + + + diff --git a/internal/db/code_session_mapper.go b/internal/db/code_session_mapper.go new file mode 100644 index 00000000..ec3a4964 --- /dev/null +++ b/internal/db/code_session_mapper.go @@ -0,0 +1,44 @@ +package db + +import ( + "context" + "time" + + "github.com/google/uuid" +) + +//go:generate go tool sqlmapgen -mapper CodeSessionMapper -sql ./code_session_mapper.xml -dialect postgres + +// CodeSessionMapper contains queries whose primary table is code_sessions. +type CodeSessionMapper interface { + LockCodeSessionByExternalID( + ctx context.Context, + codeSessionExternalID string, + ) (codeSessionRow, bool, error) + + LockInitializingCodeSession( + ctx context.Context, + workspaceUUID uuid.UUID, + codeSessionUUID uuid.UUID, + ) (codeSessionRow, bool, error) + + UpdateCodeSessionInboundSequence( + ctx context.Context, + codeSessionUUID uuid.UUID, + sequenceNum int64, + now time.Time, + ) (int64, error) + + UpdateCodeSessionOutboundSequence( + ctx context.Context, + codeSessionUUID uuid.UUID, + sequenceNum int64, + now time.Time, + ) (int64, error) + + ActivateCodeSession( + ctx context.Context, + codeSessionUUID uuid.UUID, + now time.Time, + ) (int64, error) +} diff --git a/internal/db/code_session_mapper.xml b/internal/db/code_session_mapper.xml new file mode 100644 index 00000000..b1ef9a52 --- /dev/null +++ b/internal/db/code_session_mapper.xml @@ -0,0 +1,83 @@ + + + + + + uuid, + external_id, + organization_uuid, + workspace_uuid, + session_uuid, + session_external_id, + environment_uuid, + environment_external_id, + work_dir, + permission_mode, + model, + status, + metadata, + connection_status, + last_inbound_sequence_num, + last_outbound_sequence_num, + last_internal_sequence_num, + last_worker_connected_at, + last_worker_activity_at, + current_worker_epoch, + worker_lease_expires_at, + worker_registered_at, + worker_last_heartbeat_at, + worker_token_session_id, + worker_binding, + worker_status, + worker_external_metadata, + worker_requires_action_details, + created_at, + updated_at, + deleted_at + + + + + + + + UPDATE code_sessions + SET last_inbound_sequence_num = #{sequenceNum}, + updated_at = #{now} + WHERE uuid = #{codeSessionUUID} + + + + UPDATE code_sessions + SET last_outbound_sequence_num = #{sequenceNum}, + updated_at = #{now} + WHERE uuid = #{codeSessionUUID} + + + + UPDATE code_sessions + SET status = 'active', + updated_at = #{now} + WHERE uuid = #{codeSessionUUID} + AND status = 'initializing' + AND deleted_at IS NULL + + diff --git a/internal/db/code_session_outbound_event_mapper.go b/internal/db/code_session_outbound_event_mapper.go new file mode 100644 index 00000000..a63d1588 --- /dev/null +++ b/internal/db/code_session_outbound_event_mapper.go @@ -0,0 +1,24 @@ +package db + +import ( + "context" + + "github.com/google/uuid" +) + +//go:generate go tool sqlmapgen -mapper CodeSessionOutboundEventMapper -sql ./code_session_outbound_event_mapper.xml -dialect postgres + +// CodeSessionOutboundEventMapper contains queries whose primary table is +// code_session_outbound_events. +type CodeSessionOutboundEventMapper interface { + GetCodeSessionOutboundEventByIdempotencyKey( + ctx context.Context, + workspaceUUID uuid.UUID, + idempotencyKey string, + ) (codeSessionEventRow, bool, error) + + InsertCodeSessionOutboundEvent( + ctx context.Context, + row codeSessionOutboundEventInsertRow, + ) (codeSessionEventRow, error) +} diff --git a/internal/db/code_session_outbound_event_mapper.xml b/internal/db/code_session_outbound_event_mapper.xml new file mode 100644 index 00000000..ef97ad5e --- /dev/null +++ b/internal/db/code_session_outbound_event_mapper.xml @@ -0,0 +1,93 @@ + + + + + + uuid, + external_id, + organization_uuid, + workspace_uuid, + code_session_uuid, + code_session_external_id, + sequence_num, + event_type, + event_subtype, + payload_uuid, + request_id, + payload, + payload_hash, + idempotency_key, + CAST('' AS text) AS delivery_status, + source, + CAST(null AS timestamptz) AS sent_at, + CAST(null AS bigint) AS delivery_worker_epoch, + CAST(null AS timestamptz) AS received_at, + CAST(null AS timestamptz) AS processing_at, + CAST(null AS timestamptz) AS processed_at, + CAST(null AS timestamptz) AS last_delivery_attempt_at, + CAST(null AS timestamptz) AS last_delivery_update_at, + CAST(0 AS integer) AS delivery_attempts, + ephemeral, + created_at, + updated_at, + deleted_at + + + + + + INSERT INTO code_session_outbound_events ( + external_id, + organization_uuid, + workspace_uuid, + code_session_uuid, + code_session_external_id, + sequence_num, + event_type, + event_subtype, + payload_uuid, + request_id, + payload, + payload_hash, + idempotency_key, + source, + ephemeral, + created_at, + updated_at + ) VALUES ( + #{row.ExternalID}, + #{row.OrganizationUUID}, + #{row.WorkspaceUUID}, + #{row.CodeSessionUUID}, + #{row.CodeSessionExternalID}, + #{row.SequenceNum}, + #{row.EventType}, + #{row.EventSubtype}, + #{row.PayloadUUID}, + #{row.RequestID}, + CAST(#{row.Payload} AS jsonb), + #{row.PayloadHash}, + #{row.IdempotencyKey}, + #{row.Source}, + #{row.Ephemeral}, + #{row.CreatedAt}, + #{row.CreatedAt} + ) + RETURNING + + + diff --git a/internal/db/code_sessions.go b/internal/db/code_sessions.go index a7086ee3..1565e429 100644 --- a/internal/db/code_sessions.go +++ b/internal/db/code_sessions.go @@ -8,6 +8,9 @@ import ( "errors" "strings" "time" + + "github.com/samber/lo" + "github.com/superduck-ai/yourbatis" ) type CodeSession struct { @@ -276,6 +279,215 @@ func (d *DB) CreateCodeSession(ctx context.Context, input CreateCodeSessionInput }) } +func (tx ManagedAgentActivationTx) LockInitializingCodeSession( + ctx context.Context, + workspaceUUID string, + codeSessionUUID string, +) (CodeSession, error) { + parsedWorkspaceUUID, err := parseDBUUID("workspace_uuid", workspaceUUID) + if err != nil { + return CodeSession{}, err + } + parsedUUID, err := parseDBUUID("code_session_uuid", codeSessionUUID) + if err != nil { + return CodeSession{}, err + } + row, found, err := tx.codeSessionMapper.LockInitializingCodeSession(ctx, parsedWorkspaceUUID, parsedUUID) + if err != nil { + return CodeSession{}, err + } + if !found { + return CodeSession{}, ErrNotFound + } + return row.session(), nil +} + +const managedAgentActivationInboundBatchSize = 500 + +func (tx ManagedAgentActivationTx) AppendCodeSessionInboundEvents( + ctx context.Context, + codeSession CodeSession, + inputs []AppendCodeSessionEventInput, +) error { + if len(inputs) == 0 { + return nil + } + existing, err := listExistingActivationInboundEvents( + ctx, + tx.codeSessionInboundEventMapper, + codeSession, + inputs, + ) + if err != nil { + return err + } + rows, lastSequence, err := activationInboundEventInsertRows(codeSession, inputs, existing) + if err != nil { + return err + } + if len(rows) == 0 { + return nil + } + for start := 0; start < len(rows); start += managedAgentActivationInboundBatchSize { + end := min(start+managedAgentActivationInboundBatchSize, len(rows)) + inserted, err := tx.codeSessionInboundEventMapper.InsertCodeSessionInboundEvents( + ctx, + rows[start:end], + ) + if err != nil { + return err + } + if inserted != int64(end-start) { + return ErrInvalidState + } + } + codeSessionUUID, err := parseDBUUID("code_session_uuid", codeSession.UUID) + if err != nil { + return err + } + updated, err := tx.codeSessionMapper.UpdateCodeSessionInboundSequence( + ctx, + codeSessionUUID, + lastSequence, + time.Now().UTC(), + ) + if err != nil { + return err + } + if updated != 1 { + return ErrInvalidState + } + return nil +} + +func listExistingActivationInboundEvents( + ctx context.Context, + codeSessionInboundEventMapper CodeSessionInboundEventMapper, + codeSession CodeSession, + inputs []AppendCodeSessionEventInput, +) (map[string]struct{}, error) { + idempotencyKeys := lo.Uniq(lo.FilterMap( + inputs, + func(input AppendCodeSessionEventInput, _ int) (string, bool) { + return input.IdempotencyKey, input.IdempotencyKey != "" + }, + )) + if len(idempotencyKeys) == 0 { + return map[string]struct{}{}, nil + } + + organizationUUID, err := parseDBUUID("organization_uuid", codeSession.OrganizationUUID) + if err != nil { + return nil, err + } + workspaceUUID, err := parseDBUUID("workspace_uuid", codeSession.WorkspaceUUID) + if err != nil { + return nil, err + } + existing := make(map[string]struct{}, len(idempotencyKeys)) + for start := 0; start < len(idempotencyKeys); start += managedAgentActivationInboundBatchSize { + end := min(start+managedAgentActivationInboundBatchSize, len(idempotencyKeys)) + rows, err := codeSessionInboundEventMapper.ListExistingActivationInboundEvents( + ctx, + organizationUUID, + workspaceUUID, + idempotencyKeys[start:end], + ) + if err != nil { + return nil, err + } + for _, row := range rows { + if row.CodeSessionExternalID != codeSession.ExternalID { + return nil, ErrInvalidState + } + existing[row.IdempotencyKey] = struct{}{} + } + } + return existing, nil +} + +func activationInboundEventInsertRows( + codeSession CodeSession, + inputs []AppendCodeSessionEventInput, + existing map[string]struct{}, +) ([]codeSessionInboundEventInsertRow, int64, error) { + organizationUUID, err := parseDBUUID("organization_uuid", codeSession.OrganizationUUID) + if err != nil { + return nil, 0, err + } + workspaceUUID, err := parseDBUUID("workspace_uuid", codeSession.WorkspaceUUID) + if err != nil { + return nil, 0, err + } + codeSessionUUID, err := parseDBUUID("code_session_uuid", codeSession.UUID) + if err != nil { + return nil, 0, err + } + + now := time.Now().UTC() + sequence := codeSession.LastInboundSequenceNum + rows := make([]codeSessionInboundEventInsertRow, 0, len(inputs)) + seen := make(map[string]struct{}, len(inputs)) + for _, input := range inputs { + if input.RequiredWorkerEpoch != nil && codeSession.CurrentWorkerEpoch != *input.RequiredWorkerEpoch { + return nil, 0, ErrWorkerEpochMismatch + } + if input.IdempotencyKey != "" { + if _, ok := existing[input.IdempotencyKey]; ok { + continue + } + if _, ok := seen[input.IdempotencyKey]; ok { + continue + } + seen[input.IdempotencyKey] = struct{}{} + } + createdAt := input.CreatedAt + if createdAt.IsZero() { + createdAt = now + } + deliveryStatus := input.DeliveryStatus + if deliveryStatus == "" { + deliveryStatus = "queued" + } + sequence++ + rows = append(rows, codeSessionInboundEventInsertRow{ + ExternalID: input.ExternalID, + OrganizationUUID: organizationUUID, + WorkspaceUUID: workspaceUUID, + CodeSessionUUID: codeSessionUUID, + CodeSessionExternalID: codeSession.ExternalID, + SequenceNum: sequence, + EventType: input.EventType, + EventSubtype: input.EventSubtype, + PayloadUUID: input.PayloadUUID, + RequestID: input.RequestID, + Payload: []byte(input.Payload), + PayloadHash: input.PayloadHash, + IdempotencyKey: input.IdempotencyKey, + DeliveryStatus: deliveryStatus, + Source: input.Source, + CreatedAt: createdAt, + }) + } + return rows, sequence, nil +} + +func (tx ManagedAgentActivationTx) ActivateCodeSession( + ctx context.Context, + codeSessionUUID string, + now time.Time, +) (bool, error) { + parsedUUID, err := parseDBUUID("code_session_uuid", codeSessionUUID) + if err != nil { + return false, err + } + updated, err := tx.codeSessionMapper.ActivateCodeSession(ctx, parsedUUID, now) + if err != nil { + return false, err + } + return updated == 1, nil +} + // codeSessionCredentialContextSelect 查询 code session 的鉴权身份信息。 // OAuth token 鉴权和 session-ingress JWT 签发都会使用这些信息。 // JOIN 中同时校验 organization、workspace 和 session 的归属,防止跨租户查询。 @@ -811,135 +1023,125 @@ func (d *DB) appendCodeSessionEvent(ctx context.Context, direction string, codeS if input.RequiredWorkerEpoch != nil && *input.RequiredWorkerEpoch <= 0 { return CodeSessionEvent{}, false, ErrWorkerEpochMismatch } - tx, err := d.sql.BeginTxx(ctx, nil) - if err != nil { - return CodeSessionEvent{}, false, err - } - defer tx.Rollback() + var event CodeSessionEvent + var duplicate bool + err := d.mapperDB.Transaction(ctx, func(executor yourbatis.Executor) error { + codeSessionMapper := NewCodeSessionMapper(executor) + inboundMapper := NewCodeSessionInboundEventMapper(executor) + outboundMapper := NewCodeSessionOutboundEventMapper(executor) - session, err := getCodeSessionSQLX(ctx, tx, ` - select `+codeSessionColumns()+` - from code_sessions - where external_id = :external_id and deleted_at is null - for update - `, map[string]any{"external_id": codeSessionExternalID}) - if err != nil { - return CodeSessionEvent{}, false, err - } - if input.RequiredWorkerEpoch != nil && session.CurrentWorkerEpoch != *input.RequiredWorkerEpoch { - return CodeSessionEvent{}, false, ErrWorkerEpochMismatch - } - if input.IdempotencyKey != "" { - existing, err := d.getCodeSessionEventTx(ctx, tx, direction, session.WorkspaceUUID, input.IdempotencyKey) - if err == nil { - return existing, true, tx.Commit() + codeSession, found, err := codeSessionMapper.LockCodeSessionByExternalID(ctx, codeSessionExternalID) + if err != nil { + return err } - if !errors.Is(err, ErrNotFound) { - return CodeSessionEvent{}, false, err + if !found { + return ErrNotFound + } + if input.RequiredWorkerEpoch != nil && codeSession.CurrentWorkerEpoch != *input.RequiredWorkerEpoch { + return ErrWorkerEpochMismatch + } + // 有幂等键时先按 workspace 查是否已写入;命中则返回已有事件并标记 duplicate,避免重复插入。 + if input.IdempotencyKey != "" { + var existing codeSessionEventRow + if direction == "outbound" { + existing, found, err = outboundMapper.GetCodeSessionOutboundEventByIdempotencyKey( + ctx, codeSession.WorkspaceUUID, input.IdempotencyKey, + ) + } else { + existing, found, err = inboundMapper.GetCodeSessionInboundEventByIdempotencyKey( + ctx, codeSession.WorkspaceUUID, input.IdempotencyKey, + ) + } + if err != nil { + return err + } + if found { + event = existing.event() + duplicate = true + return nil + } } - } - now := input.CreatedAt - if now.IsZero() { - now = time.Now().UTC() - } - sequence := session.LastInboundSequenceNum + 1 - sequenceColumn := "last_inbound_sequence_num" - if direction == "outbound" { - sequence = session.LastOutboundSequenceNum + 1 - sequenceColumn = "last_outbound_sequence_num" - } - deliveryStatus := input.DeliveryStatus - if deliveryStatus == "" && direction == "inbound" { - deliveryStatus = "queued" - } + now := input.CreatedAt + if now.IsZero() { + now = time.Now().UTC() + } + // outbound:写入 worker 产出事件,并推进 last_outbound_sequence_num。 + if direction == "outbound" { + sequence := codeSession.LastOutboundSequenceNum + 1 + inserted, err := outboundMapper.InsertCodeSessionOutboundEvent(ctx, codeSessionOutboundEventInsertRow{ + ExternalID: input.ExternalID, + OrganizationUUID: codeSession.OrganizationUUID, + WorkspaceUUID: codeSession.WorkspaceUUID, + CodeSessionUUID: codeSession.UUID, + CodeSessionExternalID: codeSession.ExternalID, + SequenceNum: sequence, + EventType: input.EventType, + EventSubtype: input.EventSubtype, + PayloadUUID: input.PayloadUUID, + RequestID: input.RequestID, + Payload: input.Payload, + PayloadHash: input.PayloadHash, + IdempotencyKey: input.IdempotencyKey, + Source: input.Source, + Ephemeral: input.Ephemeral, + CreatedAt: now, + }) + if err != nil { + return err + } + updated, err := codeSessionMapper.UpdateCodeSessionOutboundSequence(ctx, codeSession.UUID, sequence, now) + if err != nil { + return err + } + if updated != 1 { + return ErrInvalidState + } + event = inserted.event() + return nil + } - var event CodeSessionEvent - eventArguments := map[string]any{ - "external_id": input.ExternalID, - "organization_uuid": dbUUID(session.OrganizationUUID), - "workspace_uuid": dbUUID(session.WorkspaceUUID), - "code_session_uuid": dbUUID(session.UUID), - "code_session_external_id": session.ExternalID, - "sequence_num": sequence, - "event_type": input.EventType, - "event_subtype": input.EventSubtype, - "payload_uuid": input.PayloadUUID, - "request_id": input.RequestID, - "payload": jsonArg(input.Payload), - "payload_hash": input.PayloadHash, - "idempotency_key": input.IdempotencyKey, - "source": input.Source, - "created_at": now, - } - if direction == "inbound" { - eventArguments["delivery_status"] = deliveryStatus - event, err = getCodeSessionEventSQLX(ctx, tx, ` - insert into code_session_inbound_events ( - external_id, organization_uuid, workspace_uuid, code_session_uuid, code_session_external_id, - sequence_num, event_type, event_subtype, payload_uuid, request_id, payload, - payload_hash, idempotency_key, delivery_status, source, created_at, updated_at - ) - values ( - :external_id, :organization_uuid, :workspace_uuid, :code_session_uuid, - :code_session_external_id, :sequence_num, :event_type, :event_subtype, - :payload_uuid, :request_id, CAST(:payload AS jsonb), :payload_hash, - :idempotency_key, :delivery_status, :source, :created_at, :created_at - ) - returning `+codeSessionInboundEventColumns()+` - `, eventArguments) - } else { - eventArguments["ephemeral"] = input.Ephemeral - event, err = getCodeSessionEventSQLX(ctx, tx, ` - insert into code_session_outbound_events ( - external_id, organization_uuid, workspace_uuid, code_session_uuid, code_session_external_id, - sequence_num, event_type, event_subtype, payload_uuid, request_id, payload, - payload_hash, idempotency_key, source, ephemeral, created_at, updated_at - ) - values ( - :external_id, :organization_uuid, :workspace_uuid, :code_session_uuid, - :code_session_external_id, :sequence_num, :event_type, :event_subtype, - :payload_uuid, :request_id, CAST(:payload AS jsonb), :payload_hash, - :idempotency_key, :source, :ephemeral, :created_at, :created_at - ) - returning `+codeSessionOutboundEventColumns()+` - `, eventArguments) - } + // inbound:写入投递给 worker 的事件(默认 delivery_status=queued),并推进 last_inbound_sequence_num。 + deliveryStatus := input.DeliveryStatus + if deliveryStatus == "" { + deliveryStatus = "queued" + } + sequence := codeSession.LastInboundSequenceNum + 1 + inserted, err := inboundMapper.InsertCodeSessionInboundEvent(ctx, codeSessionInboundEventInsertRow{ + ExternalID: input.ExternalID, + OrganizationUUID: codeSession.OrganizationUUID, + WorkspaceUUID: codeSession.WorkspaceUUID, + CodeSessionUUID: codeSession.UUID, + CodeSessionExternalID: codeSession.ExternalID, + SequenceNum: sequence, + EventType: input.EventType, + EventSubtype: input.EventSubtype, + PayloadUUID: input.PayloadUUID, + RequestID: input.RequestID, + Payload: input.Payload, + PayloadHash: input.PayloadHash, + IdempotencyKey: input.IdempotencyKey, + DeliveryStatus: deliveryStatus, + Source: input.Source, + CreatedAt: now, + }) + if err != nil { + return err + } + updated, err := codeSessionMapper.UpdateCodeSessionInboundSequence(ctx, codeSession.UUID, sequence, now) + if err != nil { + return err + } + if updated != 1 { + return ErrInvalidState + } + event = inserted.event() + return nil + }) if err != nil { return CodeSessionEvent{}, false, err } - if _, err := namedExecContext(ctx, tx, `update code_sessions set `+sequenceColumn+` = :sequence_num, updated_at = :now where uuid = :uuid`, map[string]any{ - "sequence_num": sequence, - "now": now, - "uuid": dbUUID(session.UUID), - }); err != nil { - return CodeSessionEvent{}, false, err - } - if err := tx.Commit(); err != nil { - return CodeSessionEvent{}, false, err - } - return event, false, nil -} - -func (d *DB) getCodeSessionEventTx(ctx context.Context, tx sqlxNamedQueryer, direction string, workspaceUUID string, idempotencyKey string) (CodeSessionEvent, error) { - arguments := map[string]any{ - "workspace_uuid": dbUUID(workspaceUUID), - "idempotency_key": idempotencyKey, - } - if direction == "outbound" { - return getCodeSessionEventSQLX(ctx, tx, ` - select `+codeSessionOutboundEventColumns()+` - from code_session_outbound_events - where workspace_uuid = :workspace_uuid and idempotency_key = :idempotency_key and deleted_at is null - limit 1 - `, arguments) - } - return getCodeSessionEventSQLX(ctx, tx, ` - select `+codeSessionInboundEventColumns()+` - from code_session_inbound_events - where workspace_uuid = :workspace_uuid and idempotency_key = :idempotency_key and deleted_at is null - limit 1 - `, arguments) + return event, duplicate, nil } func (d *DB) ListCodeSessionInternalEventsPage(ctx context.Context, params ListCodeSessionInternalEventsPageParams) ([]CodeSessionInternalEvent, bool, error) { diff --git a/internal/db/code_sessions_sqlx.go b/internal/db/code_sessions_sqlx.go index 97dd5766..aa7abf32 100644 --- a/internal/db/code_sessions_sqlx.go +++ b/internal/db/code_sessions_sqlx.go @@ -75,6 +75,49 @@ type codeSessionEventRow struct { DeletedAt *time.Time `db:"deleted_at"` } +type codeSessionInboundEventIdentityRow struct { + CodeSessionExternalID string `db:"code_session_external_id"` + IdempotencyKey string `db:"idempotency_key"` +} + +type codeSessionInboundEventInsertRow struct { + ExternalID string `db:"external_id"` + OrganizationUUID uuid.UUID `db:"organization_uuid"` + WorkspaceUUID uuid.UUID `db:"workspace_uuid"` + CodeSessionUUID uuid.UUID `db:"code_session_uuid"` + CodeSessionExternalID string `db:"code_session_external_id"` + SequenceNum int64 `db:"sequence_num"` + EventType string `db:"event_type"` + EventSubtype string `db:"event_subtype"` + PayloadUUID *string `db:"payload_uuid"` + RequestID *string `db:"request_id"` + Payload []byte `db:"payload"` + PayloadHash string `db:"payload_hash"` + IdempotencyKey string `db:"idempotency_key"` + DeliveryStatus string `db:"delivery_status"` + Source string `db:"source"` + CreatedAt time.Time `db:"created_at"` +} + +type codeSessionOutboundEventInsertRow struct { + ExternalID string `db:"external_id"` + OrganizationUUID uuid.UUID `db:"organization_uuid"` + WorkspaceUUID uuid.UUID `db:"workspace_uuid"` + CodeSessionUUID uuid.UUID `db:"code_session_uuid"` + CodeSessionExternalID string `db:"code_session_external_id"` + SequenceNum int64 `db:"sequence_num"` + EventType string `db:"event_type"` + EventSubtype string `db:"event_subtype"` + PayloadUUID *string `db:"payload_uuid"` + RequestID *string `db:"request_id"` + Payload []byte `db:"payload"` + PayloadHash string `db:"payload_hash"` + IdempotencyKey string `db:"idempotency_key"` + Source string `db:"source"` + Ephemeral bool `db:"ephemeral"` + CreatedAt time.Time `db:"created_at"` +} + type codeSessionInternalEventRow struct { UUID uuid.UUID `db:"uuid"` ExternalID string `db:"external_id"` diff --git a/internal/db/managed_agent_activation.go b/internal/db/managed_agent_activation.go new file mode 100644 index 00000000..bd1414e9 --- /dev/null +++ b/internal/db/managed_agent_activation.go @@ -0,0 +1,53 @@ +package db + +import ( + "context" + + "github.com/superduck-ai/yourbatis" +) + +// ManagedAgentActivationTx exposes the resource-scoped SQL operations used by +// the code-session service to atomically hand off startup events. +type ManagedAgentActivationTx struct { + codeSessionMapper CodeSessionMapper + codeSessionInboundEventMapper CodeSessionInboundEventMapper + sessionMapper SessionMapper + sessionEventMapper SessionEventMapper +} + +// WithManagedAgentActivationTx owns the database transaction lifecycle while +// leaving activation ordering and business decisions to the code-session service. +func (d *DB) WithManagedAgentActivationTx( + ctx context.Context, + fn func(ManagedAgentActivationTx) error, +) error { + return d.mapperDB.Transaction(ctx, func(executor yourbatis.Executor) error { + return fn(ManagedAgentActivationTx{ + codeSessionMapper: NewCodeSessionMapper(executor), + codeSessionInboundEventMapper: NewCodeSessionInboundEventMapper(executor), + sessionMapper: NewSessionMapper(executor), + sessionEventMapper: NewSessionEventMapper(executor), + }) + }) +} + +// ListSessionEventsForActivation returns complete public history in insertion order. +func (tx ManagedAgentActivationTx) ListSessionEventsForActivation( + ctx context.Context, + session Session, +) ([]SessionEvent, error) { + rows, err := tx.sessionEventMapper.ListSessionEventsForActivation( + ctx, + session.OrganizationUUID, + session.WorkspaceUUID, + session.UUID, + ) + if err != nil { + return nil, err + } + events := make([]SessionEvent, len(rows)) + for i, row := range rows { + events[i] = row.event() + } + return events, nil +} diff --git a/internal/db/session_event_mapper.go b/internal/db/session_event_mapper.go new file mode 100644 index 00000000..56f0cceb --- /dev/null +++ b/internal/db/session_event_mapper.go @@ -0,0 +1,17 @@ +package db + +import ( + "context" +) + +//go:generate go tool sqlmapgen -mapper SessionEventMapper -sql ./session_event_mapper.xml -dialect postgres + +// SessionEventMapper contains queries whose primary table is session_events. +type SessionEventMapper interface { + ListSessionEventsForActivation( + ctx context.Context, + organizationUUID string, + workspaceUUID string, + sessionUUID string, + ) ([]sessionEventRow, error) +} diff --git a/internal/db/session_event_mapper.xml b/internal/db/session_event_mapper.xml new file mode 100644 index 00000000..a25ff569 --- /dev/null +++ b/internal/db/session_event_mapper.xml @@ -0,0 +1,33 @@ + + + + + + uuid, + external_id, + organization_uuid, + workspace_uuid, + session_uuid, + session_external_id, + thread_uuid, + thread_external_id, + event_type, + payload, + processed_at, + created_at, + deleted_at + + + + diff --git a/internal/db/session_mapper.go b/internal/db/session_mapper.go new file mode 100644 index 00000000..9abdf916 --- /dev/null +++ b/internal/db/session_mapper.go @@ -0,0 +1,16 @@ +package db + +import ( + "context" +) + +//go:generate go tool sqlmapgen -mapper SessionMapper -sql ./session_mapper.xml -dialect postgres + +// SessionMapper contains queries whose primary table is sessions. +type SessionMapper interface { + LockSessionForEvents( + ctx context.Context, + workspaceUUID string, + sessionExternalID string, + ) (sessionRow, bool, error) +} diff --git a/internal/db/session_mapper.xml b/internal/db/session_mapper.xml new file mode 100644 index 00000000..6fd5abce --- /dev/null +++ b/internal/db/session_mapper.xml @@ -0,0 +1,39 @@ + + + + uuid, + external_id, + organization_uuid, + workspace_uuid, + created_by_api_key_uuid, + environment_uuid, + environment_external_id, + agent_uuid, + agent_external_id, + agent_version, + agent_snapshot, + deployment_uuid, + deployment_external_id, + title, + metadata, + vault_ids, + status, + usage, + stats, + outcome_evaluations, + created_at, + updated_at, + archived_at, + deleted_at + + + + diff --git a/internal/db/sessions.go b/internal/db/sessions.go index b66f298c..9494bf1d 100644 --- a/internal/db/sessions.go +++ b/internal/db/sessions.go @@ -634,7 +634,13 @@ func (d *DB) DeleteSessionResource(ctx context.Context, workspaceUUID string, se return tx.Commit() } -func (d *DB) AppendSessionEvents(ctx context.Context, workspaceUUID string, sessionExternalID string, events []SessionEvent) ([]SessionEvent, error) { +func (d *DB) AppendSessionEvents( + ctx context.Context, + workspaceUUID string, + sessionExternalID string, + events []SessionEvent, + outcomeEvaluations json.RawMessage, +) ([]SessionEvent, error) { tx, err := d.sql.BeginTxx(ctx, nil) if err != nil { return nil, err @@ -652,6 +658,15 @@ func (d *DB) AppendSessionEvents(ctx context.Context, workspaceUUID string, sess if err != nil { return nil, err } + if len(outcomeEvaluations) > 0 { + if _, err := getSessionSQLX(ctx, tx, setSessionOutcomeEvaluationsQuery, map[string]any{ + "workspace_uuid": dbUUID(session.WorkspaceUUID), + "session_external_id": session.ExternalID, + "outcome_evaluations": jsonArg(outcomeEvaluations), + }); err != nil { + return nil, err + } + } if err := tx.Commit(); err != nil { return nil, err } diff --git a/internal/db/sessions_sqlx.go b/internal/db/sessions_sqlx.go index 37f2dd80..ecf0c0d9 100644 --- a/internal/db/sessions_sqlx.go +++ b/internal/db/sessions_sqlx.go @@ -349,6 +349,25 @@ func sessionLookupArguments(workspaceUUID string, sessionExternalID string) map[ } } +func (tx ManagedAgentActivationTx) LockSessionForEvents( + ctx context.Context, + workspaceUUID string, + sessionExternalID string, +) (Session, error) { + row, found, err := tx.sessionMapper.LockSessionForEvents( + ctx, + workspaceUUID, + sessionExternalID, + ) + if err != nil { + return Session{}, err + } + if !found { + return Session{}, ErrNotFound + } + return row.session(), nil +} + func getSessionSQLX( ctx context.Context, database sqlxNamedQueryer, diff --git a/internal/db/yourbatis_mappers_test.go b/internal/db/yourbatis_mappers_test.go new file mode 100644 index 00000000..d6943b49 --- /dev/null +++ b/internal/db/yourbatis_mappers_test.go @@ -0,0 +1,182 @@ +package db + +import ( + "context" + "fmt" + "strings" + "testing" + "time" + + "github.com/google/uuid" + yourbatis "github.com/superduck-ai/yourbatis" +) + +func TestTableMappersBuildDynamicQueries(t *testing.T) { + organizationUUID := uuid.MustParse("11111111-1111-4111-8111-111111111111") + workspaceUUID := uuid.MustParse("22222222-2222-4222-8222-222222222222") + codeSessionUUID := uuid.MustParse("33333333-3333-4333-8333-333333333333") + + t.Run("single code session event append", func(t *testing.T) { + lockBound := buildCodeSessionMapperLockCodeSessionByExternalID( + yourbatis.DialectPostgres, + "cse_test", + ) + assertMapperSQLContains(t, lockBound, "WHERE external_id = $1 AND deleted_at IS NULL FOR UPDATE") + if strings.Contains(lockBound.SQL, "status = 'initializing'") { + t.Fatalf("daily append lock unexpectedly restricts status: %s", lockBound.SQL) + } + + inboundBound := buildCodeSessionInboundEventMapperGetCodeSessionInboundEventByIdempotencyKey( + yourbatis.DialectPostgres, + workspaceUUID, + "idem-inbound", + ) + assertMapperSQLContains(t, inboundBound, "workspace_uuid = $1 AND idempotency_key = $2 AND deleted_at IS NULL") + + outboundBound := buildCodeSessionOutboundEventMapperGetCodeSessionOutboundEventByIdempotencyKey( + yourbatis.DialectPostgres, + workspaceUUID, + "idem-outbound", + ) + assertMapperSQLContains(t, outboundBound, "workspace_uuid = $1 AND idempotency_key = $2 AND deleted_at IS NULL") + }) + + t.Run("activation code session lock", func(t *testing.T) { + bound := buildCodeSessionMapperLockInitializingCodeSession( + yourbatis.DialectPostgres, + workspaceUUID, + codeSessionUUID, + ) + assertMapperSQLContains(t, bound, "WHERE workspace_uuid = $1 AND uuid = $2") + }) + + t.Run("idempotency keys", func(t *testing.T) { + bound := buildCodeSessionInboundEventMapperListExistingActivationInboundEvents( + yourbatis.DialectPostgres, + organizationUUID, + workspaceUUID, + []string{"idem-one", "idem-two"}, + ) + assertMapperSQLContains(t, bound, "idempotency_key IN ( $3 , $4 )") + assertMapperArgumentNames(t, bound, []string{ + "organizationUUID", + "workspaceUUID", + "idempotencyKey", + "idempotencyKey", + }) + }) +} + +func TestTableMappersBuildWrites(t *testing.T) { + organizationUUID := uuid.MustParse("11111111-1111-4111-8111-111111111111") + workspaceUUID := uuid.MustParse("22222222-2222-4222-8222-222222222222") + codeSessionUUID := uuid.MustParse("33333333-3333-4333-8333-333333333333") + createdAt := time.Date(2026, time.August, 2, 12, 0, 0, 0, time.UTC) + row := codeSessionInboundEventInsertRow{ + ExternalID: "evt_test", + OrganizationUUID: organizationUUID, + WorkspaceUUID: workspaceUUID, + CodeSessionUUID: codeSessionUUID, + CodeSessionExternalID: "codeses_test", + SequenceNum: 1, + EventType: "user", + EventSubtype: "message", + Payload: []byte(`{"type":"user"}`), + PayloadHash: "hash", + IdempotencyKey: "idem", + DeliveryStatus: "queued", + Source: "session_event", + CreatedAt: createdAt, + } + bound := buildCodeSessionInboundEventMapperInsertCodeSessionInboundEvents( + yourbatis.DialectPostgres, + []codeSessionInboundEventInsertRow{row, row}, + ) + assertMapperSQLContains(t, bound, "CAST($11 AS jsonb)") + assertMapperSQLContains(t, bound, "CAST($28 AS jsonb)") + if len(bound.Args) != 34 { + t.Fatalf("inbound batch argument count = %d, want 34", len(bound.Args)) + } + + singleInboundBound := buildCodeSessionInboundEventMapperInsertCodeSessionInboundEvent( + yourbatis.DialectPostgres, + row, + ) + assertMapperSQLContains(t, singleInboundBound, "INSERT INTO code_session_inbound_events") + assertMapperSQLContains(t, singleInboundBound, "RETURNING uuid, external_id") + + outboundRow := codeSessionOutboundEventInsertRow{ + ExternalID: "evt_outbound_test", + OrganizationUUID: organizationUUID, + WorkspaceUUID: workspaceUUID, + CodeSessionUUID: codeSessionUUID, + CodeSessionExternalID: "codeses_test", + SequenceNum: 1, + EventType: "assistant", + EventSubtype: "message", + Payload: []byte(`{"type":"assistant"}`), + PayloadHash: "hash-outbound", + IdempotencyKey: "idem-outbound", + Source: "worker", + Ephemeral: true, + CreatedAt: createdAt, + } + singleOutboundBound := buildCodeSessionOutboundEventMapperInsertCodeSessionOutboundEvent( + yourbatis.DialectPostgres, + outboundRow, + ) + assertMapperSQLContains(t, singleOutboundBound, "INSERT INTO code_session_outbound_events") + assertMapperSQLContains(t, singleOutboundBound, "RETURNING uuid, external_id") +} + +func TestListExistingActivationInboundEventsBatchesLookup(t *testing.T) { + executor := newMapperTestExecutor(t, mapperTestResponse{}) + inputs := make([]AppendCodeSessionEventInput, managedAgentActivationInboundBatchSize+1) + for index := range inputs { + inputs[index].IdempotencyKey = fmt.Sprintf("idem-%d", index) + } + + _, err := listExistingActivationInboundEvents( + context.Background(), + NewCodeSessionInboundEventMapper(executor), + CodeSession{ + ExternalID: "cse_test", + OrganizationUUID: "11111111-1111-4111-8111-111111111111", + WorkspaceUUID: "22222222-2222-4222-8222-222222222222", + }, + inputs, + ) + if err != nil { + t.Fatalf("list existing activation inbound events: %v", err) + } + if executor.queryCallCount != 2 { + t.Fatalf("activation idempotency lookup calls = %d, want 2", executor.queryCallCount) + } +} + +func assertMapperSQLContains( + t *testing.T, + bound yourbatis.BoundSQL, + want string, +) { + t.Helper() + compact := strings.Join(strings.Fields(bound.SQL), " ") + if !strings.Contains(compact, want) { + t.Fatalf("SQL does not contain %q:\n%s", want, compact) + } +} + +func assertMapperArgumentNames( + t *testing.T, + bound yourbatis.BoundSQL, + want []string, +) { + t.Helper() + got := make([]string, len(bound.Args)) + for index, argument := range bound.Args { + got[index] = argument.Name + } + if strings.Join(got, ",") != strings.Join(want, ",") { + t.Fatalf("argument names = %v, want %v", got, want) + } +} diff --git a/internal/environments/runner.go b/internal/environments/runner.go index d46de5e9..0ae031ba 100644 --- a/internal/environments/runner.go +++ b/internal/environments/runner.go @@ -77,7 +77,6 @@ type Runner struct { type managedAgentLaunchPreparation struct { Session db.Session - InitialEvents []json.RawMessage SessionConfig json.RawMessage WorkDir string Title string @@ -483,10 +482,6 @@ func (r *Runner) prepareManagedAgentLaunch( if err != nil { return nil, err } - events, err := r.sessionEventPayloads(ctx, session) - if err != nil { - return nil, err - } runtimeSkills, err := r.resolveRuntimeSkills(ctx, session) if err != nil { return nil, err @@ -503,7 +498,6 @@ func (r *Runner) prepareManagedAgentLaunch( } return &managedAgentLaunchPreparation{ Session: session, - InitialEvents: events, SessionConfig: sessionConfig, WorkDir: workDir, Title: title, @@ -526,7 +520,6 @@ func (r *Runner) createManagedAgentRuntimeLaunch( PermissionMode: "bypassPermissions", DangerouslySkipPermissions: true, Config: preparation.SessionConfig, - InitialEvents: preparation.InitialEvents, }) if err != nil { return managedAgentRuntimeLaunch{}, err @@ -805,31 +798,6 @@ func (r *Runner) replaceRuntimeSkillArchives( return nil } -func (r *Runner) sessionEventPayloads(ctx context.Context, session db.Session) ([]json.RawMessage, error) { - var out []json.RawMessage - var cursor *db.SessionEventPageCursor - for { - events, hasMore, err := r.db.ListSessionEventsPage(ctx, db.ListSessionEventsPageParams{ - WorkspaceUUID: session.WorkspaceUUID, - SessionExternalID: session.ExternalID, - Limit: 100, - Cursor: cursor, - Order: "asc", - }) - if err != nil { - return nil, err - } - for _, event := range events { - out = append(out, append(json.RawMessage(nil), event.Payload...)) - } - if !hasMore || len(events) == 0 { - return out, nil - } - last := events[len(events)-1] - cursor = &db.SessionEventPageCursor{CreatedAt: last.CreatedAt, UUID: last.UUID} - } -} - func sessionIDFromEnvironmentWork(work db.EnvironmentWork) (string, bool) { var data struct { Type string `json:"type"` diff --git a/internal/sessions/code_event_bridge.go b/internal/sessions/code_event_bridge.go index da74172b..c1d23032 100644 --- a/internal/sessions/code_event_bridge.go +++ b/internal/sessions/code_event_bridge.go @@ -13,7 +13,7 @@ import ( ) func (h *Handler) appendAndBroadcastInternal(r *http.Request, sessionID string, events []db.SessionEvent) { - created, err := h.db.AppendSessionEvents(r.Context(), workspaceUUIDFromRequest(r), sessionID, events) + created, err := h.db.AppendSessionEvents(r.Context(), workspaceUUIDFromRequest(r), sessionID, events, nil) if err != nil { h.logger.ErrorContext(r.Context(), "append internal session events", "session_id", sessionID, "error", err) return diff --git a/internal/sessions/service.go b/internal/sessions/service.go index d397cb84..95c44e29 100644 --- a/internal/sessions/service.go +++ b/internal/sessions/service.go @@ -596,16 +596,24 @@ func (h *Handler) sendEventsRoute(w http.ResponseWriter, r *http.Request) { now := time.Now().UTC() events := make([]db.SessionEvent, 0, len(inputs)) var outcomesChanged bool + normalizedSession := session for _, raw := range inputs { - event, changed, err := h.normalizeInputEvent(r.Context(), session, raw, now) + event, outcomes, changed, err := normalizeInputEvent(normalizedSession, raw, now) if err != nil { writeBadRequest(w, r, err) return } - outcomesChanged = outcomesChanged || changed + if changed { + normalizedSession.OutcomeEvaluations = outcomes + outcomesChanged = true + } events = append(events, event) } - created, err := h.db.AppendSessionEvents(r.Context(), session.WorkspaceUUID, session.ExternalID, events) + var outcomeEvaluations json.RawMessage + if outcomesChanged { + outcomeEvaluations = normalizedSession.OutcomeEvaluations + } + created, err := h.db.AppendSessionEvents(r.Context(), session.WorkspaceUUID, session.ExternalID, events, outcomeEvaluations) if err != nil { if errors.Is(err, db.ErrInvalidState) { writeBadRequest(w, r, errors.New("archived sessions do not accept new events")) diff --git a/internal/sessions/service_helpers.go b/internal/sessions/service_helpers.go index 4bb340d2..b7da326e 100644 --- a/internal/sessions/service_helpers.go +++ b/internal/sessions/service_helpers.go @@ -1,7 +1,6 @@ package sessions import ( - "context" "encoding/json" "errors" "fmt" @@ -223,21 +222,25 @@ func (h *Handler) resourceFromFields( }, nil } -func (h *Handler) normalizeInputEvent(ctx context.Context, session db.Session, raw json.RawMessage, now time.Time) (db.SessionEvent, bool, error) { +func normalizeInputEvent( + session db.Session, + raw json.RawMessage, + now time.Time, +) (db.SessionEvent, json.RawMessage, bool, error) { var payload map[string]any if err := json.Unmarshal(raw, &payload); err != nil { - return db.SessionEvent{}, false, errors.New("event must be an object") + return db.SessionEvent{}, nil, false, errors.New("event must be an object") } eventType, _ := payload["type"].(string) if !allowedPublicEventType(eventType) { - return db.SessionEvent{}, false, errors.New("event type is not accepted by this endpoint") + return db.SessionEvent{}, nil, false, errors.New("event type is not accepted by this endpoint") } if err := validatePublicInputEvent(eventType, payload); err != nil { - return db.SessionEvent{}, false, err + return db.SessionEvent{}, nil, false, err } eventID, err := ids.New("sevt_") if err != nil { - return db.SessionEvent{}, false, err + return db.SessionEvent{}, nil, false, err } payload["id"] = eventID payload["processed_at"] = now.Format(time.RFC3339) @@ -253,7 +256,7 @@ func (h *Handler) normalizeInputEvent(ctx context.Context, session db.Session, r if outcomeID == "" { outcomeID, err = ids.New("outc_") if err != nil { - return db.SessionEvent{}, false, err + return db.SessionEvent{}, nil, false, err } payload["outcome_id"] = outcomeID } @@ -262,21 +265,19 @@ func (h *Handler) normalizeInputEvent(ctx context.Context, session db.Session, r maxIterations = int(rawMax) } if maxIterations > 20 { - return db.SessionEvent{}, false, errors.New("max_iterations must be at most 20") + return db.SessionEvent{}, nil, false, errors.New("max_iterations must be at most 20") } payload["max_iterations"] = maxIterations outcomes, err := appendOutcomeEvaluation(session.OutcomeEvaluations, outcomeID, maxIterations, now) if err != nil { - return db.SessionEvent{}, false, err - } - if _, err := h.db.SetSessionOutcomeEvaluations(ctx, session.WorkspaceUUID, session.ExternalID, outcomes); err != nil { - return db.SessionEvent{}, false, err + return db.SessionEvent{}, nil, false, err } + session.OutcomeEvaluations = outcomes outcomesChanged = true } payloadRaw, err := httpapi.MarshalRaw(payload) if err != nil { - return db.SessionEvent{}, false, err + return db.SessionEvent{}, nil, false, err } return db.SessionEvent{ UUID: uuid.NewString(), @@ -290,7 +291,7 @@ func (h *Handler) normalizeInputEvent(ctx context.Context, session db.Session, r Payload: payloadRaw, ProcessedAt: now, CreatedAt: now, - }, outcomesChanged, nil + }, session.OutcomeEvaluations, outcomesChanged, nil } func allowedPublicEventType(eventType string) bool { diff --git a/justfile b/justfile index f9e60e39..b43a4136 100644 --- a/justfile +++ b/justfile @@ -8,7 +8,7 @@ help: # Generate ignored Go sources required by builds, tests, and static analysis. generate: - go generate ./... + ./scripts/generate-go.sh # Create the gitignored Docker Compose runtime config without overwriting an existing secret-bearing file. init-compose-config: diff --git a/scripts/generate-go.sh b/scripts/generate-go.sh new file mode 100755 index 00000000..096edaac --- /dev/null +++ b/scripts/generate-go.sh @@ -0,0 +1,7 @@ +#!/usr/bin/env bash +set -euo pipefail + +repo_root="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +cd "$repo_root" + +go generate ./internal/db diff --git a/scripts/go-complexity.sh b/scripts/go-complexity.sh index 2ad95178..9f9d7c97 100755 --- a/scripts/go-complexity.sh +++ b/scripts/go-complexity.sh @@ -9,6 +9,8 @@ if ! command -v golangci-lint >/dev/null 2>&1; then exit 1 fi +./scripts/generate-go.sh + packages=() while IFS= read -r directory; do [[ "$directory" == "$repo_root/web/node_modules/"* ]] && continue diff --git a/scripts/go-dead-code.sh b/scripts/go-dead-code.sh index 17fcb253..4fb66773 100755 --- a/scripts/go-dead-code.sh +++ b/scripts/go-dead-code.sh @@ -9,6 +9,8 @@ if ! command -v golangci-lint >/dev/null 2>&1; then exit 1 fi +./scripts/generate-go.sh + packages=() while IFS= read -r directory; do [[ "$directory" == "$repo_root/web/node_modules/"* ]] && continue diff --git a/scripts/pre-commit-go-lint.sh b/scripts/pre-commit-go-lint.sh index a56f6c3f..ab2ee23b 100755 --- a/scripts/pre-commit-go-lint.sh +++ b/scripts/pre-commit-go-lint.sh @@ -1,18 +1,30 @@ #!/usr/bin/env bash set -euo pipefail +repo_root="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +cd "$repo_root" + if ! command -v golangci-lint >/dev/null 2>&1; then echo 'golangci-lint is required for staged Go files. Install the version used by .github/workflows/lint.yml.' >&2 exit 1 fi +./scripts/generate-go.sh + packages=() for file in "$@"; do [[ -f "$file" ]] || continue - directory="$(dirname "$file")" - package="./${directory}" - [[ "$directory" == '.' ]] && package='.' + case "$file" in + go.mod|go.sum|scripts/generate-go.sh) + package='./...' + ;; + *) + directory="$(dirname "$file")" + package="./${directory}" + [[ "$directory" == '.' ]] && package='.' + ;; + esac seen=false for existing in "${packages[@]:-}"; do diff --git a/scripts/restart-server.sh b/scripts/restart-server.sh index 467e859b..e993f948 100755 --- a/scripts/restart-server.sh +++ b/scripts/restart-server.sh @@ -60,6 +60,7 @@ stop_listeners() { fi } +"$REPO_ROOT/scripts/generate-go.sh" stop_listeners echo "Starting claude-api-server with $CONFIG_FILE in foreground" diff --git a/tests/deployments_api_test.go b/tests/deployments_api_test.go index aaea3c98..d9a0aa57 100644 --- a/tests/deployments_api_test.go +++ b/tests/deployments_api_test.go @@ -240,6 +240,41 @@ func TestDeploymentsAPI(t *testing.T) { ) }) + t.Run("success initial user messages replay in order", func(t *testing.T) { + agent := createAgent(t, app, `{"model":"claude-opus-4-6","name":"deployment-initial-history-agent"}`) + defer cleanupAgentRows(t, app.db, agent.ID) + env := createEnvironment(t, app, `{"name":"deployment-initial-history-env"}`) + defer cleanupEnvironmentRows(t, app.db, env.ID) + deployment := createDeployment(t, app, `{ + "agent":`+quoteJSON(agent.ID)+`, + "environment_id":`+quoteJSON(env.ID)+`, + "name":"deployment initial history", + "initial_events":[ + {"type":"user.message","content":[{"type":"text","text":"deployment first"}]}, + {"type":"system.message","content":[{"type":"text","text":"public only"}]}, + {"type":"user.message","content":[{"type":"text","text":"deployment second"}]} + ] + }`) + defer cleanupDeploymentRows(t, app, deployment.ID) + run := runDeployment(t, app, deployment.ID) + if run.SessionID == nil || *run.SessionID == "" { + t.Fatalf("deployment initial history Session ID = nil: %+v", run) + } + defer deleteSession(t, app, *run.SessionID) + + codeSessionID := launchLocalCodeSession(t, app, *run.SessionID) + inbound, err := app.db.ListQueuedCodeSessionInboundEvents(context.Background(), codeSessionID) + if err != nil { + t.Fatalf("list deployment startup inbound: %v", err) + } + if len(inbound) != 3 || + inbound[0].EventSubtype != "initialize" || + !strings.Contains(string(inbound[1].Payload), "deployment first") || + !strings.Contains(string(inbound[2].Payload), "deployment second") { + t.Fatalf("deployment startup inbound = %#v, want initialize, first, second", inbound) + } + }) + t.Run("success lifecycle manual run session events and run filters", func(t *testing.T) { agent := createAgent(t, app, `{"model":"claude-opus-4-6","name":"deployments-api-agent"}`) defer cleanupAgentRows(t, app.db, agent.ID) diff --git a/tests/environments_runner_cloud_test.go b/tests/environments_runner_cloud_test.go index 65217882..e0610fba 100644 --- a/tests/environments_runner_cloud_test.go +++ b/tests/environments_runner_cloud_test.go @@ -1,6 +1,7 @@ package tests import ( + "bytes" "context" "encoding/json" "errors" @@ -318,6 +319,77 @@ func TestEnvironmentRunnerLaunchesManagedAgentCloudSession(t *testing.T) { } } +// TestEnvironmentRunnerDeliversMessageAcceptedBeforeCodeSessionCreation 对应 #189: +// Session 已可发消息,但 Code Session 在 Runner 后半段才创建。 +// 在 provider.Create 时发送(prepare 已完成、Code Session 尚不存在): +// 消息应进入 session_events,且 Runner 结束后也应出现在 code_session_inbound_events +// (initialize 之后)。修复前该窗口内只有 session_events,inbound 往往只有 initialize。 +func TestEnvironmentRunnerDeliversMessageAcceptedBeforeCodeSessionCreation(t *testing.T) { + ctx := context.Background() + cfg, err := config.Load() + if err != nil { + t.Fatalf("load config: %v", err) + } + app := newTestAppWithStore(t, &cfg, newFakeStore("runner-startup-message-bucket")) + defer app.close() + cfg.CodeSession.SandboxAPIBaseURL = app.baseURL + + agent := createAgent(t, app, `{"model":"claude-opus-4-8","name":"Runner Startup Message Agent"}`) + defer archiveAgent(t, app, agent.ID) + env := createEnvironment(t, app, `{"name":"runner-startup-message-`+strings.ReplaceAll(time.Now().Format("150405.000000000"), ".", "")+`"}`) + defer cleanupEnvironmentRows(t, app.db, env.ID) + session := createSession(t, app, `{"agent":`+quoteJSON(agent.ID)+`,"environment_id":`+quoteJSON(env.ID)+`}`) + defer deleteSession(t, app, session.ID) + + const prompt = "startup-window message must reach inbound" + ids := getDefaultDBIDs(t, app.db) + var acceptedEventID, interruptEventID string + provider := &recordingRunnerProvider{ + sandboxID: "sandbox-runner-startup-message", + beforeCreate: func() { + if _, err := app.db.GetCodeSessionBySessionExternalID(ctx, ids.WorkspaceUUID, session.ID); !errors.Is(err, db.ErrNotFound) { + t.Fatalf("code session at provider.Create = %v, want ErrNotFound", err) + } + sent := sendSessionEvents(t, app, session.ID, `{"events":[{"type":"user.message","content":[{"type":"text","text":`+quoteJSON(prompt)+`}]}]}`, defaultTestKey) + if len(sent.Data) != 1 { + t.Fatalf("accepted events = %#v, want one", sent.Data) + } + acceptedEventID = sessionEventStringField(t, sent.Data[0], "id") + interrupt := sendSessionEvents(t, app, session.ID, `{"events":[{"type":"user.interrupt"}]}`, defaultTestKey) + if len(interrupt.Data) != 1 { + t.Fatalf("accepted interrupt = %#v, want one", interrupt.Data) + } + interruptEventID = sessionEventStringField(t, interrupt.Data[0], "id") + }, + } + + processed, err := newManagedAgentRunner(t, app, provider, cfg).RunOnce(ctx, "runner-startup-message-test") + if err != nil || !processed { + t.Fatalf("RunOnce() = (%t, %v), want success", processed, err) + } + + publicEvents := listSessionEvents(t, app, session.ID, "order=asc&limit=100", defaultTestKey) + if !eventPageContains(publicEvents, acceptedEventID) { + t.Fatalf("session_events missing accepted event %q: %+v", acceptedEventID, publicEvents.Data) + } + codeSession, err := app.db.GetCodeSessionBySessionExternalID(ctx, ids.WorkspaceUUID, session.ID) + if err != nil { + t.Fatalf("load Code Session after runner startup: %v", err) + } + queued, err := app.db.ListQueuedCodeSessionInboundEvents(ctx, codeSession.ExternalID) + if err != nil { + t.Fatalf("list queued inbound events: %v", err) + } + if len(queued) != 3 || + queued[0].EventSubtype != "initialize" || + queued[1].EventType != "user" || + !bytes.Contains(queued[1].Payload, []byte(acceptedEventID)) || + !bytes.Contains(queued[1].Payload, []byte(prompt)) || + !bytes.Contains(queued[2].Payload, []byte(interruptEventID)) { + t.Fatalf("inbound = %#v, want initialize, accepted user message, interrupt", queued) + } +} + func TestEnvironmentRunnerPackageProvisioning(t *testing.T) { t.Run("failure terminates sandbox before manager startup", func(t *testing.T) { provider, processed, err := runPackageEnvironment(t, packageRunnerCase{commandErr: errors.New("gem install failed")}) diff --git a/tests/sessions_api_test.go b/tests/sessions_api_test.go index 6c1f14d2..5557d1b5 100644 --- a/tests/sessions_api_test.go +++ b/tests/sessions_api_test.go @@ -685,6 +685,200 @@ func TestSessionClaudeCodeTaskEventsMapToCanonicalThreads(t *testing.T) { } } +func TestManagedAgentActivationReplaysStartupHistory(t *testing.T) { + ctx := context.Background() + app := newTestAppWithStore(t, nil, newFakeStore("sessions-managed-agent-activation-bucket")) + defer app.close() + + agent := createAgent(t, app, `{"model":"claude-opus-4-6","name":"sessions-managed-agent-activation-agent"}`) + defer cleanupAgentRows(t, app.db, agent.ID) + env := createEnvironment(t, app, `{"name":"sessions-managed-agent-activation-env"}`) + defer cleanupEnvironmentRows(t, app.db, env.ID) + session := createSession(t, app, `{"agent":`+quoteJSON(agent.ID)+`,"environment_id":`+quoteJSON(env.ID)+`}`) + defer deleteSession(t, app, session.ID) + codeSessionID := launchLocalCodeSession(t, app, session.ID) + + if _, err := app.db.Pool.Exec(ctx, ` + update code_sessions + set status = 'initializing' + where external_id = $1 + `, codeSessionID); err != nil { + t.Fatalf("set code session initializing: %v", err) + } + + codeSession, err := app.db.GetCodeSession(ctx, codeSessionID) + if err != nil { + t.Fatalf("load initializing code session: %v", err) + } + codeSessionService := codesessions.NewServiceWithCredentials(app.db, app.credentials, nil) + + accepted := sendSessionEvents(t, app, session.ID, `{"events":[ + {"type":"user.message","content":[{"type":"text","text":"startup message one"}]}, + {"type":"user.message","content":[{"type":"text","text":"startup message two"}]} + ]}`, defaultTestKey) + if len(accepted.Data) != 2 { + t.Fatalf("accepted session events = %#v, want two", accepted.Data) + } + + if err := codeSessionService.ActivateManagedAgentCodeSession(ctx, codeSession); err != nil { + t.Fatalf("activate with startup history: %v", err) + } + codeSession, err = app.db.GetCodeSession(ctx, codeSessionID) + if err != nil { + t.Fatalf("reload activated code session: %v", err) + } + if codeSession.Status != "active" { + t.Fatalf("status after activation = %q, want active", codeSession.Status) + } + inbound, err := app.db.ListQueuedCodeSessionInboundEvents(ctx, codeSessionID) + if err != nil { + t.Fatalf("list inbound after activation: %v", err) + } + if len(inbound) != 3 || + !bytes.Contains(inbound[1].Payload, []byte("startup message one")) || + !bytes.Contains(inbound[2].Payload, []byte("startup message two")) { + t.Fatalf("inbound after activation = %#v, want initialize and both startup messages", inbound) + } + + sent := sendSessionEvents(t, app, session.ID, `{"events":[{"type":"user.message","content":[{"type":"text","text":"post-cutover message"}]}]}`, defaultTestKey) + if len(sent.Data) != 1 { + t.Fatalf("post-cutover events = %#v, want one", sent.Data) + } + inbound, err = app.db.ListQueuedCodeSessionInboundEvents(ctx, codeSessionID) + if err != nil { + t.Fatalf("list post-cutover inbound: %v", err) + } + if len(inbound) != 4 || !bytes.Contains(inbound[3].Payload, []byte("post-cutover message")) { + t.Fatalf("post-cutover inbound = %#v, want realtime message after startup history", inbound) + } +} + +func TestManagedAgentActivationPreservesLargeHistoryOrder(t *testing.T) { + ctx := context.Background() + app := newTestAppWithStore(t, nil, newFakeStore("sessions-managed-agent-large-history-bucket")) + defer app.close() + + agent := createAgent(t, app, `{"model":"claude-opus-4-6","name":"sessions-managed-agent-large-history-agent"}`) + defer cleanupAgentRows(t, app.db, agent.ID) + env := createEnvironment(t, app, `{"name":"sessions-managed-agent-large-history-env"}`) + defer cleanupEnvironmentRows(t, app.db, env.ID) + sessionResponse := createSession(t, app, `{"agent":`+quoteJSON(agent.ID)+`,"environment_id":`+quoteJSON(env.ID)+`}`) + defer deleteSession(t, app, sessionResponse.ID) + codeSessionID := launchLocalCodeSession(t, app, sessionResponse.ID) + if _, err := app.db.Pool.Exec(ctx, `update code_sessions set status = 'initializing' where external_id = $1`, codeSessionID); err != nil { + t.Fatalf("set Code Session initializing: %v", err) + } + + session, err := app.db.GetSession(ctx, getDefaultDBIDs(t, app.db).WorkspaceUUID, sessionResponse.ID) + if err != nil { + t.Fatalf("load Session: %v", err) + } + codeSession, err := app.db.GetCodeSession(ctx, codeSessionID) + if err != nil { + t.Fatalf("load Code Session: %v", err) + } + const eventCount = 501 + createdAt := time.Now().UTC() + uuidPrefix := uuid.New() + eventIDs := make([]string, 0, eventCount) + events := make([]db.SessionEvent, 0, eventCount) + for i := range eventCount { + eventID := fmt.Sprintf("sevt_large_history_%03d_%s", i, strings.TrimPrefix(codeSessionID, "cse_")) + eventUUID := uuidPrefix + eventUUID[14] = byte((eventCount - i) >> 8) + eventUUID[15] = byte(eventCount - i) + eventIDs = append(eventIDs, eventID) + events = append(events, db.SessionEvent{ + UUID: eventUUID.String(), + ExternalID: eventID, + EventType: "user.interrupt", + Payload: json.RawMessage(`{"type":"user.interrupt","id":` + quoteJSON(eventID) + `}`), + ProcessedAt: createdAt, + CreatedAt: createdAt, + }) + } + if _, err := app.db.AppendSessionEvents(ctx, session.WorkspaceUUID, session.ExternalID, events, nil); err != nil { + t.Fatalf("append large history: %v", err) + } + if err := codesessions.NewServiceWithCredentials(app.db, app.credentials, nil).ActivateManagedAgentCodeSession(ctx, codeSession); err != nil { + t.Fatalf("activate Code Session: %v", err) + } + + inbound, err := app.db.ListQueuedCodeSessionInboundEvents(ctx, codeSessionID) + if err != nil { + t.Fatalf("list inbound: %v", err) + } + if len(inbound) != eventCount+1 { + t.Fatalf("inbound count = %d, want %d", len(inbound), eventCount+1) + } + for i, eventID := range eventIDs { + if !bytes.Contains(inbound[i+1].Payload, []byte(eventID)) { + t.Fatalf("inbound %d is out of order: %s", i, inbound[i+1].Payload) + } + } +} + +func TestManagedAgentActivationRollsBackOnHistoryConversionFailure(t *testing.T) { + ctx := context.Background() + app := newTestAppWithStore(t, nil, newFakeStore("sessions-history-activation-rollback-bucket")) + defer app.close() + + agent := createAgent(t, app, `{"model":"claude-opus-4-6","name":"sessions-history-activation-rollback-agent"}`) + defer cleanupAgentRows(t, app.db, agent.ID) + env := createEnvironment(t, app, `{"name":"sessions-history-activation-rollback-env"}`) + defer cleanupEnvironmentRows(t, app.db, env.ID) + sessionResponse := createSession(t, app, `{"agent":`+quoteJSON(agent.ID)+`,"environment_id":`+quoteJSON(env.ID)+`}`) + defer deleteSession(t, app, sessionResponse.ID) + codeSessionID := launchLocalCodeSession(t, app, sessionResponse.ID) + + if _, err := app.db.Pool.Exec(ctx, ` + update code_sessions + set status = 'initializing' + where external_id = $1 + `, codeSessionID); err != nil { + t.Fatalf("set rollback code session initializing: %v", err) + } + sendSessionEvents(t, app, sessionResponse.ID, `{"events":[{"type":"user.message","content":[{"type":"text","text":"history must remain durable"}]}]}`, defaultTestKey) + session, err := app.db.GetSession(ctx, getDefaultDBIDs(t, app.db).WorkspaceUUID, sessionResponse.ID) + if err != nil { + t.Fatalf("load rollback Session: %v", err) + } + codeSession, err := app.db.GetCodeSession(ctx, codeSessionID) + if err != nil { + t.Fatalf("load rollback Code Session: %v", err) + } + invalidAt := time.Now().UTC().Add(time.Second) + if _, err := app.db.AppendSessionEvents(ctx, session.WorkspaceUUID, session.ExternalID, []db.SessionEvent{{ + UUID: uuid.NewString(), + ExternalID: "sevt_activation_invalid_" + strings.TrimPrefix(codeSessionID, "cse_"), + EventType: "user.interrupt", + Payload: json.RawMessage(`[]`), + ProcessedAt: invalidAt, + CreatedAt: invalidAt, + }}, nil); err != nil { + t.Fatalf("append invalid forwardable history event: %v", err) + } + codeSessionService := codesessions.NewServiceWithCredentials(app.db, app.credentials, nil) + before, err := app.db.ListQueuedCodeSessionInboundEvents(ctx, codeSessionID) + if err != nil || len(before) != 1 { + t.Fatalf("rollback inbound before delivery = (%#v, %v), want initialize", before, err) + } + if err := codeSessionService.ActivateManagedAgentCodeSession(ctx, codeSession); err == nil { + t.Fatal("activation with invalid forwardable history succeeded") + } + after, listErr := app.db.ListQueuedCodeSessionInboundEvents(ctx, codeSessionID) + if listErr != nil || len(after) != 1 { + t.Fatalf("rollback inbound after delivery = (%#v, %v), want unchanged initialize", after, listErr) + } + reloaded, err := app.db.GetCodeSession(ctx, codeSessionID) + if err != nil { + t.Fatalf("reload rollback Code Session: %v", err) + } + if reloaded.Status != "initializing" { + t.Fatalf("rollback Code Session status = %q, want initializing", reloaded.Status) + } +} + func TestSessionClaudeCodeSubagentInternalEventsPublishToChildThread(t *testing.T) { app := newTestAppWithStore(t, nil, newFakeStore("sessions-claude-code-subagent-internal-bucket")) defer app.close() @@ -1314,7 +1508,7 @@ func TestSessionEventsListHidesLegacyEnvManagerLog(t *testing.T) { ProcessedAt: now.Add(time.Second), CreatedAt: now.Add(time.Second), }, - }); err != nil { + }, nil); err != nil { t.Fatalf("append legacy env manager log: %v", err) } @@ -2235,6 +2429,78 @@ func TestCodeSessionWorkerEventAppendChecksEpochInsideTransaction(t *testing.T) } } +func TestCodeSessionEventAppendPreservesIdempotencyAndDirectionSequences(t *testing.T) { + app := newTestAppWithStore(t, nil, newFakeStore("sessions-code-event-append-sequences-bucket")) + defer app.close() + + agent := createAgent(t, app, `{"model":"claude-opus-4-6","name":"sessions-code-event-append-sequences-agent"}`) + defer cleanupAgentRows(t, app.db, agent.ID) + env := createEnvironment(t, app, `{"name":"sessions-code-event-append-sequences-env"}`) + defer cleanupEnvironmentRows(t, app.db, env.ID) + session := createSession(t, app, `{"agent":`+quoteJSON(agent.ID)+`,"environment_id":`+quoteJSON(env.ID)+`}`) + codeSessionID := launchLocalCodeSession(t, app, session.ID) + ctx := context.Background() + + before, err := app.db.GetCodeSession(ctx, codeSessionID) + if err != nil { + t.Fatalf("load Code Session before append: %v", err) + } + suffix := strings.TrimPrefix(codeSessionID, "cse_") + inboundInput := db.AppendCodeSessionEventInput{ + ExternalID: "csev_inbound_sequence_" + suffix, + EventType: "user", + Payload: json.RawMessage(`{"type":"user"}`), + PayloadHash: "inbound-sequence", + IdempotencyKey: "inbound-sequence:" + suffix, + Source: "test", + } + inbound, duplicate, err := app.db.AppendCodeSessionInboundEvent(ctx, codeSessionID, inboundInput) + if err != nil || duplicate { + t.Fatalf("append inbound event = (%+v, duplicate=%v, err=%v)", inbound, duplicate, err) + } + if inbound.SequenceNum != before.LastInboundSequenceNum+1 { + t.Fatalf("inbound sequence = %d, want %d", inbound.SequenceNum, before.LastInboundSequenceNum+1) + } + duplicateInbound, duplicate, err := app.db.AppendCodeSessionInboundEvent(ctx, codeSessionID, inboundInput) + if err != nil || !duplicate || duplicateInbound.UUID != inbound.UUID { + t.Fatalf("duplicate inbound event = (%+v, duplicate=%v, err=%v), want UUID %q", duplicateInbound, duplicate, err, inbound.UUID) + } + + outboundInput := db.AppendCodeSessionEventInput{ + ExternalID: "csev_outbound_sequence_" + suffix, + EventType: "assistant", + Payload: json.RawMessage(`{"type":"assistant"}`), + PayloadHash: "outbound-sequence", + IdempotencyKey: "outbound-sequence:" + suffix, + Source: "test", + } + outbound, duplicate, err := app.db.AppendCodeSessionOutboundEvent(ctx, codeSessionID, outboundInput) + if err != nil || duplicate { + t.Fatalf("append outbound event = (%+v, duplicate=%v, err=%v)", outbound, duplicate, err) + } + if outbound.SequenceNum != before.LastOutboundSequenceNum+1 { + t.Fatalf("outbound sequence = %d, want %d", outbound.SequenceNum, before.LastOutboundSequenceNum+1) + } + duplicateOutbound, duplicate, err := app.db.AppendCodeSessionOutboundEvent(ctx, codeSessionID, outboundInput) + if err != nil || !duplicate || duplicateOutbound.UUID != outbound.UUID { + t.Fatalf("duplicate outbound event = (%+v, duplicate=%v, err=%v), want UUID %q", duplicateOutbound, duplicate, err, outbound.UUID) + } + + after, err := app.db.GetCodeSession(ctx, codeSessionID) + if err != nil { + t.Fatalf("load Code Session after append: %v", err) + } + if after.LastInboundSequenceNum != inbound.SequenceNum || after.LastOutboundSequenceNum != outbound.SequenceNum { + t.Fatalf( + "stored sequences = inbound %d/outbound %d, want inbound %d/outbound %d", + after.LastInboundSequenceNum, + after.LastOutboundSequenceNum, + inbound.SequenceNum, + outbound.SequenceNum, + ) + } +} + func TestCodeSessionWorkerConnectionStateUpdatesAreEpochScoped(t *testing.T) { app := newTestAppWithStore(t, nil, newFakeStore("sessions-code-worker-epoch-state-bucket")) defer app.close() @@ -3227,6 +3493,40 @@ func TestCodeSessionWorkerStreamReplaysUnprocessedEventsForNewEpoch(t *testing.T } } +func TestSessionEventRejectedBatchHasNoSideEffects(t *testing.T) { + app := newTestAppWithStore(t, nil, newFakeStore("sessions-events-atomicity-bucket")) + defer app.close() + + agent := createAgent(t, app, `{"model":"claude-opus-4-6","name":"sessions-events-atomicity-agent"}`) + defer cleanupAgentRows(t, app.db, agent.ID) + env := createEnvironment(t, app, `{"name":"sessions-events-atomicity-env"}`) + defer cleanupEnvironmentRows(t, app.db, env.ID) + session := createSession(t, app, `{"agent":`+quoteJSON(agent.ID)+`,"environment_id":`+quoteJSON(env.ID)+`}`) + defer deleteSession(t, app, session.ID) + + response := doSessionRequest( + t, + app, + http.MethodPost, + "/v1/sessions/"+session.ID+"/events?beta=true", + strings.NewReader(`{"events":[ + {"type":"user.define_outcome","description":"must not persist","rubric":{"type":"text","text":"must pass"}}, + {"type":"user.message","content":[]} + ]}`), + defaultTestKey, + true, + ) + assertError(t, response, http.StatusBadRequest, "invalid_request_error") + + retrieved := retrieveSession(t, app, session.ID, defaultTestKey) + if string(retrieved.OutcomeEvaluations) != "[]" { + t.Fatalf("outcomes after rejected batch = %s, want []", retrieved.OutcomeEvaluations) + } + if events := listSessionEvents(t, app, session.ID, "order=asc&limit=100", defaultTestKey); len(events.Data) != 0 { + t.Fatalf("rejected batch public events = %#v, want empty", events.Data) + } +} + func TestSessionEventInputValidation(t *testing.T) { app := newTestAppWithStore(t, nil, newFakeStore("sessions-events-validation-bucket")) defer app.close() @@ -3250,6 +3550,10 @@ func TestSessionEventInputValidation(t *testing.T) { if len(valid.Data) != 2 || !bytes.Contains(valid.Data[0], []byte(`"type":"user.custom_tool_result"`)) || !bytes.Contains(valid.Data[1], []byte(`"type":"user.define_outcome"`)) { t.Fatalf("unexpected valid events response: %+v", valid) } + retrieved := retrieveSession(t, app, session.ID, defaultTestKey) + if !bytes.Contains(retrieved.OutcomeEvaluations, []byte(`"max_iterations":2`)) { + t.Fatalf("outcomes after accepted batch = %s, want committed evaluation", retrieved.OutcomeEvaluations) + } } func TestSessionWebhooks(t *testing.T) { diff --git a/tests/uuid_boundary_postgres_test.go b/tests/uuid_boundary_postgres_test.go index e51c6226..ec9dc389 100644 --- a/tests/uuid_boundary_postgres_test.go +++ b/tests/uuid_boundary_postgres_test.go @@ -561,7 +561,7 @@ func TestTypedUUIDSessionsAndRuntimePostgres(t *testing.T) { Payload: []byte(`{"typed_uuid":true}`), ProcessedAt: now, CreatedAt: now, - }}) + }}, nil) if err != nil || len(events) != 1 || events[0].ThreadUUID == nil || *events[0].ThreadUUID != thread.UUID { t.Fatalf("append Session event with inferred typed thread UUID = (%+v, %v)", events, err) }