diff --git a/cmd/gomodel/docs/docs.go b/cmd/gomodel/docs/docs.go index 89295320f..9e5f1dbfe 100644 --- a/cmd/gomodel/docs/docs.go +++ b/cmd/gomodel/docs/docs.go @@ -4521,7 +4521,7 @@ const docTemplate = `{ ] }, "post": { - "description": "Replaces the conversation metadata in full.", + "description": "Merges supplied keys into the conversation metadata.", "consumes": [ "application/json" ], @@ -4632,6 +4632,288 @@ const docTemplate = `{ ] } }, + "/v1/conversations/{id}/items": { + "get": { + "produces": [ + "application/json" + ], + "tags": [ + "conversations" + ], + "summary": "List conversation items", + "parameters": [ + { + "type": "string", + "description": "Conversation ID", + "name": "id", + "in": "path", + "required": true + }, + { + "type": "string", + "description": "Return items after this item ID", + "name": "after", + "in": "query" + }, + { + "type": "array", + "items": { + "type": "string" + }, + "collectionFormat": "csv", + "description": "Additional fields to include", + "name": "include", + "in": "query" + }, + { + "maximum": 100, + "minimum": 1, + "type": "integer", + "default": 20, + "description": "Maximum items", + "name": "limit", + "in": "query" + }, + { + "enum": [ + "asc", + "desc" + ], + "type": "string", + "default": "desc", + "description": "Sort order", + "name": "order", + "in": "query" + } + ], + "responses": { + "200": { + "description": "OK", + "schema": { + "$ref": "#/definitions/core.ConversationItemListResponse" + } + }, + "400": { + "description": "Bad Request", + "schema": { + "$ref": "#/definitions/core.OpenAIErrorEnvelope" + } + }, + "401": { + "description": "Unauthorized", + "schema": { + "$ref": "#/definitions/core.OpenAIErrorEnvelope" + } + }, + "404": { + "description": "Not Found", + "schema": { + "$ref": "#/definitions/core.OpenAIErrorEnvelope" + } + } + }, + "security": [ + { + "BearerAuth": [] + } + ] + }, + "post": { + "consumes": [ + "application/json" + ], + "produces": [ + "application/json" + ], + "tags": [ + "conversations" + ], + "summary": "Create conversation items", + "parameters": [ + { + "type": "string", + "description": "Conversation ID", + "name": "id", + "in": "path", + "required": true + }, + { + "type": "array", + "items": { + "type": "string" + }, + "collectionFormat": "csv", + "description": "Additional fields to include", + "name": "include", + "in": "query" + }, + { + "description": "Conversation items", + "name": "request", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/core.ConversationItemCreateRequest" + } + } + ], + "responses": { + "200": { + "description": "OK", + "schema": { + "$ref": "#/definitions/core.ConversationItemListResponse" + } + }, + "400": { + "description": "Bad Request", + "schema": { + "$ref": "#/definitions/core.OpenAIErrorEnvelope" + } + }, + "401": { + "description": "Unauthorized", + "schema": { + "$ref": "#/definitions/core.OpenAIErrorEnvelope" + } + }, + "404": { + "description": "Not Found", + "schema": { + "$ref": "#/definitions/core.OpenAIErrorEnvelope" + } + } + }, + "security": [ + { + "BearerAuth": [] + } + ] + } + }, + "/v1/conversations/{id}/items/{item_id}": { + "get": { + "produces": [ + "application/json" + ], + "tags": [ + "conversations" + ], + "summary": "Get a conversation item", + "parameters": [ + { + "type": "string", + "description": "Conversation ID", + "name": "id", + "in": "path", + "required": true + }, + { + "type": "string", + "description": "Conversation item ID", + "name": "item_id", + "in": "path", + "required": true + }, + { + "type": "array", + "items": { + "type": "string" + }, + "collectionFormat": "csv", + "description": "Additional fields to include", + "name": "include", + "in": "query" + } + ], + "responses": { + "200": { + "description": "OK", + "schema": { + "type": "object" + } + }, + "400": { + "description": "Bad Request", + "schema": { + "$ref": "#/definitions/core.OpenAIErrorEnvelope" + } + }, + "401": { + "description": "Unauthorized", + "schema": { + "$ref": "#/definitions/core.OpenAIErrorEnvelope" + } + }, + "404": { + "description": "Not Found", + "schema": { + "$ref": "#/definitions/core.OpenAIErrorEnvelope" + } + } + }, + "security": [ + { + "BearerAuth": [] + } + ] + }, + "delete": { + "produces": [ + "application/json" + ], + "tags": [ + "conversations" + ], + "summary": "Delete a conversation item", + "parameters": [ + { + "type": "string", + "description": "Conversation ID", + "name": "id", + "in": "path", + "required": true + }, + { + "type": "string", + "description": "Conversation item ID", + "name": "item_id", + "in": "path", + "required": true + } + ], + "responses": { + "200": { + "description": "OK", + "schema": { + "$ref": "#/definitions/core.Conversation" + } + }, + "400": { + "description": "Bad Request", + "schema": { + "$ref": "#/definitions/core.OpenAIErrorEnvelope" + } + }, + "401": { + "description": "Unauthorized", + "schema": { + "$ref": "#/definitions/core.OpenAIErrorEnvelope" + } + }, + "404": { + "description": "Not Found", + "schema": { + "$ref": "#/definitions/core.OpenAIErrorEnvelope" + } + } + }, + "security": [ + { + "BearerAuth": [] + } + ] + } + }, "/v1/embeddings": { "post": { "consumes": [ @@ -8446,6 +8728,45 @@ const docTemplate = `{ } } }, + "core.ConversationItemCreateRequest": { + "type": "object", + "required": [ + "items" + ], + "properties": { + "items": { + "type": "array", + "items": { + "type": "object" + } + } + } + }, + "core.ConversationItemListResponse": { + "type": "object", + "properties": { + "data": { + "type": "array", + "items": { + "type": "object" + } + }, + "first_id": { + "type": "string", + "x-nullable": true + }, + "has_more": { + "type": "boolean" + }, + "last_id": { + "type": "string", + "x-nullable": true + }, + "object": { + "type": "string" + } + } + }, "core.ConversationUpdateRequest": { "type": "object", "required": [ diff --git a/docs/advanced/api-endpoints.mdx b/docs/advanced/api-endpoints.mdx index 26dff44c6..edee1c58b 100644 --- a/docs/advanced/api-endpoints.mdx +++ b/docs/advanced/api-endpoints.mdx @@ -27,8 +27,12 @@ For request and response details, see the dedicated guides: | `/v1/responses/compact` | POST | Compact a Responses conversation (provider-native where supported) | | `/v1/conversations` | POST | Create a conversation (gateway-managed) | | `/v1/conversations/{id}` | GET | Retrieve a conversation | -| `/v1/conversations/{id}` | POST | Replace conversation metadata in full | +| `/v1/conversations/{id}` | POST | Merge conversation metadata | | `/v1/conversations/{id}` | DELETE | Delete a conversation | +| `/v1/conversations/{id}/items` | POST | Add items to a conversation | +| `/v1/conversations/{id}/items` | GET | List conversation items with cursor pagination | +| `/v1/conversations/{id}/items/{item_id}` | GET | Retrieve a conversation item | +| `/v1/conversations/{id}/items/{item_id}` | DELETE | Delete a conversation item and return the conversation | | `/v1/embeddings` | POST | Text embeddings | | `/v1/models` | GET | List available models | | `/v1/audio/speech` | POST | Text-to-speech, returning binary audio | diff --git a/docs/advanced/conversations-api.mdx b/docs/advanced/conversations-api.mdx index 64ddd703e..85b1cb85a 100644 --- a/docs/advanced/conversations-api.mdx +++ b/docs/advanced/conversations-api.mdx @@ -1,6 +1,6 @@ --- title: "Conversations API" -description: "Create, retrieve, update, and delete OpenAI-compatible conversations through GoModel." +description: "Manage OpenAI-compatible conversations and their items through GoModel." icon: "messages-square" tag: "Beta" --- @@ -22,8 +22,12 @@ provider configuration. | --- | --- | | `POST /v1/conversations` | Creates a conversation. Accepts optional `items` and `metadata`. | | `GET /v1/conversations/{id}` | Returns a stored conversation. | -| `POST /v1/conversations/{id}` | Replaces the conversation `metadata` in full. | +| `POST /v1/conversations/{id}` | Merges supplied keys into the conversation `metadata`. | | `DELETE /v1/conversations/{id}` | Deletes a stored conversation. | +| `POST /v1/conversations/{id}/items` | Adds up to 20 items and returns the added items as a list. | +| `GET /v1/conversations/{id}/items` | Lists items with `after`, `include`, `limit`, and `order`. | +| `GET /v1/conversations/{id}/items/{item_id}` | Returns one conversation item. | +| `DELETE /v1/conversations/{id}/items/{item_id}` | Deletes one item and returns the conversation. | ## Conversation object @@ -47,22 +51,48 @@ compatible: - `metadata` — at most 16 key-value pairs; keys up to 64 characters; values up to 512 characters. -The `items` array is accepted and stored with the conversation. It is not yet -exposed through a conversation items listing endpoint. +Item payloads follow the Responses API input/output item union. GoModel +normalizes convenient message strings, assigns stable item IDs, and preserves +new or provider-specific item fields so they can be listed and replayed later. + +Item lists default to newest-first (`order=desc`) with 20 items per page. Use +the returned `last_id` as the next request's `after` cursor. `limit` accepts up +to 100 items. The optional `include` parameter accepts the same values as the +OpenAI Conversations API, including `reasoning.encrypted_content` and +`message.output_text.logprobs`. GoModel stores complete item payloads, so +optional fields remain available for later requests while only the fields named +by `include` are returned. ## Storage and retention -Conversations are held in an in-memory store. They survive across requests but -not process restarts. Retention is bounded: conversations expire after 30 days -and the store keeps at most 10,000 conversations, evicting the oldest first. +The normal GoModel application stores conversations in the configured shared +storage backend (`sqlite`, `postgresql`, or `mongodb`), so they survive process +restarts. Conversation snapshots expire after 30 days. + +An in-memory fallback is used only when embedding the HTTP server without an +application storage configuration, such as lightweight tests. That fallback +also expires entries after 30 days and keeps at most 10,000 conversations, +evicting the oldest first. ## Errors GoModel returns OpenAI-compatible errors: -- `400 invalid_request_error` — invalid body, missing `metadata` on update, or a - limit exceeded (the `param` field names the offending field). -- `404 not_found_error` — the conversation ID does not exist. +- `400 invalid_request_error` — invalid body, missing `metadata` on update, an + invalid item, or a limit exceeded (the `param` field names the offending field). +- `404 not_found_error` — the conversation or requested item does not exist. +- `404 invalid_request_error` with `param: "after"` — the pagination cursor does + not identify an item in the conversation. + +Conversation persistence is part of a successful `/v1/responses` turn. If the +provider succeeds but the completed exchange cannot be stored, a non-streaming +request returns `500`; a streaming request ends without emitting +`response.completed`. GoModel does not currently deduplicate a client retry of +that failed request, so retrying can invoke the provider and incur its cost +again. The failed exchange is absent from conversation history, so inspecting +history cannot reliably determine whether the provider call already ran. Avoid +automatic retries, or coordinate them through an independent durable +idempotency mechanism when duplicate calls have material cost or side effects. ## Examples @@ -84,7 +114,7 @@ curl http://localhost:8080/v1/conversations/conv_abc123 \ -H "Authorization: Bearer $GOMODEL_MASTER_KEY" ``` -Update its metadata (replaces all metadata): +Update its metadata (existing keys not supplied here are retained): ```bash curl http://localhost:8080/v1/conversations/conv_abc123 \ @@ -95,6 +125,36 @@ curl http://localhost:8080/v1/conversations/conv_abc123 \ }' ``` +Add advanced Responses items: + +```bash +curl 'http://localhost:8080/v1/conversations/conv_abc123/items?include=reasoning.encrypted_content' \ + -H "Authorization: Bearer $GOMODEL_MASTER_KEY" \ + -H "Content-Type: application/json" \ + -d '{ + "items": [ + { + "type": "function_call", + "call_id": "call_lookup_1", + "name": "lookup_order", + "arguments": "{\"order_id\":\"A-42\"}" + }, + { + "type": "function_call_output", + "call_id": "call_lookup_1", + "output": "{\"status\":\"shipped\"}" + } + ] + }' +``` + +List the next page in chronological order: + +```bash +curl 'http://localhost:8080/v1/conversations/conv_abc123/items?order=asc&limit=10&after=msg_abc123' \ + -H "Authorization: Bearer $GOMODEL_MASTER_KEY" +``` + Delete it: ```bash diff --git a/docs/openapi.json b/docs/openapi.json index 7e4130fea..09171feb3 100644 --- a/docs/openapi.json +++ b/docs/openapi.json @@ -6760,7 +6760,7 @@ } }, "post": { - "description": "Replaces the conversation metadata in full.", + "description": "Merges supplied keys into the conversation metadata.", "tags": [ "conversations" ], @@ -6910,6 +6910,386 @@ } } }, + "/v1/conversations/{id}/items": { + "get": { + "tags": [ + "conversations" + ], + "summary": "List conversation items", + "parameters": [ + { + "description": "Conversation ID", + "name": "id", + "in": "path", + "required": true, + "schema": { + "type": "string" + } + }, + { + "description": "Return items after this item ID", + "name": "after", + "in": "query", + "schema": { + "type": "string" + } + }, + { + "description": "Additional fields to include", + "name": "include", + "in": "query", + "style": "form", + "explode": false, + "schema": { + "type": "array", + "items": { + "type": "string" + } + } + }, + { + "description": "Maximum items", + "name": "limit", + "in": "query", + "schema": { + "type": "integer", + "minimum": 1, + "maximum": 100, + "default": 20 + } + }, + { + "description": "Sort order", + "name": "order", + "in": "query", + "schema": { + "type": "string", + "enum": [ + "asc", + "desc" + ], + "default": "desc" + } + } + ], + "responses": { + "200": { + "description": "OK", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.ConversationItemListResponse" + } + } + } + }, + "400": { + "description": "Bad Request", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.OpenAIErrorEnvelope" + } + } + } + }, + "401": { + "description": "Unauthorized", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.OpenAIErrorEnvelope" + } + } + } + }, + "404": { + "description": "Not Found", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.OpenAIErrorEnvelope" + } + } + } + } + }, + "security": [ + { + "BearerAuth": [] + } + ], + "x-mint": { + "metadata": { + "sidebarTitle": "/v1/conversations/{id}/items" + } + } + }, + "post": { + "tags": [ + "conversations" + ], + "summary": "Create conversation items", + "parameters": [ + { + "description": "Conversation ID", + "name": "id", + "in": "path", + "required": true, + "schema": { + "type": "string" + } + }, + { + "description": "Additional fields to include", + "name": "include", + "in": "query", + "style": "form", + "explode": false, + "schema": { + "type": "array", + "items": { + "type": "string" + } + } + } + ], + "requestBody": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.ConversationItemCreateRequest" + } + } + }, + "description": "Conversation items", + "required": true + }, + "responses": { + "200": { + "description": "OK", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.ConversationItemListResponse" + } + } + } + }, + "400": { + "description": "Bad Request", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.OpenAIErrorEnvelope" + } + } + } + }, + "401": { + "description": "Unauthorized", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.OpenAIErrorEnvelope" + } + } + } + }, + "404": { + "description": "Not Found", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.OpenAIErrorEnvelope" + } + } + } + } + }, + "security": [ + { + "BearerAuth": [] + } + ], + "x-mint": { + "metadata": { + "sidebarTitle": "/v1/conversations/{id}/items" + } + } + } + }, + "/v1/conversations/{id}/items/{item_id}": { + "get": { + "tags": [ + "conversations" + ], + "summary": "Get a conversation item", + "parameters": [ + { + "description": "Conversation ID", + "name": "id", + "in": "path", + "required": true, + "schema": { + "type": "string" + } + }, + { + "description": "Conversation item ID", + "name": "item_id", + "in": "path", + "required": true, + "schema": { + "type": "string" + } + }, + { + "description": "Additional fields to include", + "name": "include", + "in": "query", + "style": "form", + "explode": false, + "schema": { + "type": "array", + "items": { + "type": "string" + } + } + } + ], + "responses": { + "200": { + "description": "OK", + "content": { + "application/json": { + "schema": { + "type": "object" + } + } + } + }, + "400": { + "description": "Bad Request", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.OpenAIErrorEnvelope" + } + } + } + }, + "401": { + "description": "Unauthorized", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.OpenAIErrorEnvelope" + } + } + } + }, + "404": { + "description": "Not Found", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.OpenAIErrorEnvelope" + } + } + } + } + }, + "security": [ + { + "BearerAuth": [] + } + ], + "x-mint": { + "metadata": { + "sidebarTitle": "/v1/conversations/{id}/items/{item_id}" + } + } + }, + "delete": { + "tags": [ + "conversations" + ], + "summary": "Delete a conversation item", + "parameters": [ + { + "description": "Conversation ID", + "name": "id", + "in": "path", + "required": true, + "schema": { + "type": "string" + } + }, + { + "description": "Conversation item ID", + "name": "item_id", + "in": "path", + "required": true, + "schema": { + "type": "string" + } + } + ], + "responses": { + "200": { + "description": "OK", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.Conversation" + } + } + } + }, + "400": { + "description": "Bad Request", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.OpenAIErrorEnvelope" + } + } + } + }, + "401": { + "description": "Unauthorized", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.OpenAIErrorEnvelope" + } + } + } + }, + "404": { + "description": "Not Found", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.OpenAIErrorEnvelope" + } + } + } + } + }, + "security": [ + { + "BearerAuth": [] + } + ], + "x-mint": { + "metadata": { + "sidebarTitle": "/v1/conversations/{id}/items/{item_id}" + } + } + } + }, "/v1/embeddings": { "post": { "tags": [ @@ -11475,6 +11855,45 @@ } } }, + "core.ConversationItemCreateRequest": { + "type": "object", + "required": [ + "items" + ], + "properties": { + "items": { + "type": "array", + "items": { + "type": "object" + } + } + } + }, + "core.ConversationItemListResponse": { + "type": "object", + "properties": { + "data": { + "type": "array", + "items": { + "type": "object" + } + }, + "first_id": { + "type": "string", + "nullable": true + }, + "has_more": { + "type": "boolean" + }, + "last_id": { + "type": "string", + "nullable": true + }, + "object": { + "type": "string" + } + } + }, "core.ConversationUpdateRequest": { "type": "object", "required": [ diff --git a/internal/conversationstore/store.go b/internal/conversationstore/store.go index fd154f101..1f374f34b 100644 --- a/internal/conversationstore/store.go +++ b/internal/conversationstore/store.go @@ -15,8 +15,19 @@ import ( "github.com/enterpilot/gomodel/internal/core" ) -// ErrNotFound indicates a requested conversation was not found. -var ErrNotFound = errors.New("conversation not found") +var ( + // ErrNotFound indicates a requested conversation was not found. + ErrNotFound = errors.New("conversation not found") + // ErrItemNotFound indicates a requested item was not present in an existing + // conversation. + ErrItemNotFound = errors.New("conversation item not found") + // ErrDuplicateItem indicates an append would introduce an item id already + // present in the conversation. + ErrDuplicateItem = errors.New("conversation item id already exists") + // ErrMetadataLimitExceeded indicates a patch would make the final metadata + // object exceed the OpenAI Conversations limit. + ErrMetadataLimitExceeded = errors.New("conversation metadata limit exceeded") +) // StoredConversation keeps the public conversation snapshot separate from // gateway-only metadata (initial items, owning user path, request id). @@ -33,15 +44,49 @@ type StoredConversation struct { type Store interface { Create(ctx context.Context, conversation *StoredConversation) error Get(ctx context.Context, id string) (*StoredConversation, error) - Update(ctx context.Context, conversation *StoredConversation) error + // MergeMetadata atomically overlays metadata without rewriting items that a + // concurrently completing Responses turn may be appending. + MergeMetadata(ctx context.Context, id string, metadata map[string]string) (*StoredConversation, error) // AppendItems atomically appends items to an existing conversation, so two // concurrently completing turns cannot overwrite each other's exchange the // way a Get-then-Update would. AppendItems(ctx context.Context, id string, items []json.RawMessage) error + // DeleteItem atomically removes one item and returns the updated snapshot. + DeleteItem(ctx context.Context, id, itemID string) (*StoredConversation, error) Delete(ctx context.Context, id string) error Close() error } +func itemID(raw json.RawMessage) string { + var item struct { + ID string `json:"id"` + } + if err := json.Unmarshal(raw, &item); err != nil { + return "" + } + return strings.TrimSpace(item.ID) +} + +func duplicateItemID(existing, added []json.RawMessage) string { + ids := make(map[string]struct{}, len(existing)+len(added)) + for _, raw := range existing { + if id := itemID(raw); id != "" { + ids[id] = struct{}{} + } + } + for _, raw := range added { + id := itemID(raw) + if id == "" { + continue + } + if _, exists := ids[id]; exists { + return id + } + ids[id] = struct{}{} + } + return "" +} + func cloneConversation(src *StoredConversation) (*StoredConversation, error) { dst, _, err := cloneConversationWithSize(src) return dst, err diff --git a/internal/conversationstore/store_memory.go b/internal/conversationstore/store_memory.go index 44579c27d..b0cca2f28 100644 --- a/internal/conversationstore/store_memory.go +++ b/internal/conversationstore/store_memory.go @@ -3,6 +3,7 @@ package conversationstore import ( "context" "fmt" + "maps" "sort" "sync" "time" @@ -90,16 +91,20 @@ func (s *MemoryStore) Create(_ context.Context, conversation *StoredConversation return fmt.Errorf("conversation id is required") } - c, size, err := cloneConversationWithSize(conversation) + c, err := cloneConversation(conversation) if err != nil { return err } - if err := s.checkByteBudget(size); err != nil { - return err - } now := time.Now().UTC() prepareStoredConversationForMemory(c, now, s.ttl) + c, size, err := cloneConversationWithSize(c) + if err != nil { + return err + } + if err := s.checkByteBudget(size); err != nil { + return err + } s.mu.Lock() defer s.mu.Unlock() @@ -122,60 +127,57 @@ func (s *MemoryStore) Create(_ context.Context, conversation *StoredConversation func (s *MemoryStore) Get(_ context.Context, id string) (*StoredConversation, error) { now := time.Now().UTC() s.mu.Lock() + defer s.mu.Unlock() s.cleanupExpiredLocked(now) conversation, ok := s.items[id] if !ok { - s.mu.Unlock() return nil, ErrNotFound } if conversationExpired(conversation, now) { s.removeLocked(id) - s.mu.Unlock() return nil, ErrNotFound } - s.mu.Unlock() + // Clone while the lock is held so future store changes cannot accidentally + // expose an internal snapshot or reintroduce a read/write race. return cloneConversation(conversation) } -// Update replaces an existing conversation snapshot. -func (s *MemoryStore) Update(_ context.Context, conversation *StoredConversation) error { - if conversation == nil || conversation.Conversation == nil || conversation.Conversation.ID == "" { - return fmt.Errorf("conversation id is required") - } - c, size, err := cloneConversationWithSize(conversation) - if err != nil { - return err - } - if err := s.checkByteBudget(size); err != nil { - return err - } - +// MergeMetadata overlays metadata while holding the same lock that protects +// item appends, so a metadata update cannot restore a stale item slice. +func (s *MemoryStore) MergeMetadata(_ context.Context, id string, metadata map[string]string) (*StoredConversation, error) { now := time.Now().UTC() s.mu.Lock() defer s.mu.Unlock() s.cleanupExpiredLocked(now) - existing, exists := s.items[c.Conversation.ID] - if !exists { - return ErrNotFound + existing, exists := s.items[id] + if !exists || conversationExpired(existing, now) { + if exists { + s.removeLocked(id) + } + return nil, ErrNotFound } - if conversationExpired(existing, now) { - s.removeLocked(c.Conversation.ID) - return ErrNotFound + + candidate, err := cloneConversation(existing) + if err != nil { + return nil, err } - if c.StoredAt.IsZero() { - c.StoredAt = existing.StoredAt + if candidate.Conversation.Metadata == nil { + candidate.Conversation.Metadata = make(map[string]string, len(metadata)) } - if c.ExpiresAt.IsZero() { - c.ExpiresAt = existing.ExpiresAt + maps.Copy(candidate.Conversation.Metadata, metadata) + if len(candidate.Conversation.Metadata) > core.MaxConversationMetadataPairs { + return nil, ErrMetadataLimitExceeded } - prepareStoredConversationForMemory(c, now, s.ttl) - if conversationExpired(c, now) { - s.removeLocked(c.Conversation.ID) - return ErrNotFound + _, size, err := cloneConversationWithSize(candidate) + if err != nil { + return nil, err } - s.putLocked(c.Conversation.ID, c, size) - s.enforceBoundsLocked(c.Conversation.ID) - return nil + if err := s.checkByteBudget(size); err != nil { + return nil, err + } + s.putLocked(id, candidate, size) + s.enforceBoundsLocked(id) + return cloneConversation(candidate) } // AppendItems atomically appends items to an existing conversation snapshot. @@ -183,11 +185,9 @@ func (s *MemoryStore) AppendItems(_ context.Context, id string, items []json.Raw if len(items) == 0 { return nil } - var added int64 - for _, item := range items { - added += int64(len(item)) + if duplicateItemID(nil, items) != "" { + return ErrDuplicateItem } - now := time.Now().UTC() s.mu.Lock() defer s.mu.Unlock() @@ -200,21 +200,68 @@ func (s *MemoryStore) AppendItems(_ context.Context, id string, items []json.Raw s.removeLocked(id) return ErrNotFound } - // Reject growth past the byte budget before mutating, mirroring Create and - // Update; otherwise bound enforcement would have to drop the very - // conversation the caller believes was just persisted. - if s.maxBytes > 0 && s.sizes[id]+added > s.maxBytes { - return fmt.Errorf("conversation snapshot would grow to %d bytes, exceeding the in-memory store budget of %d bytes", s.sizes[id]+added, s.maxBytes) + if duplicateItemID(conversation.Items, items) != "" { + return ErrDuplicateItem + } + candidate, err := cloneConversation(conversation) + if err != nil { + return err } for _, item := range items { - conversation.Items = append(conversation.Items, core.CloneRawJSON(item)) + candidate.Items = append(candidate.Items, core.CloneRawJSON(item)) } - s.sizes[id] += added - s.totalBytes += added + candidate, size, err := cloneConversationWithSize(candidate) + if err != nil { + return err + } + // Reject growth past the byte budget before mutating, mirroring Create; + // otherwise bound enforcement would have to drop the very + // conversation the caller believes was just persisted. + if err := s.checkByteBudget(size); err != nil { + return err + } + s.putLocked(id, candidate, size) s.enforceBoundsLocked(id) return nil } +// DeleteItem removes one item while holding the append lock. +func (s *MemoryStore) DeleteItem(_ context.Context, id, targetItemID string) (*StoredConversation, error) { + now := time.Now().UTC() + s.mu.Lock() + defer s.mu.Unlock() + s.cleanupExpiredLocked(now) + existing, exists := s.items[id] + if !exists || conversationExpired(existing, now) { + if exists { + s.removeLocked(id) + } + return nil, ErrNotFound + } + + index := -1 + for i, raw := range existing.Items { + if itemID(raw) == targetItemID { + index = i + break + } + } + if index < 0 { + return nil, ErrItemNotFound + } + candidate, err := cloneConversation(existing) + if err != nil { + return nil, err + } + candidate.Items = append(candidate.Items[:index], candidate.Items[index+1:]...) + _, size, err := cloneConversationWithSize(candidate) + if err != nil { + return nil, err + } + s.putLocked(id, candidate, size) + return cloneConversation(candidate) +} + // Delete removes one conversation snapshot by id. func (s *MemoryStore) Delete(_ context.Context, id string) error { s.mu.Lock() diff --git a/internal/conversationstore/store_memory_test.go b/internal/conversationstore/store_memory_test.go index 81706d8d4..5e220ca63 100644 --- a/internal/conversationstore/store_memory_test.go +++ b/internal/conversationstore/store_memory_test.go @@ -24,7 +24,7 @@ func storedConversation(id string, storedAt time.Time) *StoredConversation { } } -func TestMemoryStoreCreateGetUpdateDelete(t *testing.T) { +func TestMemoryStoreCreateGetDelete(t *testing.T) { ctx := context.Background() store := NewMemoryStore() @@ -40,18 +40,6 @@ func TestMemoryStoreCreateGetUpdateDelete(t *testing.T) { t.Fatalf("id = %q, want conv_1", got.Conversation.ID) } - got.Conversation.Metadata = map[string]string{"k": "v"} - if err := store.Update(ctx, got); err != nil { - t.Fatalf("Update() error = %v", err) - } - updated, err := store.Get(ctx, "conv_1") - if err != nil { - t.Fatalf("Get() after update error = %v", err) - } - if updated.Conversation.Metadata["k"] != "v" { - t.Fatalf("metadata[k] = %q, want v", updated.Conversation.Metadata["k"]) - } - if err := store.Delete(ctx, "conv_1"); err != nil { t.Fatalf("Delete() error = %v", err) } @@ -72,18 +60,55 @@ func TestMemoryStoreCreateRejectsDuplicate(t *testing.T) { } } -func TestMemoryStoreUpdateMissingReturnsNotFound(t *testing.T) { +func TestMemoryStoreDeleteMissingReturnsNotFound(t *testing.T) { + if err := NewMemoryStore().Delete(context.Background(), "conv_missing"); !errors.Is(err, ErrNotFound) { + t.Fatalf("Delete() error = %v, want ErrNotFound", err) + } +} + +func TestMemoryStoreConcurrentAppendRejectsDuplicateItemID(t *testing.T) { ctx := context.Background() store := NewMemoryStore() - - if err := store.Update(ctx, storedConversation("conv_missing", time.Time{})); !errors.Is(err, ErrNotFound) { - t.Fatalf("Update() error = %v, want ErrNotFound", err) + if err := store.Create(ctx, storedConversation("conv_duplicate_items", time.Time{})); err != nil { + t.Fatalf("Create() error = %v", err) } -} -func TestMemoryStoreDeleteMissingReturnsNotFound(t *testing.T) { - if err := NewMemoryStore().Delete(context.Background(), "conv_missing"); !errors.Is(err, ErrNotFound) { - t.Fatalf("Delete() error = %v, want ErrNotFound", err) + item := json.RawMessage(`{"id":"msg_shared","type":"message","role":"user","content":[]}`) + const writers = 32 + start := make(chan struct{}) + errs := make(chan error, writers) + var wg sync.WaitGroup + for range writers { + wg.Go(func() { + <-start + errs <- store.AppendItems(ctx, "conv_duplicate_items", []json.RawMessage{item}) + }) + } + close(start) + wg.Wait() + close(errs) + + succeeded := 0 + duplicates := 0 + for err := range errs { + switch { + case err == nil: + succeeded++ + case errors.Is(err, ErrDuplicateItem): + duplicates++ + default: + t.Fatalf("AppendItems() unexpected error = %v", err) + } + } + if succeeded != 1 || duplicates != writers-1 { + t.Fatalf("append results = %d success, %d duplicate; want 1/%d", succeeded, duplicates, writers-1) + } + got, err := store.Get(ctx, "conv_duplicate_items") + if err != nil { + t.Fatalf("Get() error = %v", err) + } + if len(got.Items) != 1 { + t.Fatalf("stored items = %d, want 1", len(got.Items)) } } @@ -198,6 +223,61 @@ func TestMemoryStoreAppendItems(t *testing.T) { } } +func TestMemoryStoreMergeMetadataAndDeleteItem(t *testing.T) { + store := NewMemoryStore() + conv := storedConversation("conv_items", time.Time{}) + conv.Conversation.Metadata = map[string]string{"existing": "kept"} + conv.Items = []json.RawMessage{ + json.RawMessage(`{"id":"msg_1","type":"message"}`), + json.RawMessage(`{"id":"msg_2","type":"message"}`), + } + if err := store.Create(context.Background(), conv); err != nil { + t.Fatalf("Create() error = %v", err) + } + + merged, err := store.MergeMetadata(context.Background(), "conv_items", map[string]string{"new": "value"}) + if err != nil { + t.Fatalf("MergeMetadata() error = %v", err) + } + if merged.Conversation.Metadata["existing"] != "kept" || merged.Conversation.Metadata["new"] != "value" || len(merged.Items) != 2 { + t.Fatalf("merged = %+v, want merged metadata and preserved items", merged) + } + + updated, err := store.DeleteItem(context.Background(), "conv_items", "msg_1") + if err != nil { + t.Fatalf("DeleteItem() error = %v", err) + } + if len(updated.Items) != 1 || itemID(updated.Items[0]) != "msg_2" { + t.Fatalf("items = %s, want msg_2 only", updated.Items) + } + if _, err := store.DeleteItem(context.Background(), "conv_items", "missing"); !errors.Is(err, ErrItemNotFound) { + t.Fatalf("DeleteItem(missing) error = %v, want ErrItemNotFound", err) + } +} + +func TestMemoryStoreMergeMetadataRejectsOversizedResult(t *testing.T) { + store := NewMemoryStore() + conv := storedConversation("conv_metadata_limit", time.Time{}) + conv.Conversation.Metadata = make(map[string]string, core.MaxConversationMetadataPairs) + for index := range core.MaxConversationMetadataPairs { + conv.Conversation.Metadata[fmt.Sprintf("key_%d", index)] = "value" + } + if err := store.Create(context.Background(), conv); err != nil { + t.Fatalf("Create() error = %v", err) + } + + if _, err := store.MergeMetadata(context.Background(), conv.Conversation.ID, map[string]string{"extra": "value"}); !errors.Is(err, ErrMetadataLimitExceeded) { + t.Fatalf("MergeMetadata() error = %v, want ErrMetadataLimitExceeded", err) + } + got, err := store.Get(context.Background(), conv.Conversation.ID) + if err != nil { + t.Fatalf("Get() error = %v", err) + } + if len(got.Conversation.Metadata) != core.MaxConversationMetadataPairs { + t.Fatalf("metadata size = %d, want %d", len(got.Conversation.Metadata), core.MaxConversationMetadataPairs) + } +} + func TestMemoryStoreAppendItems_ConcurrentAppendsAllSurvive(t *testing.T) { store := NewMemoryStore() conv := &StoredConversation{Conversation: &core.Conversation{ID: "conv_race", Object: "conversation"}} @@ -285,7 +365,7 @@ func TestMemoryStoreAppendItemsCountsTowardByteBudget(t *testing.T) { ctx := context.Background() now := time.Now().UTC() - store := NewMemoryStore(WithTTL(0), WithMaxEntries(0), WithMaxBytes(2000)) + store := NewMemoryStore(WithTTL(0), WithMaxEntries(0), WithMaxBytes(2100)) if err := store.Create(ctx, storedConversation("conv_old", now.Add(-time.Minute))); err != nil { t.Fatalf("Create(conv_old) error = %v", err) } @@ -310,8 +390,15 @@ func TestMemoryStoreAppendItemsCountsTowardByteBudget(t *testing.T) { if len(grown.Items) != 1 { t.Fatalf("conv_grow items = %d, want 1", len(grown.Items)) } - if store.totalBytes > 2000 { - t.Fatalf("totalBytes = %d, want <= 2000", store.totalBytes) + if store.totalBytes > 2100 { + t.Fatalf("totalBytes = %d, want <= 2100", store.totalBytes) + } + _, exactSize, err := cloneConversationWithSize(grown) + if err != nil { + t.Fatalf("measure grown conversation: %v", err) + } + if store.sizes["conv_grow"] != exactSize { + t.Fatalf("recorded size = %d, exact serialized size = %d", store.sizes["conv_grow"], exactSize) } } @@ -319,7 +406,7 @@ func TestMemoryStoreAppendItemsRejectsOversizeGrowth(t *testing.T) { ctx := context.Background() now := time.Now().UTC() - store := NewMemoryStore(WithTTL(0), WithMaxEntries(0), WithMaxBytes(2000)) + store := NewMemoryStore(WithTTL(0), WithMaxEntries(0), WithMaxBytes(2100)) if err := store.Create(ctx, storedConversation("conv_other", now.Add(-time.Minute))); err != nil { t.Fatalf("Create(conv_other) error = %v", err) } @@ -350,7 +437,7 @@ func TestMemoryStoreAppendItemsNeverEvictsAppendedConversation(t *testing.T) { ctx := context.Background() now := time.Now().UTC() - store := NewMemoryStore(WithTTL(0), WithMaxEntries(0), WithMaxBytes(2000)) + store := NewMemoryStore(WithTTL(0), WithMaxEntries(0), WithMaxBytes(2100)) // conv_grow is the OLDEST entry — without protection, oldest-first // eviction would drop it right after its own successful append. if err := store.Create(ctx, storedConversation("conv_grow", now.Add(-time.Minute))); err != nil { diff --git a/internal/conversationstore/store_mongodb.go b/internal/conversationstore/store_mongodb.go index d6a02fef9..40a455c67 100644 --- a/internal/conversationstore/store_mongodb.go +++ b/internal/conversationstore/store_mongodb.go @@ -5,6 +5,8 @@ import ( "errors" "fmt" "log/slog" + "maps" + "math/rand/v2" "sync" "time" @@ -13,6 +15,7 @@ import ( "go.mongodb.org/mongo-driver/v2/mongo" "go.mongodb.org/mongo-driver/v2/mongo/options" + "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/storage" ) @@ -34,6 +37,12 @@ type MongoDBStore struct { closeOnce sync.Once } +const ( + mongoMutationMaxAttempts = 64 + mongoMutationInitialBackoff = time.Millisecond + mongoMutationMaxBackoff = 20 * time.Millisecond +) + // NewMongoDBStore creates collection indexes if needed and starts the hourly // expired-snapshot sweep. func NewMongoDBStore(database *mongo.Database) (*MongoDBStore, error) { @@ -111,50 +120,147 @@ func (s *MongoDBStore) Get(ctx context.Context, id string) (*StoredConversation, return stored, nil } -// Update replaces an existing, unexpired conversation snapshot including its -// items. Zero StoredAt or ExpiresAt values preserve the stored retention fields. -func (s *MongoDBStore) Update(ctx context.Context, conversation *StoredConversation) error { - now := time.Now().UTC() - normalized, data, _, err := prepareStoredConversationForStorage(conversation, now, s.ttl, false) - if err != nil { - return err - } - set := bson.M{ - "data": string(data), - "items": itemsToStrings(normalized.Items), - } - if !normalized.StoredAt.IsZero() { - set["stored_at"] = normalized.StoredAt.Unix() +// MergeMetadata uses an optimistic compare-and-swap on the serialized +// snapshot. Items live in a separate field, so they are never rewritten. +func (s *MongoDBStore) MergeMetadata(ctx context.Context, id string, metadata map[string]string) (*StoredConversation, error) { + for attempt := range mongoMutationMaxAttempts { + var doc mongoConversationDocument + // id is encoded by the driver as a BSON string value; it cannot add + // query keys or operators. lgtm[go/sql-injection] + if err := s.collection.FindOne(ctx, bson.M{"_id": id}).Decode(&doc); err != nil { + if errors.Is(err, mongo.ErrNoDocuments) { + return nil, ErrNotFound + } + return nil, fmt.Errorf("query conversation snapshot: %w", err) + } + now := time.Now().UTC() + if doc.ExpiresAt > 0 && doc.ExpiresAt <= now.Unix() { + return nil, ErrNotFound + } + stored, err := decodeStoredConversation([]byte(doc.Data), nil, doc.StoredAt, doc.ExpiresAt) + if err != nil { + return nil, err + } + if stored.Conversation.Metadata == nil { + stored.Conversation.Metadata = make(map[string]string, len(metadata)) + } + maps.Copy(stored.Conversation.Metadata, metadata) + if len(stored.Conversation.Metadata) > core.MaxConversationMetadataPairs { + return nil, ErrMetadataLimitExceeded + } + _, data, _, err := prepareStoredConversationForStorage(stored, now, s.ttl, false) + if err != nil { + return nil, err + } + filter := storage.MongoUnexpiredFilter(id, now) + filter["data"] = doc.Data + // Filter/update keys and operators are fixed above; all variable data is + // encoded as BSON scalar/array values. lgtm[go/sql-injection] + result, err := s.collection.UpdateOne(ctx, filter, bson.M{"$set": bson.M{"data": string(data)}}) + if err != nil { + return nil, fmt.Errorf("merge conversation metadata: %w", err) + } + if result.MatchedCount == 1 { + return s.Get(ctx, id) + } + if err := waitForMongoMutationRetry(ctx, attempt); err != nil { + return nil, fmt.Errorf("merge conversation metadata: %w", err) + } } - if !normalized.ExpiresAt.IsZero() { - set["expires_at"] = normalized.ExpiresAt.Unix() - } - result, err := s.collection.UpdateOne(ctx, storage.MongoUnexpiredFilter(normalized.Conversation.ID, now), bson.M{"$set": set}) - if err != nil { - return fmt.Errorf("update conversation snapshot: %w", err) - } - if result.MatchedCount == 0 { - return ErrNotFound - } - return nil + return nil, fmt.Errorf("merge conversation metadata: concurrent updates did not settle") } -// AppendItems atomically appends items to an existing, unexpired conversation -// via $push, so two concurrently completing turns cannot overwrite each -// other's exchange. +// AppendItems atomically appends items to an existing, unexpired conversation. +// Items are stored as JSON strings, so an optimistic compare-and-swap keeps id +// uniqueness and the append in the same atomic operation. func (s *MongoDBStore) AppendItems(ctx context.Context, id string, items []json.RawMessage) error { if len(items) == 0 { return nil } - update := bson.M{"$push": bson.M{"items": bson.M{"$each": itemsToStrings(items)}}} - result, err := s.collection.UpdateOne(ctx, storage.MongoUnexpiredFilter(id, time.Now()), update) - if err != nil { - return fmt.Errorf("append conversation items: %w", err) + if duplicateItemID(nil, items) != "" { + return ErrDuplicateItem } - if result.MatchedCount == 0 { - return ErrNotFound + for attempt := range mongoMutationMaxAttempts { + var doc mongoConversationDocument + // id is encoded by the driver as a BSON string value; it cannot add + // query keys or operators. lgtm[go/sql-injection] + if err := s.collection.FindOne(ctx, bson.M{"_id": id}).Decode(&doc); err != nil { + if errors.Is(err, mongo.ErrNoDocuments) { + return ErrNotFound + } + return fmt.Errorf("query conversation snapshot: %w", err) + } + now := time.Now() + if doc.ExpiresAt > 0 && doc.ExpiresAt <= now.Unix() { + return ErrNotFound + } + if duplicateItemID(itemsFromStrings(doc.Items), items) != "" { + return ErrDuplicateItem + } + filter := storage.MongoUnexpiredFilter(id, now) + filter["items"] = doc.Items + update := bson.M{"$push": bson.M{"items": bson.M{"$each": itemsToStrings(items)}}} + // Filter/update keys and operators are fixed above; all variable data is + // encoded as BSON scalar/array values. lgtm[go/sql-injection] + result, err := s.collection.UpdateOne(ctx, filter, update) + if err != nil { + return fmt.Errorf("append conversation items: %w", err) + } + if result.MatchedCount == 1 { + return nil + } + if err := waitForMongoMutationRetry(ctx, attempt); err != nil { + return fmt.Errorf("append conversation items: %w", err) + } } - return nil + return fmt.Errorf("append conversation items: concurrent updates did not settle") +} + +// DeleteItem uses an optimistic compare-and-swap because MongoDB stores each +// raw item as a JSON string to preserve its exact shape. +func (s *MongoDBStore) DeleteItem(ctx context.Context, id, targetItemID string) (*StoredConversation, error) { + for attempt := range mongoMutationMaxAttempts { + var doc mongoConversationDocument + // id is encoded by the driver as a BSON string value; it cannot add + // query keys or operators. lgtm[go/sql-injection] + if err := s.collection.FindOne(ctx, bson.M{"_id": id}).Decode(&doc); err != nil { + if errors.Is(err, mongo.ErrNoDocuments) { + return nil, ErrNotFound + } + return nil, fmt.Errorf("query conversation snapshot: %w", err) + } + now := time.Now() + if doc.ExpiresAt > 0 && doc.ExpiresAt <= now.Unix() { + return nil, ErrNotFound + } + index := -1 + for i, raw := range doc.Items { + if itemID(json.RawMessage(raw)) == targetItemID { + index = i + break + } + } + if index < 0 { + return nil, ErrItemNotFound + } + updatedItems := append([]string(nil), doc.Items[:index]...) + updatedItems = append(updatedItems, doc.Items[index+1:]...) + filter := storage.MongoUnexpiredFilter(id, now) + filter["items"] = doc.Items + // Filter/update keys and operators are fixed above; all variable data is + // encoded as BSON scalar/array values. lgtm[go/sql-injection] + result, err := s.collection.UpdateOne(ctx, filter, bson.M{"$set": bson.M{"items": updatedItems}}) + if err != nil { + return nil, fmt.Errorf("delete conversation item: %w", err) + } + if result.MatchedCount == 1 { + return s.Get(ctx, id) + } + if err := waitForMongoMutationRetry(ctx, attempt); err != nil { + return nil, fmt.Errorf("delete conversation item: %w", err) + } + } + return nil, fmt.Errorf("delete conversation item: concurrent updates did not settle") } // Delete removes one unexpired conversation snapshot by id. @@ -197,6 +303,25 @@ func itemsFromStrings(items []string) []json.RawMessage { return decoded } +func waitForMongoMutationRetry(ctx context.Context, attempt int) error { + if attempt >= mongoMutationMaxAttempts-1 { + return nil + } + shift := min(attempt, 5) + delay := min(mongoMutationInitialBackoff< $6) - `, string(data), string(items), storage.UnixOrZero(normalized.StoredAt), storage.UnixOrZero(normalized.ExpiresAt), normalized.Conversation.ID, now.Unix()) + UPDATE conversation_snapshots SET data = jsonb_set( + data::jsonb, + '{conversation,metadata}', + COALESCE(data::jsonb #> '{conversation,metadata}', '{}'::jsonb) || $2::jsonb + )::text + WHERE id = $1 AND (expires_at = 0 OR expires_at > $3) + AND ( + SELECT COUNT(*) + FROM jsonb_object_keys( + COALESCE(data::jsonb #> '{conversation,metadata}', '{}'::jsonb) || $2::jsonb + ) + ) <= $4 + `, id, string(patch), time.Now().Unix(), core.MaxConversationMetadataPairs) if err != nil { - return fmt.Errorf("update conversation snapshot: %w", err) + return nil, fmt.Errorf("merge conversation metadata: %w", err) } if cmd.RowsAffected() == 0 { - return ErrNotFound + if _, getErr := s.Get(ctx, id); getErr != nil { + return nil, getErr + } + return nil, ErrMetadataLimitExceeded } - return nil + return s.Get(ctx, id) } // AppendItems atomically appends items to an existing, unexpired conversation @@ -125,6 +134,9 @@ func (s *PostgreSQLStore) AppendItems(ctx context.Context, id string, items []js if len(items) == 0 { return nil } + if duplicateItemID(nil, items) != "" { + return ErrDuplicateItem + } appended, err := json.Marshal(items) if err != nil { return fmt.Errorf("marshal conversation items: %w", err) @@ -132,16 +144,60 @@ func (s *PostgreSQLStore) AppendItems(ctx context.Context, id string, items []js cmd, err := s.pool.Exec(ctx, ` UPDATE conversation_snapshots SET items = items || $2::jsonb WHERE id = $1 AND (expires_at = 0 OR expires_at > $3) + AND NOT EXISTS ( + SELECT 1 + FROM jsonb_array_elements(items) AS existing(value) + JOIN jsonb_array_elements($2::jsonb) AS added(value) + ON existing.value ->> 'id' = added.value ->> 'id' + ) `, id, string(appended), time.Now().Unix()) if err != nil { return fmt.Errorf("append conversation items: %w", err) } if cmd.RowsAffected() == 0 { + stored, getErr := s.Get(ctx, id) + if getErr != nil { + return getErr + } + if duplicateItemID(stored.Items, items) != "" { + return ErrDuplicateItem + } return ErrNotFound } return nil } +// DeleteItem atomically removes the first matching item from the JSONB array. +func (s *PostgreSQLStore) DeleteItem(ctx context.Context, id, targetItemID string) (*StoredConversation, error) { + cmd, err := s.pool.Exec(ctx, ` + UPDATE conversation_snapshots c + SET items = c.items - ( + SELECT (element.ordinality - 1)::int + FROM jsonb_array_elements(c.items) WITH ORDINALITY AS element(value, ordinality) + WHERE element.value ->> 'id' = $2 + ORDER BY element.ordinality + LIMIT 1 + ) + WHERE c.id = $1 + AND (c.expires_at = 0 OR c.expires_at > $3) + AND EXISTS ( + SELECT 1 + FROM jsonb_array_elements(c.items) AS element(value) + WHERE element.value ->> 'id' = $2 + ) + `, id, targetItemID, time.Now().Unix()) + if err != nil { + return nil, fmt.Errorf("delete conversation item: %w", err) + } + if cmd.RowsAffected() == 0 { + if _, getErr := s.Get(ctx, id); getErr != nil { + return nil, getErr + } + return nil, ErrItemNotFound + } + return s.Get(ctx, id) +} + // Delete removes one unexpired conversation snapshot by id. func (s *PostgreSQLStore) Delete(ctx context.Context, id string) error { cmd, err := s.pool.Exec(ctx, ` diff --git a/internal/conversationstore/store_sqlite.go b/internal/conversationstore/store_sqlite.go index 8f9df96b8..3064f2f93 100644 --- a/internal/conversationstore/store_sqlite.go +++ b/internal/conversationstore/store_sqlite.go @@ -11,6 +11,7 @@ import ( "github.com/goccy/go-json" + "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/storage" ) @@ -94,35 +95,41 @@ func (s *SQLiteStore) Get(ctx context.Context, id string) (*StoredConversation, `, id), sql.ErrNoRows) } -// Update replaces an existing, unexpired conversation snapshot including its -// items. Zero StoredAt or ExpiresAt values preserve the stored retention columns. -func (s *SQLiteStore) Update(ctx context.Context, conversation *StoredConversation) error { - now := time.Now().UTC() - normalized, data, items, err := prepareStoredConversationForStorage(conversation, now, s.ttl, false) +// MergeMetadata atomically overlays metadata in the snapshot JSON while +// leaving the independently stored item array untouched. +func (s *SQLiteStore) MergeMetadata(ctx context.Context, id string, metadata map[string]string) (*StoredConversation, error) { + patch, err := json.Marshal(metadata) if err != nil { - return err + return nil, fmt.Errorf("marshal conversation metadata: %w", err) } - storedAt := storage.UnixOrZero(normalized.StoredAt) - expiresAt := storage.UnixOrZero(normalized.ExpiresAt) + now := time.Now().Unix() result, err := s.db.ExecContext(ctx, ` - UPDATE conversation_snapshots SET - data = ?, - items = ?, - stored_at = CASE WHEN ? = 0 THEN stored_at ELSE ? END, - expires_at = CASE WHEN ? = 0 THEN expires_at ELSE ? END + UPDATE conversation_snapshots SET data = json_set( + data, + '$.conversation.metadata', + json_patch(COALESCE(json_extract(data, '$.conversation.metadata'), json('{}')), json(?)) + ) WHERE id = ? AND (expires_at = 0 OR expires_at > ?) - `, string(data), string(items), storedAt, storedAt, expiresAt, expiresAt, normalized.Conversation.ID, now.Unix()) + AND ( + SELECT COUNT(*) FROM json_each( + json_patch(COALESCE(json_extract(data, '$.conversation.metadata'), json('{}')), json(?)) + ) + ) <= ? + `, string(patch), id, now, string(patch), core.MaxConversationMetadataPairs) if err != nil { - return fmt.Errorf("update conversation snapshot: %w", err) + return nil, fmt.Errorf("merge conversation metadata: %w", err) } affected, err := result.RowsAffected() if err != nil { - return fmt.Errorf("read update rows affected: %w", err) + return nil, fmt.Errorf("read metadata merge rows affected: %w", err) } if affected == 0 { - return ErrNotFound + if _, getErr := s.Get(ctx, id); getErr != nil { + return nil, getErr + } + return nil, ErrMetadataLimitExceeded } - return nil + return s.Get(ctx, id) } // AppendItems atomically appends items to an existing, unexpired conversation. @@ -132,20 +139,33 @@ func (s *SQLiteStore) AppendItems(ctx context.Context, id string, items []json.R if len(items) == 0 { return nil } + if duplicateItemID(nil, items) != "" { + return ErrDuplicateItem + } var expr strings.Builder expr.WriteString("json_insert(items") - args := make([]any, 0, len(items)+2) + args := make([]any, 0, len(items)+3) for _, item := range items { expr.WriteString(", '$[#]', json(?)") args = append(args, string(item)) } expr.WriteString(")") - args = append(args, id, time.Now().Unix()) + added, err := json.Marshal(items) + if err != nil { + return fmt.Errorf("marshal conversation items: %w", err) + } + args = append(args, id, time.Now().Unix(), string(added)) result, err := s.db.ExecContext(ctx, "UPDATE conversation_snapshots SET items = "+expr.String()+ - " WHERE id = ? AND (expires_at = 0 OR expires_at > ?)", args...) + ` WHERE id = ? AND (expires_at = 0 OR expires_at > ?) + AND NOT EXISTS ( + SELECT 1 + FROM json_each(items) AS existing + JOIN json_each(json(?)) AS added + ON json_extract(existing.value, '$.id') = json_extract(added.value, '$.id') + )`, args...) if err != nil { return fmt.Errorf("append conversation items: %w", err) } @@ -154,11 +174,50 @@ func (s *SQLiteStore) AppendItems(ctx context.Context, id string, items []json.R return fmt.Errorf("read append rows affected: %w", err) } if affected == 0 { + stored, getErr := s.Get(ctx, id) + if getErr != nil { + return getErr + } + if duplicateItemID(stored.Items, items) != "" { + return ErrDuplicateItem + } return ErrNotFound } return nil } +// DeleteItem atomically removes the first item with the requested id. +func (s *SQLiteStore) DeleteItem(ctx context.Context, id, targetItemID string) (*StoredConversation, error) { + now := time.Now().Unix() + result, err := s.db.ExecContext(ctx, ` + UPDATE conversation_snapshots + SET items = json_remove(items, '$[' || ( + SELECT key FROM json_each(items) + WHERE json_extract(value, '$.id') = ? + LIMIT 1 + ) || ']') + WHERE id = ? AND (expires_at = 0 OR expires_at > ?) + AND EXISTS ( + SELECT 1 FROM json_each(items) + WHERE json_extract(value, '$.id') = ? + ) + `, targetItemID, id, now, targetItemID) + if err != nil { + return nil, fmt.Errorf("delete conversation item: %w", err) + } + affected, err := result.RowsAffected() + if err != nil { + return nil, fmt.Errorf("read item delete rows affected: %w", err) + } + if affected == 0 { + if _, getErr := s.Get(ctx, id); getErr != nil { + return nil, getErr + } + return nil, ErrItemNotFound + } + return s.Get(ctx, id) +} + // Delete removes one unexpired conversation snapshot by id. func (s *SQLiteStore) Delete(ctx context.Context, id string) error { result, err := s.db.ExecContext(ctx, ` diff --git a/internal/conversationstore/store_sqlite_test.go b/internal/conversationstore/store_sqlite_test.go index fdc4bba63..b39892a25 100644 --- a/internal/conversationstore/store_sqlite_test.go +++ b/internal/conversationstore/store_sqlite_test.go @@ -3,6 +3,7 @@ package conversationstore import ( "context" "errors" + "fmt" "path/filepath" "strings" "testing" @@ -126,38 +127,81 @@ func TestSQLiteConversationAppendItemsMissingReturnsNotFound(t *testing.T) { } } -func TestSQLiteConversationUpdateReplacesItemsAndPreservesRetention(t *testing.T) { +func TestSQLiteConversationAppendItemsRejectsDuplicateID(t *testing.T) { store := newSQLiteTestStore(t) ctx := context.Background() - - if err := store.Create(ctx, testStoredConversation("conv-1")); err != nil { + conv := testStoredConversation("conv-duplicate-items") + conv.Items = []json.RawMessage{json.RawMessage(`{"id":"msg_existing","type":"message"}`)} + if err := store.Create(ctx, conv); err != nil { t.Fatalf("create: %v", err) } - created, err := store.Get(ctx, "conv-1") + + err := store.AppendItems(ctx, conv.Conversation.ID, []json.RawMessage{ + json.RawMessage(`{"id":"msg_existing","type":"message","content":"duplicate"}`), + }) + if !errors.Is(err, ErrDuplicateItem) { + t.Fatalf("append duplicate err = %v, want ErrDuplicateItem", err) + } + got, err := store.Get(ctx, conv.Conversation.ID) if err != nil { - t.Fatalf("get created: %v", err) + t.Fatalf("get: %v", err) } - - updated := testStoredConversation("conv-1") - updated.Conversation.Metadata = map[string]string{"topic": "changed"} - updated.Items = []json.RawMessage{json.RawMessage(`{"type":"message","content":"replaced"}`)} - if err := store.Update(ctx, updated); err != nil { - t.Fatalf("update: %v", err) + if len(got.Items) != 1 { + t.Fatalf("stored items = %d, want unchanged length 1", len(got.Items)) } +} - got, err := store.Get(ctx, "conv-1") +func TestSQLiteConversationMergeMetadataAndDeleteItem(t *testing.T) { + store := newSQLiteTestStore(t) + ctx := context.Background() + conv := testStoredConversation("conv-items") + conv.Conversation.Metadata = map[string]string{"existing": "kept"} + conv.Items = []json.RawMessage{ + json.RawMessage(`{"id":"msg_1","type":"message"}`), + json.RawMessage(`{"id":"msg_2","type":"message"}`), + } + if err := store.Create(ctx, conv); err != nil { + t.Fatalf("create: %v", err) + } + merged, err := store.MergeMetadata(ctx, "conv-items", map[string]string{"new": "value"}) + if err != nil { + t.Fatalf("merge metadata: %v", err) + } + if merged.Conversation.Metadata["existing"] != "kept" || merged.Conversation.Metadata["new"] != "value" || len(merged.Items) != 2 { + t.Fatalf("merged = %+v, want merged metadata and preserved items", merged) + } + updated, err := store.DeleteItem(ctx, "conv-items", "msg_1") if err != nil { - t.Fatalf("get updated: %v", err) + t.Fatalf("delete item: %v", err) } - if got.Conversation.Metadata["topic"] != "changed" { - t.Fatalf("metadata = %v, want topic=changed", got.Conversation.Metadata) + if len(updated.Items) != 1 || itemID(updated.Items[0]) != "msg_2" { + t.Fatalf("items = %s, want msg_2 only", updated.Items) } - if len(got.Items) != 1 || !strings.Contains(string(got.Items[0]), "replaced") { - t.Fatalf("items = %v, want replaced item only", got.Items) + if _, err := store.DeleteItem(ctx, "conv-items", "missing"); !errors.Is(err, ErrItemNotFound) { + t.Fatalf("delete missing item error = %v, want ErrItemNotFound", err) + } +} + +func TestSQLiteConversationMergeMetadataRejectsOversizedResult(t *testing.T) { + store := newSQLiteTestStore(t) + conv := testStoredConversation("conv_sqlite_metadata_limit") + conv.Conversation.Metadata = make(map[string]string, core.MaxConversationMetadataPairs) + for index := range core.MaxConversationMetadataPairs { + conv.Conversation.Metadata[fmt.Sprintf("key_%d", index)] = "value" + } + if err := store.Create(context.Background(), conv); err != nil { + t.Fatalf("Create() error = %v", err) + } + + if _, err := store.MergeMetadata(context.Background(), conv.Conversation.ID, map[string]string{"extra": "value"}); !errors.Is(err, ErrMetadataLimitExceeded) { + t.Fatalf("MergeMetadata() error = %v, want ErrMetadataLimitExceeded", err) + } + got, err := store.Get(context.Background(), conv.Conversation.ID) + if err != nil { + t.Fatalf("Get() error = %v", err) } - if !got.StoredAt.Equal(created.StoredAt) || !got.ExpiresAt.Equal(created.ExpiresAt) { - t.Fatalf("retention changed: stored %v→%v expires %v→%v", - created.StoredAt, got.StoredAt, created.ExpiresAt, got.ExpiresAt) + if len(got.Conversation.Metadata) != core.MaxConversationMetadataPairs { + t.Fatalf("metadata size = %d, want %d", len(got.Conversation.Metadata), core.MaxConversationMetadataPairs) } } diff --git a/internal/core/conversations.go b/internal/core/conversations.go index f564995e8..b7fe1ad2a 100644 --- a/internal/core/conversations.go +++ b/internal/core/conversations.go @@ -21,8 +21,10 @@ const ( // MaxConversationInitialItems caps the items array accepted by // POST /v1/conversations. MaxConversationInitialItems = 20 + // MaxConversationMetadataPairs caps the final metadata object after create + // or update. + MaxConversationMetadataPairs = 16 - maxConversationMetadataPairs = 16 maxConversationMetadataKeyLength = 64 maxConversationMetadataValueLength = 512 ) @@ -58,6 +60,32 @@ type ConversationUpdateRequest struct { Metadata *map[string]string `json:"metadata" binding:"required"` } +// ConversationItemCreateRequest is accepted by +// POST /v1/conversations/{id}/items. +type ConversationItemCreateRequest struct { + Items []json.RawMessage `json:"items" binding:"required" swaggertype:"array,object"` +} + +// ConversationItemListResponse is returned when creating or listing +// conversation items. Data stays raw so new OpenAI item variants can pass +// through the gateway without waiting for a typed Go model. +type ConversationItemListResponse struct { + Object string `json:"object"` + Data []json.RawMessage `json:"data" swaggertype:"array,object"` + FirstID *string `json:"first_id" extensions:"x-nullable"` + LastID *string `json:"last_id" extensions:"x-nullable"` + HasMore bool `json:"has_more"` +} + +// ConversationItemListParams contains query parameters accepted by +// GET /v1/conversations/{id}/items. +type ConversationItemListParams struct { + After string + Include []string + Limit int + Order string +} + // DecodeConversationCreateRequest parses a conversation create body. An empty // body is treated as an empty request (a conversation with no items/metadata). func DecodeConversationCreateRequest(data []byte) (*ConversationCreateRequest, error) { @@ -84,13 +112,25 @@ func DecodeConversationUpdateRequest(data []byte) (*ConversationUpdateRequest, e return req, nil } +// DecodeConversationItemCreateRequest parses a conversation item batch. +func DecodeConversationItemCreateRequest(data []byte) (*ConversationItemCreateRequest, error) { + req := &ConversationItemCreateRequest{} + if len(bytes.TrimSpace(data)) == 0 { + return req, nil + } + if err := json.Unmarshal(data, req); err != nil { + return nil, err + } + return req, nil +} + // ValidateConversationMetadata enforces the OpenAI metadata limits (at most 16 // pairs, keys up to 64 characters, values up to 512 characters). It returns nil // when the metadata is acceptable. func ValidateConversationMetadata(metadata map[string]string) *GatewayError { - if len(metadata) > maxConversationMetadataPairs { + if len(metadata) > MaxConversationMetadataPairs { return NewInvalidRequestError( - fmt.Sprintf("metadata supports at most %d key-value pairs", maxConversationMetadataPairs), nil, + fmt.Sprintf("metadata supports at most %d key-value pairs", MaxConversationMetadataPairs), nil, ).WithParam("metadata") } for key, value := range metadata { diff --git a/internal/core/conversations_test.go b/internal/core/conversations_test.go index 5a59e6809..6c1331c40 100644 --- a/internal/core/conversations_test.go +++ b/internal/core/conversations_test.go @@ -69,8 +69,8 @@ func TestValidateConversationMetadata(t *testing.T) { }) t.Run("too many pairs", func(t *testing.T) { - metadata := make(map[string]string, maxConversationMetadataPairs+1) - for i := 0; i <= maxConversationMetadataPairs; i++ { + metadata := make(map[string]string, MaxConversationMetadataPairs+1) + for i := 0; i <= MaxConversationMetadataPairs; i++ { metadata[string(rune('a'+i))] = "v" } err := ValidateConversationMetadata(metadata) diff --git a/internal/core/endpoints_test.go b/internal/core/endpoints_test.go index b8c30a7fd..34c6df5c0 100644 --- a/internal/core/endpoints_test.go +++ b/internal/core/endpoints_test.go @@ -20,6 +20,8 @@ func TestDescribeEndpointPath(t *testing.T) { {path: "/v1/responses/resp_1/input_items", managed: true, dialect: "openai_compat", operation: OperationResponses, bodyMode: BodyModeNone, interaction: true}, {path: "/v1/conversations", managed: true, dialect: "openai_compat", operation: OperationConversations, bodyMode: BodyModeNone, interaction: true}, {path: "/v1/conversations/conv_1", managed: true, dialect: "openai_compat", operation: OperationConversations, bodyMode: BodyModeNone, interaction: true}, + {path: "/v1/conversations/conv_1/items", managed: true, dialect: "openai_compat", operation: OperationConversations, bodyMode: BodyModeNone, interaction: true}, + {path: "/v1/conversations/conv_1/items/msg_1", managed: true, dialect: "openai_compat", operation: OperationConversations, bodyMode: BodyModeNone, interaction: true}, {path: "/v1/batches", managed: true, dialect: "openai_compat", operation: OperationBatches, bodyMode: BodyModeNone, interaction: true}, {path: "/v1/messages/batches", managed: true, dialect: "anthropic", operation: OperationBatches, bodyMode: BodyModeNone, interaction: true}, {path: "/v1/messages/batches/msgbatch_1/results", managed: true, dialect: "anthropic", operation: OperationBatches, bodyMode: BodyModeNone, interaction: true}, @@ -75,6 +77,10 @@ func TestDescribeEndpoint_UsesMethodForBodyMode(t *testing.T) { {method: http.MethodPost, path: "/v1/conversations/conv_1", bodyMode: BodyModeJSON}, {method: http.MethodGet, path: "/v1/conversations/conv_1", bodyMode: BodyModeNone}, {method: http.MethodDelete, path: "/v1/conversations/conv_1", bodyMode: BodyModeNone}, + {method: http.MethodPost, path: "/v1/conversations/conv_1/items", bodyMode: BodyModeJSON}, + {method: http.MethodGet, path: "/v1/conversations/conv_1/items", bodyMode: BodyModeNone}, + {method: http.MethodGet, path: "/v1/conversations/conv_1/items/msg_1", bodyMode: BodyModeNone}, + {method: http.MethodDelete, path: "/v1/conversations/conv_1/items/msg_1", bodyMode: BodyModeNone}, {method: http.MethodPost, path: "/v1/files", bodyMode: BodyModeMultipart}, {method: http.MethodPost, path: "/v1/files/", bodyMode: BodyModeMultipart}, {method: http.MethodGet, path: "/v1/files/file_1", bodyMode: BodyModeNone}, diff --git a/internal/core/responses.go b/internal/core/responses.go index 44728d2b4..2fcca8c65 100644 --- a/internal/core/responses.go +++ b/internal/core/responses.go @@ -205,6 +205,10 @@ type ResponsesOutputItem struct { Name string `json:"name,omitempty"` Arguments string `json:"arguments,omitempty"` Content []ResponsesContentItem `json:"content,omitempty"` + // Preserve fields belonging to newer or variant-specific output items, such + // as reasoning.summary, reasoning.encrypted_content, hosted-tool payloads, + // and provider extensions. Conversation replay depends on these fields. + ExtraFields UnknownJSONFields `json:"-" swaggerignore:"true"` } // ResponsesContentItem represents a content item in the output. diff --git a/internal/core/responses_json.go b/internal/core/responses_json.go index 95e2674e9..37e3503a2 100644 --- a/internal/core/responses_json.go +++ b/internal/core/responses_json.go @@ -14,6 +14,7 @@ import ( var ( responsesRequestFields = jsonFieldNames(ResponsesRequest{}) responsesUtilityRequestFields = jsonFieldNames(ResponseInputTokensRequest{}) + responsesOutputItemFields = jsonFieldNames(ResponsesOutputItem{}) ) // responsesExtrasAndInput finishes a responses-shaped decode: it captures @@ -318,6 +319,31 @@ func (e ResponsesInputElement) MarshalJSON() ([]byte, error) { } } +// UnmarshalJSON preserves variant-specific Responses output item fields. This +// is required for lossless Responses passthrough and for replaying reasoning +// and hosted-tool items from a gateway-managed conversation. +func (i *ResponsesOutputItem) UnmarshalJSON(data []byte) error { + type alias ResponsesOutputItem + var decoded alias + if err := json.Unmarshal(data, &decoded); err != nil { + return err + } + extraFields, err := extractUnknownJSONFields(data, responsesOutputItemFields...) + if err != nil { + return err + } + *i = ResponsesOutputItem(decoded) + i.ExtraFields = extraFields + return nil +} + +// MarshalJSON emits typed output fields together with every unknown field +// retained during decoding. +func (i ResponsesOutputItem) MarshalJSON() ([]byte, error) { + type alias ResponsesOutputItem + return marshalWithUnknownJSONFields(alias(i), i.ExtraFields) +} + // cloneRawMessage returns a detached, whitespace-trimmed copy of a raw JSON // value so stored Raw fields stay independent of the decoder's backing buffer. func cloneRawMessage(data []byte) json.RawMessage { diff --git a/internal/core/responses_json_test.go b/internal/core/responses_json_test.go index 457347ea8..02e273bd0 100644 --- a/internal/core/responses_json_test.go +++ b/internal/core/responses_json_test.go @@ -707,6 +707,28 @@ func TestResponsesInputElementMarshalJSON_MergesRawUnknownItemExtras(t *testing. } } +func TestResponsesOutputItemJSONPreservesReasoningFields(t *testing.T) { + raw := []byte(`{"id":"rs_123","type":"reasoning","summary":[],"encrypted_content":"opaque","provider_trace":{"id":"trace_1"}}`) + var item ResponsesOutputItem + if err := json.Unmarshal(raw, &item); err != nil { + t.Fatalf("json.Unmarshal() error = %v", err) + } + body, err := json.Marshal(item) + if err != nil { + t.Fatalf("json.Marshal() error = %v", err) + } + var decoded map[string]any + if err := json.Unmarshal(body, &decoded); err != nil { + t.Fatalf("json.Unmarshal(round trip) error = %v", err) + } + if _, ok := decoded["summary"].([]any); !ok || decoded["encrypted_content"] != "opaque" { + t.Fatalf("round-tripped item = %#v, want reasoning fields", decoded) + } + if trace, ok := decoded["provider_trace"].(map[string]any); !ok || trace["id"] != "trace_1" { + t.Fatalf("provider_trace = %#v, want trace_1", decoded["provider_trace"]) + } +} + func TestResponsesRequestJSON_PreservesVariantSpecificUnknownFields(t *testing.T) { var req ResponsesRequest if err := json.Unmarshal([]byte(`{ diff --git a/internal/server/conversation_handlers.go b/internal/server/conversation_handlers.go index fe6fc2232..65d131385 100644 --- a/internal/server/conversation_handlers.go +++ b/internal/server/conversation_handlers.go @@ -44,7 +44,7 @@ func (h *Handler) GetConversation(c *echo.Context) error { // UpdateConversation handles POST /v1/conversations/{id}. // // @Summary Update a conversation -// @Description Replaces the conversation metadata in full. +// @Description Merges supplied keys into the conversation metadata. // @Tags conversations // @Accept json // @Produce json @@ -75,3 +75,77 @@ func (h *Handler) UpdateConversation(c *echo.Context) error { func (h *Handler) DeleteConversation(c *echo.Context) error { return h.conversations().DeleteConversation(c) } + +// CreateConversationItems handles POST /v1/conversations/{id}/items. +// +// @Summary Create conversation items +// @Tags conversations +// @Accept json +// @Produce json +// @Security BearerAuth +// @Param id path string true "Conversation ID" +// @Param include query []string false "Additional fields to include" +// @Param request body core.ConversationItemCreateRequest true "Conversation items" +// @Success 200 {object} core.ConversationItemListResponse +// @Failure 400 {object} core.OpenAIErrorEnvelope +// @Failure 401 {object} core.OpenAIErrorEnvelope +// @Failure 404 {object} core.OpenAIErrorEnvelope +// @Router /v1/conversations/{id}/items [post] +func (h *Handler) CreateConversationItems(c *echo.Context) error { + return h.conversations().CreateConversationItems(c) +} + +// ListConversationItems handles GET /v1/conversations/{id}/items. +// +// @Summary List conversation items +// @Tags conversations +// @Produce json +// @Security BearerAuth +// @Param id path string true "Conversation ID" +// @Param after query string false "Return items after this item ID" +// @Param include query []string false "Additional fields to include" +// @Param limit query int false "Maximum items" minimum(1) maximum(100) default(20) +// @Param order query string false "Sort order" Enums(asc,desc) default(desc) +// @Success 200 {object} core.ConversationItemListResponse +// @Failure 400 {object} core.OpenAIErrorEnvelope +// @Failure 401 {object} core.OpenAIErrorEnvelope +// @Failure 404 {object} core.OpenAIErrorEnvelope +// @Router /v1/conversations/{id}/items [get] +func (h *Handler) ListConversationItems(c *echo.Context) error { + return h.conversations().ListConversationItems(c) +} + +// GetConversationItem handles GET /v1/conversations/{id}/items/{item_id}. +// +// @Summary Get a conversation item +// @Tags conversations +// @Produce json +// @Security BearerAuth +// @Param id path string true "Conversation ID" +// @Param item_id path string true "Conversation item ID" +// @Param include query []string false "Additional fields to include" +// @Success 200 {object} object +// @Failure 400 {object} core.OpenAIErrorEnvelope +// @Failure 401 {object} core.OpenAIErrorEnvelope +// @Failure 404 {object} core.OpenAIErrorEnvelope +// @Router /v1/conversations/{id}/items/{item_id} [get] +func (h *Handler) GetConversationItem(c *echo.Context) error { + return h.conversations().GetConversationItem(c) +} + +// DeleteConversationItem handles DELETE /v1/conversations/{id}/items/{item_id}. +// +// @Summary Delete a conversation item +// @Tags conversations +// @Produce json +// @Security BearerAuth +// @Param id path string true "Conversation ID" +// @Param item_id path string true "Conversation item ID" +// @Success 200 {object} core.Conversation +// @Failure 400 {object} core.OpenAIErrorEnvelope +// @Failure 401 {object} core.OpenAIErrorEnvelope +// @Failure 404 {object} core.OpenAIErrorEnvelope +// @Router /v1/conversations/{id}/items/{item_id} [delete] +func (h *Handler) DeleteConversationItem(c *echo.Context) error { + return h.conversations().DeleteConversationItem(c) +} diff --git a/internal/server/conversation_handlers_test.go b/internal/server/conversation_handlers_test.go index 7047d6ad1..cb18b1db7 100644 --- a/internal/server/conversation_handlers_test.go +++ b/internal/server/conversation_handlers_test.go @@ -7,6 +7,7 @@ import ( "net/http" "net/http/httptest" "strings" + "sync" "testing" "github.com/enterpilot/gomodel/internal/core" @@ -95,7 +96,7 @@ func TestConversationGetRoundTrip(t *testing.T) { } } -func TestConversationUpdateReplacesMetadata(t *testing.T) { +func TestConversationUpdateMergesMetadata(t *testing.T) { srv := New(&mockProvider{}, nil) created := createConversation(t, srv, `{"metadata":{"old":"value","keep":"gone"}}`) @@ -114,8 +115,319 @@ func TestConversationUpdateReplacesMetadata(t *testing.T) { if updated.Metadata["new"] != "value" { t.Fatalf("metadata[new] = %q, want value", updated.Metadata["new"]) } - if _, ok := updated.Metadata["old"]; ok { - t.Fatal("metadata still carries replaced key 'old'") + if updated.Metadata["old"] != "value" || updated.Metadata["keep"] != "gone" { + t.Fatalf("metadata = %v, want existing keys preserved", updated.Metadata) + } +} + +func TestConversationUpdateRejectsOversizedMergedMetadata(t *testing.T) { + srv := New(&mockProvider{}, nil) + metadata := make([]string, core.MaxConversationMetadataPairs) + for index := range metadata { + metadata[index] = fmt.Sprintf(`"key_%d":"value"`, index) + } + created := createConversation(t, srv, `{"metadata":{`+strings.Join(metadata, ",")+`}}`) + + req := httptest.NewRequest(http.MethodPost, "/v1/conversations/"+created.ID, + strings.NewReader(`{"metadata":{"extra":"value"}}`)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + srv.ServeHTTP(rec, req) + if rec.Code != http.StatusBadRequest { + t.Fatalf("update status = %d (%s), want 400", rec.Code, rec.Body.String()) + } + var envelope core.OpenAIErrorEnvelope + if err := json.Unmarshal(rec.Body.Bytes(), &envelope); err != nil { + t.Fatalf("decode error: %v", err) + } + if envelope.Error.Param == nil || *envelope.Error.Param != "metadata" || envelope.Error.Code == nil || *envelope.Error.Code != "metadata_max_properties_exceeded" { + t.Fatalf("error = %+v, want metadata_max_properties_exceeded", envelope.Error) + } +} + +func TestConversationItemsLifecycleAndPagination(t *testing.T) { + srv := New(&mockProvider{}, nil) + created := createConversation(t, srv, `{"items":[ + {"type":"message","role":"developer","content":"first"}, + {"type":"message","role":"user","content":"second"}, + {"type":"message","role":"assistant","content":"third"} + ]}`) + + listReq := httptest.NewRequest(http.MethodGet, "/v1/conversations/"+created.ID+"/items?order=asc&limit=2", nil) + listRec := httptest.NewRecorder() + srv.ServeHTTP(listRec, listReq) + if listRec.Code != http.StatusOK { + t.Fatalf("list status = %d (%s), want 200", listRec.Code, listRec.Body.String()) + } + var firstPage core.ConversationItemListResponse + if err := json.Unmarshal(listRec.Body.Bytes(), &firstPage); err != nil { + t.Fatalf("decode first page: %v", err) + } + if firstPage.Object != "list" || len(firstPage.Data) != 2 || !firstPage.HasMore { + t.Fatalf("first page = %+v, want two items and has_more", firstPage) + } + if firstPage.FirstID == nil || firstPage.LastID == nil { + t.Fatalf("first/last id = %v/%v, want stable ids", firstPage.FirstID, firstPage.LastID) + } + var firstItem map[string]any + if err := json.Unmarshal(firstPage.Data[0], &firstItem); err != nil { + t.Fatalf("decode first item: %v", err) + } + if firstItem["role"] != "developer" || firstItem["status"] != "completed" { + t.Fatalf("first item = %#v, want normalized developer message", firstItem) + } + + secondReq := httptest.NewRequest(http.MethodGet, "/v1/conversations/"+created.ID+"/items?order=asc&limit=2&after="+*firstPage.LastID, nil) + secondRec := httptest.NewRecorder() + srv.ServeHTTP(secondRec, secondReq) + var secondPage core.ConversationItemListResponse + if err := json.Unmarshal(secondRec.Body.Bytes(), &secondPage); err != nil { + t.Fatalf("decode second page: %v", err) + } + if secondRec.Code != http.StatusOK || len(secondPage.Data) != 1 || secondPage.HasMore { + t.Fatalf("second page status/data = %d/%+v, want final item", secondRec.Code, secondPage) + } + + createReq := httptest.NewRequest(http.MethodPost, "/v1/conversations/"+created.ID+"/items", + strings.NewReader(`{"items":[{"role":"user","content":"fourth"},{"type":"reasoning","summary":[]}]}`)) + createReq.Header.Set("Content-Type", "application/json") + createRec := httptest.NewRecorder() + srv.ServeHTTP(createRec, createReq) + if createRec.Code != http.StatusOK { + t.Fatalf("create items status = %d (%s), want 200", createRec.Code, createRec.Body.String()) + } + var added core.ConversationItemListResponse + if err := json.Unmarshal(createRec.Body.Bytes(), &added); err != nil { + t.Fatalf("decode created items: %v", err) + } + if len(added.Data) != 2 || added.HasMore || added.FirstID == nil || added.LastID == nil { + t.Fatalf("created items = %+v, want two-item list", added) + } + + getReq := httptest.NewRequest(http.MethodGet, "/v1/conversations/"+created.ID+"/items/"+*added.LastID, nil) + getRec := httptest.NewRecorder() + srv.ServeHTTP(getRec, getReq) + if getRec.Code != http.StatusOK || !strings.Contains(getRec.Body.String(), `"type":"reasoning"`) { + t.Fatalf("get item status/body = %d/%s", getRec.Code, getRec.Body.String()) + } + + deleteReq := httptest.NewRequest(http.MethodDelete, "/v1/conversations/"+created.ID+"/items/"+*added.LastID, nil) + deleteRec := httptest.NewRecorder() + srv.ServeHTTP(deleteRec, deleteReq) + if deleteRec.Code != http.StatusOK { + t.Fatalf("delete item status = %d (%s), want 200", deleteRec.Code, deleteRec.Body.String()) + } + getAfterDelete := httptest.NewRecorder() + srv.ServeHTTP(getAfterDelete, getReq) + if getAfterDelete.Code != http.StatusNotFound { + t.Fatalf("get deleted item status = %d, want 404", getAfterDelete.Code) + } +} + +func TestConversationItemsConcurrentDuplicateIDHasSingleWinner(t *testing.T) { + srv := New(&mockProvider{}, nil) + created := createConversation(t, srv, `{}`) + + const writers = 24 + start := make(chan struct{}) + statuses := make(chan int, writers) + var wg sync.WaitGroup + for range writers { + wg.Go(func() { + <-start + req := httptest.NewRequest(http.MethodPost, "/v1/conversations/"+created.ID+"/items", + strings.NewReader(`{"items":[{"id":"msg_shared","type":"message","role":"user","content":"same"}]}`)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + srv.ServeHTTP(rec, req) + statuses <- rec.Code + }) + } + close(start) + wg.Wait() + close(statuses) + + counts := map[int]int{} + for status := range statuses { + counts[status]++ + } + if counts[http.StatusOK] != 1 || counts[http.StatusBadRequest] != writers-1 { + t.Fatalf("statuses = %v, want one 200 and %d 400s", counts, writers-1) + } +} + +func TestConversationCreateRejectsNonObjectItems(t *testing.T) { + srv := New(&mockProvider{}, nil) + for _, body := range []string{`{"items":[42]}`, `{"items":["hello"]}`, `{"items":[null]}`} { + req := httptest.NewRequest(http.MethodPost, "/v1/conversations", strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + srv.ServeHTTP(rec, req) + if rec.Code != http.StatusBadRequest { + t.Fatalf("body %s status = %d (%s), want 400", body, rec.Code, rec.Body.String()) + } + var envelope core.OpenAIErrorEnvelope + if err := json.Unmarshal(rec.Body.Bytes(), &envelope); err != nil { + t.Fatalf("decode error: %v", err) + } + if envelope.Error.Param == nil || *envelope.Error.Param != "items[0]" { + t.Fatalf("body %s param = %v, want items[0]", body, envelope.Error.Param) + } + } +} + +func TestConversationCreateRejectsNullRequiredFunctionFields(t *testing.T) { + tests := []struct { + name string + item string + }{ + { + name: "function call arguments", + item: `{"type":"function_call","call_id":"call_1","name":"lookup","arguments":null}`, + }, + { + name: "function call output", + item: `{"type":"function_call_output","call_id":"call_1","output":null}`, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + srv := New(&mockProvider{}, nil) + req := httptest.NewRequest(http.MethodPost, "/v1/conversations", + strings.NewReader(`{"items":[`+tt.item+`]}`)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + srv.ServeHTTP(rec, req) + if rec.Code != http.StatusBadRequest { + t.Fatalf("status = %d (%s), want 400", rec.Code, rec.Body.String()) + } + var envelope core.OpenAIErrorEnvelope + if err := json.Unmarshal(rec.Body.Bytes(), &envelope); err != nil { + t.Fatalf("decode error: %v", err) + } + if envelope.Error.Param == nil || *envelope.Error.Param != "items[0]" { + t.Fatalf("param = %v, want items[0]", envelope.Error.Param) + } + }) + } +} + +func TestPaginateConversationItemsCapsDirectLimit(t *testing.T) { + items := make([]json.RawMessage, maxCursorListLimit+1) + for index := range items { + items[index] = json.RawMessage(fmt.Sprintf( + `{"id":"msg_%03d","type":"message","role":"user","content":"item %d"}`, + index, index, + )) + } + + page, err := paginateConversationItems(items, core.ConversationItemListParams{ + Limit: maxCursorListLimit + 1, + Order: "asc", + }) + if err != nil { + t.Fatalf("paginate: %v", err) + } + if len(page.Data) != maxCursorListLimit || !page.HasMore { + t.Fatalf("page has %d items and has_more=%v, want %d and true", len(page.Data), page.HasMore, maxCursorListLimit) + } +} + +func TestConversationItemListEmptyUsesNullCursors(t *testing.T) { + srv := New(&mockProvider{}, nil) + created := createConversation(t, srv, `{}`) + + req := httptest.NewRequest(http.MethodGet, "/v1/conversations/"+created.ID+"/items", nil) + rec := httptest.NewRecorder() + srv.ServeHTTP(rec, req) + if rec.Code != http.StatusOK { + t.Fatalf("list status = %d (%s), want 200", rec.Code, rec.Body.String()) + } + var body map[string]any + if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { + t.Fatalf("decode list: %v", err) + } + if first, exists := body["first_id"]; !exists || first != nil { + t.Fatalf("first_id = %#v (exists %v), want null", first, exists) + } + if last, exists := body["last_id"]; !exists || last != nil { + t.Fatalf("last_id = %#v (exists %v), want null", last, exists) + } +} + +func TestConversationItemListUnknownCursorReturnsOpenAICompatible404(t *testing.T) { + srv := New(&mockProvider{}, nil) + created := createConversation(t, srv, `{"items":[{"role":"user","content":"hello"}]}`) + + req := httptest.NewRequest(http.MethodGet, "/v1/conversations/"+created.ID+"/items?after=msg_missing", nil) + rec := httptest.NewRecorder() + srv.ServeHTTP(rec, req) + if rec.Code != http.StatusNotFound { + t.Fatalf("list status = %d (%s), want 404", rec.Code, rec.Body.String()) + } + var envelope core.OpenAIErrorEnvelope + if err := json.Unmarshal(rec.Body.Bytes(), &envelope); err != nil { + t.Fatalf("decode error: %v", err) + } + if envelope.Error.Type != core.ErrorTypeInvalidRequest || envelope.Error.Param == nil || *envelope.Error.Param != "after" { + t.Fatalf("error = %+v, want invalid_request_error for after", envelope.Error) + } +} + +func TestConversationItemIncludeControlsOptionalFields(t *testing.T) { + tests := []struct { + name string + item string + include string + field string + }{ + {name: "reasoning encrypted content", item: `{"type":"reasoning","summary":[],"encrypted_content":"secret"}`, include: "reasoning.encrypted_content", field: "encrypted_content"}, + {name: "output text logprobs", item: `{"type":"message","content":[{"type":"output_text","text":"ok","logprobs":[{"token":"ok"}]}]}`, include: "message.output_text.logprobs", field: "logprobs"}, + {name: "input image URL", item: `{"type":"message","content":[{"type":"input_image","image_url":"data:image/png;base64,x"}]}`, include: "message.input_image.image_url", field: "image_url"}, + {name: "file search results", item: `{"type":"file_search_call","results":[{"file_id":"file_1"}]}`, include: "file_search_call.results", field: "results"}, + {name: "web search sources", item: `{"type":"web_search_call","action":{"sources":[{"url":"https://example.com"}]}}`, include: "web_search_call.action.sources", field: "sources"}, + {name: "code interpreter outputs", item: `{"type":"code_interpreter_call","outputs":[{"type":"logs","logs":"ok"}]}`, include: "code_interpreter_call.outputs", field: "outputs"}, + {name: "computer output image", item: `{"type":"computer_call_output","output":{"type":"computer_screenshot","image_url":"data:image/png;base64,x"}}`, include: "computer_call_output.output.image_url", field: "image_url"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + without := string(conversationItemForInclude(json.RawMessage(tt.item), nil)) + if strings.Contains(without, `"`+tt.field+`"`) { + t.Fatalf("without include = %s, want %s omitted", without, tt.field) + } + with := string(conversationItemForInclude(json.RawMessage(tt.item), []string{tt.include})) + if !strings.Contains(with, `"`+tt.field+`"`) { + t.Fatalf("with include = %s, want %s retained", with, tt.field) + } + }) + } +} + +func TestConversationItemsPreserveLargeUnknownIntegers(t *testing.T) { + srv := New(&mockProvider{}, nil) + created := createConversation(t, srv, `{"items":[{"type":"future_item","opaque_integer":9007199254740993}]}`) + + req := httptest.NewRequest(http.MethodGet, "/v1/conversations/"+created.ID+"/items?order=asc", nil) + rec := httptest.NewRecorder() + srv.ServeHTTP(rec, req) + if rec.Code != http.StatusOK { + t.Fatalf("list status = %d (%s), want 200", rec.Code, rec.Body.String()) + } + if !strings.Contains(rec.Body.String(), `"opaque_integer":9007199254740993`) { + t.Fatalf("list body = %s, want exact large integer", rec.Body.String()) + } +} + +func TestConversationItemProjectionPreservesUnknownNumericFields(t *testing.T) { + raw := json.RawMessage(`{"id":"msg_1","type":"message","role":"assistant","content":[{"type":"output_text","text":"ok","logprobs":[],"opaque_integer":9007199254740993}]}`) + projected := conversationItemForInclude(raw, nil) + if !strings.Contains(string(projected), `"opaque_integer":9007199254740993`) { + t.Fatalf("projected item = %s, want exact large integer", projected) + } + if strings.Contains(string(projected), `"logprobs"`) { + t.Fatalf("projected item = %s, want logprobs omitted", projected) } } diff --git a/internal/server/conversation_item_normalization.go b/internal/server/conversation_item_normalization.go new file mode 100644 index 000000000..14b8502b0 --- /dev/null +++ b/internal/server/conversation_item_normalization.go @@ -0,0 +1,131 @@ +package server + +import ( + "bytes" + "fmt" + "strings" + + "github.com/goccy/go-json" + "github.com/google/uuid" + + "github.com/enterpilot/gomodel/internal/core" +) + +func normalizeConversationItems(items []json.RawMessage) ([]json.RawMessage, *core.GatewayError) { + normalized := make([]json.RawMessage, 0, len(items)) + seen := make(map[string]struct{}, len(items)) + for index, raw := range items { + param := fmt.Sprintf("items[%d]", index) + item, err := decodeRawJSONObject(raw) + if err != nil { + return nil, core.NewInvalidRequestError("conversation items must be JSON objects", err).WithParam(param) + } + if err := normalizeConversationItem(item, param); err != nil { + return nil, err + } + id := rawJSONString(item, "id") + if _, exists := seen[id]; exists { + return nil, core.NewInvalidRequestError("duplicate conversation item id: "+id, nil).WithParam(param) + } + seen[id] = struct{}{} + encoded, err := json.Marshal(item) + if err != nil { + return nil, core.NewInvalidRequestError("invalid conversation item", err).WithParam(param) + } + normalized = append(normalized, encoded) + } + return normalized, nil +} + +func normalizeConversationItem(item rawJSONObject, param string) *core.GatewayError { + itemType := rawJSONString(item, "type") + if itemType == "" { + if _, hasContent := item["content"]; !hasContent && rawJSONString(item, "role") == "" { + return core.NewInvalidRequestError("conversation item type or message content is required", nil).WithParam(param) + } + itemType = "message" + _ = setRawJSONValue(item, "type", itemType) + } + + switch itemType { + case "message": + role := rawJSONString(item, "role") + if role == "" { + role = "user" + _ = setRawJSONValue(item, "role", role) + } + if role != "user" && role != "system" && role != "developer" && role != "assistant" { + return core.NewInvalidRequestError("unsupported conversation message role: "+role, nil).WithParam(param) + } + content, exists := item["content"] + if !exists || !rawJSONValuePresent(item, "content") { + return core.NewInvalidRequestError("conversation message content is required", nil).WithParam(param) + } + trimmed := bytes.TrimSpace(content) + if len(trimmed) == 0 || (trimmed[0] != '"' && trimmed[0] != '[') { + return core.NewInvalidRequestError("conversation message content must be a string or array", nil).WithParam(param) + } + normalizedContent, err := normalizeResponseInputContentRaw(content) + if err != nil { + return core.NewInvalidRequestError("invalid conversation message content", err).WithParam(param) + } + item["content"] = normalizedContent + if rawJSONString(item, "status") == "" { + _ = setRawJSONValue(item, "status", "completed") + } + case "function_call": + if rawJSONString(item, "call_id") == "" || rawJSONString(item, "name") == "" { + return core.NewInvalidRequestError("function_call requires call_id and name", nil).WithParam(param) + } + arguments, exists := item["arguments"] + if !exists || !rawJSONValuePresent(item, "arguments") { + return core.NewInvalidRequestError("function_call requires arguments", nil).WithParam(param) + } + var argumentString string + if err := json.Unmarshal(arguments, &argumentString); err != nil { + if err := setRawJSONValue(item, "arguments", string(bytes.TrimSpace(arguments))); err != nil { + return core.NewInvalidRequestError("function_call arguments must be JSON", err).WithParam(param) + } + } + if rawJSONString(item, "status") == "" { + _ = setRawJSONValue(item, "status", "completed") + } + case "function_call_output": + if rawJSONString(item, "call_id") == "" { + return core.NewInvalidRequestError("function_call_output requires call_id", nil).WithParam(param) + } + if !rawJSONValuePresent(item, "output") { + return core.NewInvalidRequestError("function_call_output requires output", nil).WithParam(param) + } + if rawJSONString(item, "status") == "" { + _ = setRawJSONValue(item, "status", "completed") + } + case "reasoning": + if !rawJSONValuePresent(item, "summary") { + item["summary"] = json.RawMessage("[]") + } + default: + // Forward-compatible item variants stay opaque; only an id is added so + // list pagination, retrieval, and deletion remain well defined. + } + + if rawJSONString(item, "id") == "" { + _ = setRawJSONValue(item, "id", generatedConversationItemID(itemType)) + } + return nil +} + +func generatedConversationItemID(itemType string) string { + prefix := "item_" + switch itemType { + case "message": + prefix = "msg_" + case "function_call": + prefix = "fc_" + case "function_call_output": + prefix = "fco_" + case "reasoning": + prefix = "rs_" + } + return prefix + strings.ReplaceAll(uuid.NewString(), "-", "") +} diff --git a/internal/server/conversation_item_projection.go b/internal/server/conversation_item_projection.go new file mode 100644 index 000000000..bac5e1973 --- /dev/null +++ b/internal/server/conversation_item_projection.go @@ -0,0 +1,140 @@ +package server + +import ( + "net/http" + + "github.com/goccy/go-json" + "github.com/labstack/echo/v5" + + "github.com/enterpilot/gomodel/internal/core" +) + +func conversationItemList(items []json.RawMessage, hasMore bool, include []string) core.ConversationItemListResponse { + data := make([]json.RawMessage, 0, len(items)) + for _, item := range items { + data = append(data, conversationItemForInclude(item, include)) + } + response := core.ConversationItemListResponse{Object: "list", Data: data, HasMore: hasMore} + if len(data) > 0 { + firstID := responseInputItemID(data[0]) + lastID := responseInputItemID(data[len(data)-1]) + response.FirstID = &firstID + response.LastID = &lastID + } + return response +} + +func paginateConversationItems(items []json.RawMessage, params core.ConversationItemListParams) (core.ConversationItemListResponse, *core.GatewayError) { + count := len(items) + start := 0 + if params.After != "" { + position := -1 + for pos := range count { + if responseInputItemID(items[orderedInputItemIndex(count, pos, params.Order)]) == params.After { + position = pos + break + } + } + if position < 0 { + return core.ConversationItemListResponse{}, core.NewInvalidRequestErrorWithStatus( + http.StatusNotFound, "No item found with id '"+params.After+"'", nil, + ).WithParam("after") + } + start = position + 1 + } + limit := params.Limit + if limit <= 0 { + limit = defaultCursorListLimit + } + if limit > maxCursorListLimit { + limit = maxCursorListLimit + } + remaining := max(count-start, 0) + hasMore := remaining > limit + remaining = min(remaining, limit) + data := make([]json.RawMessage, 0, remaining) + for offset := 0; offset < remaining; offset++ { + index := orderedInputItemIndex(count, start+offset, params.Order) + data = append(data, core.CloneRawJSON(items[index])) + } + return conversationItemList(data, hasMore, params.Include), nil +} + +func conversationItemIncludes(c *echo.Context) []string { + query := c.Request().URL.Query() + return appendQueryArray(query["include"], query["include[]"]) +} + +func conversationItemForInclude(raw json.RawMessage, include []string) json.RawMessage { + requested := make(map[string]struct{}, len(include)) + for _, value := range include { + requested[value] = struct{}{} + } + has := func(value string) bool { + _, ok := requested[value] + return ok + } + + item, err := decodeRawJSONObject(raw) + if err != nil { + return core.CloneRawJSON(raw) + } + switch rawJSONString(item, "type") { + case "reasoning": + if !has("reasoning.encrypted_content") { + delete(item, "encrypted_content") + } + case "message": + var content []json.RawMessage + if err := json.Unmarshal(item["content"], &content); err == nil { + for index, rawPart := range content { + part, err := decodeRawJSONObject(rawPart) + if err != nil { + continue + } + switch rawJSONString(part, "type") { + case "input_image": + if !has("message.input_image.image_url") { + delete(part, "image_url") + } + case "output_text": + if !has("message.output_text.logprobs") { + delete(part, "logprobs") + } + } + content[index], _ = json.Marshal(part) + } + item["content"], _ = json.Marshal(content) + } + case "file_search_call": + if !has("file_search_call.results") { + delete(item, "results") + } + case "web_search_call": + if !has("web_search_call.results") { + delete(item, "results") + } + if !has("web_search_call.action.sources") { + if action, err := decodeRawJSONObject(item["action"]); err == nil { + delete(action, "sources") + item["action"], _ = json.Marshal(action) + } + } + case "code_interpreter_call": + if !has("code_interpreter_call.outputs") { + delete(item, "outputs") + } + case "computer_call_output": + if !has("computer_call_output.output.image_url") { + if output, err := decodeRawJSONObject(item["output"]); err == nil { + delete(output, "image_url") + item["output"], _ = json.Marshal(output) + } + } + } + encoded, err := json.Marshal(item) + if err != nil { + return core.CloneRawJSON(raw) + } + return encoded +} diff --git a/internal/server/conversation_persisting_stream.go b/internal/server/conversation_persisting_stream.go new file mode 100644 index 000000000..b946c5032 --- /dev/null +++ b/internal/server/conversation_persisting_stream.go @@ -0,0 +1,158 @@ +package server + +import ( + "bufio" + "bytes" + "context" + "fmt" + "io" + "log/slog" + + "github.com/goccy/go-json" +) + +// persistingStream commits the turn before releasing the provider's terminal +// event to the client. A storage failure therefore interrupts the SSE stream +// instead of reporting response.completed with history that was not saved. +func (t *conversationTurn) persistingStream(ctx context.Context, stream io.ReadCloser) io.ReadCloser { + return &conversationPersistingStream{ + upstream: stream, + reader: bufio.NewReader(stream), + observer: &conversationStreamObserver{turn: t, ctx: context.WithoutCancel(ctx)}, + } +} + +type conversationStreamObserver struct { + turn *conversationTurn + ctx context.Context + attempted bool + err error +} + +type conversationStreamEvent struct { + Type string `json:"type"` + Response *conversationStreamResponse `json:"response"` +} + +type conversationStreamResponse struct { + ID string `json:"id"` + Output []json.RawMessage `json:"output"` +} + +func (o *conversationStreamObserver) OnEvent(event conversationStreamEvent) { + if o.attempted || (event.Type != "response.completed" && event.Type != "response.done") { + return + } + if event.Response == nil { + return + } + o.attempted = true + if _, err := o.turn.appendExchange(o.ctx, event.Response.ID, event.Response.Output); err != nil { + o.err = fmt.Errorf("append streamed conversation turn: %w", err) + } +} + +func (o *conversationStreamObserver) OnStreamClose() { + if o.err != nil { + slog.Warn("conversation stream persistence failed", "conversation_id", o.turn.id, "error", o.err) + } +} + +// conversationPersistingStream buffers one complete SSE event at a time. This +// keeps ordinary event-level streaming while ensuring no prefix of a terminal +// success event escapes before its conversation turn has been committed. +type conversationPersistingStream struct { + upstream io.ReadCloser + reader *bufio.Reader + observer *conversationStreamObserver + ready []byte + readErr error + closed bool +} + +func (s *conversationPersistingStream) Read(p []byte) (int, error) { + if len(p) == 0 { + return 0, nil + } + if s.observer.err != nil { + return 0, s.observer.err + } + for len(s.ready) == 0 { + if s.readErr != nil { + return 0, s.readErr + } + + event, err := s.readEvent() + if len(event) == 0 { + return 0, err + } + s.observeEvent(event) + if s.observer.err != nil { + return 0, s.observer.err + } + s.ready = event + s.readErr = err + } + + n := copy(p, s.ready) + s.ready = s.ready[n:] + return n, nil +} + +func (s *conversationPersistingStream) Close() error { + if s.closed { + return nil + } + s.closed = true + s.observer.OnStreamClose() + return s.upstream.Close() +} + +func (s *conversationPersistingStream) readEvent() ([]byte, error) { + var event []byte + for { + line, err := s.reader.ReadBytes('\n') + event = append(event, line...) + if bytes.Equal(line, []byte("\n")) || bytes.Equal(line, []byte("\r\n")) { + return event, err + } + if err != nil { + return event, err + } + } +} + +func (s *conversationPersistingStream) observeEvent(event []byte) { + payload, ok := conversationSSEPayload(event) + if ok { + s.observer.OnEvent(payload) + } +} + +func conversationSSEPayload(event []byte) (conversationStreamEvent, bool) { + lines := bytes.Split(event, []byte("\n")) + dataLines := make([][]byte, 0, len(lines)) + for _, line := range lines { + line = bytes.TrimSuffix(line, []byte("\r")) + if !bytes.HasPrefix(line, []byte("data:")) { + continue + } + data := line[len("data:"):] + if len(data) > 0 && data[0] == ' ' { + data = data[1:] + } + dataLines = append(dataLines, data) + } + if len(dataLines) == 0 { + return conversationStreamEvent{}, false + } + data := bytes.Join(dataLines, []byte("\n")) + if bytes.Equal(bytes.TrimSpace(data), []byte("[DONE]")) { + return conversationStreamEvent{}, false + } + var payload conversationStreamEvent + if err := json.Unmarshal(data, &payload); err != nil { + return conversationStreamEvent{}, false + } + return payload, true +} diff --git a/internal/server/conversation_persisting_stream_test.go b/internal/server/conversation_persisting_stream_test.go new file mode 100644 index 000000000..30285ecb7 --- /dev/null +++ b/internal/server/conversation_persisting_stream_test.go @@ -0,0 +1,106 @@ +package server + +import ( + "context" + "errors" + "io" + "strings" + "testing" + + "github.com/enterpilot/gomodel/internal/conversationstore" + "github.com/enterpilot/gomodel/internal/core" +) + +type maxReadCloser struct { + reader *strings.Reader + max int +} + +func (r *maxReadCloser) Read(p []byte) (int, error) { + if len(p) > r.max { + p = p[:r.max] + } + return r.reader.Read(p) +} + +func (*maxReadCloser) Close() error { return nil } + +func TestConversationPersistingStreamSuppressesEntireFragmentedCompletion(t *testing.T) { + const createdEvent = "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_fragmented\"}}\n\n" + const completedEvent = "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_fragmented\",\"output\":[]}}" + + for _, suffix := range []string{"\n\n", ""} { + name := "at EOF" + if suffix != "" { + name = "with boundary" + } + t.Run(name, func(t *testing.T) { + storeErr := errors.New("append unavailable") + store := &appendFailingConversationStore{ + MemoryStore: conversationstore.NewMemoryStore(), + err: storeErr, + } + turn := &conversationTurn{store: store, id: "conv_fragmented", input: "hello"} + stream := turn.persistingStream(context.Background(), &maxReadCloser{ + reader: strings.NewReader(createdEvent + completedEvent + suffix), + max: 1, + }) + t.Cleanup(func() { _ = stream.Close() }) + + got, err := io.ReadAll(stream) + if !errors.Is(err, storeErr) { + t.Fatalf("read error = %v, want append unavailable", err) + } + if string(got) != createdEvent { + t.Fatalf("stream body = %q, want only complete non-terminal event %q", got, createdEvent) + } + }) + } +} + +func TestConversationPersistingStreamPreservesFragmentedSuccessfulStream(t *testing.T) { + const streamData = "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_fragmented\"}}\r\n\r\n" + + "data: {\"type\":\"response.completed\",\r\n" + + "data: \"response\":{\"id\":\"resp_fragmented\",\"output\":[{\"id\":\"future_1\",\"type\":\"future_item\",\"future_counter\":9007199254740993,\"future_payload\":{\"preserved\":true}}]}}\r\n\r\n" + + "data: [DONE]\r\n\r\n" + + ctx := context.Background() + store := conversationstore.NewMemoryStore() + t.Cleanup(func() { _ = store.Close() }) + if err := store.Create(ctx, &conversationstore.StoredConversation{ + Conversation: &core.Conversation{ID: "conv_fragmented", Object: "conversation", Metadata: map[string]string{}}, + }); err != nil { + t.Fatalf("create conversation: %v", err) + } + turn := &conversationTurn{store: store, id: "conv_fragmented", input: "hello"} + stream := turn.persistingStream(ctx, &maxReadCloser{ + reader: strings.NewReader(streamData), + max: 1, + }) + t.Cleanup(func() { _ = stream.Close() }) + + got, err := io.ReadAll(stream) + if err != nil { + t.Fatalf("read stream: %v", err) + } + if string(got) != streamData { + t.Fatalf("stream body changed:\n got: %q\nwant: %q", got, streamData) + } + stored, err := store.Get(ctx, "conv_fragmented") + if err != nil { + t.Fatalf("get conversation: %v", err) + } + if len(stored.Items) != 2 { + t.Fatalf("stored items = %d, want input and output items", len(stored.Items)) + } + output, err := decodeRawJSONObject(stored.Items[1]) + if err != nil { + t.Fatalf("decode stored output: %v", err) + } + if got := string(output["future_counter"]); got != "9007199254740993" { + t.Fatalf("stored future_counter = %s, want exact large integer", got) + } + if got := string(output["future_payload"]); got != `{"preserved":true}` { + t.Fatalf("stored future_payload = %s, want unknown field preserved", got) + } +} diff --git a/internal/server/conversation_responses.go b/internal/server/conversation_responses.go index 1fd5cbec8..532c0e82f 100644 --- a/internal/server/conversation_responses.go +++ b/internal/server/conversation_responses.go @@ -1,10 +1,10 @@ package server import ( + "bytes" "context" "errors" "fmt" - "log/slog" "github.com/goccy/go-json" @@ -86,7 +86,9 @@ func mergeConversationInput(history []json.RawMessage, input any) ([]any, error) merged := make([]any, 0, len(history)) for _, raw := range history { var item map[string]any - if err := json.Unmarshal(raw, &item); err != nil { + decoder := json.NewDecoder(bytes.NewReader(raw)) + decoder.UseNumber() + if err := decoder.Decode(&item); err != nil { return nil, core.NewProviderError("conversation_store", 500, "stored conversation item is not valid JSON", err) } delete(item, "id") @@ -121,75 +123,86 @@ func mergeConversationInput(history []json.RawMessage, input any) ([]any, error) } } -// appendExchange records the completed turn: the request input normalized as -// stored input items, followed by the response output. Failures only log — -// the client already has its response, so the turn must not fail after the fact. -func (t *conversationTurn) appendExchange(ctx context.Context, responseID string, output []json.RawMessage) { +// appendExchange records the completed turn and returns the output items with +// their final persisted IDs. The caller can therefore expose exactly the IDs a +// later conversation item retrieve or delete request can address. +func (t *conversationTurn) appendExchange(ctx context.Context, responseID string, output []json.RawMessage) ([]json.RawMessage, error) { items := normalizedResponseInputItems(responseID, &core.ResponsesRequest{Input: t.input}) + inputCount := len(items) items = append(items, output...) if len(items) == 0 { - return + return nil, nil } - if err := t.store.AppendItems(ctx, t.id, items); err != nil { - slog.Warn("conversation append failed", "conversation_id", t.id, "error", err) + items = normalizeConversationItemIDs(nil, items) + err := t.store.AppendItems(ctx, t.id, items) + for range 3 { + if !errors.Is(err, conversationstore.ErrDuplicateItem) { + break + } + stored, getErr := t.store.Get(ctx, t.id) + if getErr != nil { + err = getErr + break + } + items = normalizeConversationItemIDs(stored.Items, items) + err = t.store.AppendItems(ctx, t.id, items) + } + if err != nil { + return nil, err } + return items[inputCount:], nil } -// appendResponse records a completed non-streaming response on the turn. -func (t *conversationTurn) appendResponse(ctx context.Context, resp *core.ResponsesResponse) { - if resp == nil { - return +func normalizeConversationItemIDs(existing, added []json.RawMessage) []json.RawMessage { + ids := make(map[string]struct{}, len(existing)+len(added)) + for _, raw := range existing { + if id := responseInputItemID(raw); id != "" { + ids[id] = struct{}{} + } } - output := make([]json.RawMessage, 0, len(resp.Output)) - for _, item := range resp.Output { - raw, err := json.Marshal(item) + result := make([]json.RawMessage, 0, len(added)) + for _, raw := range added { + item, err := decodeRawJSONObject(raw) if err != nil { - slog.Warn("conversation append failed: marshal output", "conversation_id", t.id, "error", err) - return + result = append(result, core.CloneRawJSON(raw)) + continue } - output = append(output, raw) - } - t.appendExchange(ctx, resp.ID, output) -} - -// streamObserver returns a streaming.Observer that captures the final -// response.completed event and appends the exchange when the stream closes. -// The context is detached from the request so the append survives the client -// connection ending right after the final event. -func (t *conversationTurn) streamObserver(ctx context.Context) *conversationStreamObserver { - return &conversationStreamObserver{turn: t, ctx: context.WithoutCancel(ctx)} -} - -type conversationStreamObserver struct { - turn *conversationTurn - ctx context.Context - response map[string]any -} - -func (o *conversationStreamObserver) OnJSONEvent(payload map[string]any) { - eventType, _ := payload["type"].(string) - if eventType != "response.completed" && eventType != "response.done" { - return - } - if response, ok := payload["response"].(map[string]any); ok { - o.response = response + id := rawJSONString(item, "id") + if _, duplicate := ids[id]; id == "" || duplicate { + id = generatedConversationItemID(rawJSONString(item, "type")) + _ = setRawJSONValue(item, "id", id) + } + ids[id] = struct{}{} + encoded, err := json.Marshal(item) + if err != nil { + result = append(result, core.CloneRawJSON(raw)) + continue + } + result = append(result, encoded) } + return result } -func (o *conversationStreamObserver) OnStreamClose() { - if o.response == nil { - return +// appendResponse records a completed non-streaming response and updates its +// output IDs to the IDs committed to the conversation store. +func (t *conversationTurn) appendResponse(ctx context.Context, resp *core.ResponsesResponse) error { + if resp == nil { + return nil } - responseID, _ := o.response["id"].(string) - outputItems, _ := o.response["output"].([]any) - output := make([]json.RawMessage, 0, len(outputItems)) - for _, item := range outputItems { + output := make([]json.RawMessage, 0, len(resp.Output)) + for _, item := range resp.Output { raw, err := json.Marshal(item) if err != nil { - slog.Warn("conversation append failed: marshal streamed output", "conversation_id", o.turn.id, "error", err) - return + return fmt.Errorf("marshal conversation output: %w", err) } output = append(output, raw) } - o.turn.appendExchange(o.ctx, responseID, output) + persistedOutput, err := t.appendExchange(ctx, resp.ID, output) + if err != nil { + return err + } + for index := range min(len(resp.Output), len(persistedOutput)) { + resp.Output[index].ID = responseInputItemID(persistedOutput[index]) + } + return nil } diff --git a/internal/server/conversation_responses_test.go b/internal/server/conversation_responses_test.go index 07772e0cf..552934cf1 100644 --- a/internal/server/conversation_responses_test.go +++ b/internal/server/conversation_responses_test.go @@ -1,15 +1,27 @@ package server import ( + "context" "encoding/json" + "errors" "net/http" "net/http/httptest" "strings" "testing" + "github.com/enterpilot/gomodel/internal/conversationstore" "github.com/enterpilot/gomodel/internal/core" ) +type appendFailingConversationStore struct { + *conversationstore.MemoryStore + err error +} + +func (s *appendFailingConversationStore) AppendItems(context.Context, string, []json.RawMessage) error { + return s.err +} + func conversationTestProvider(t *testing.T) *capturingProvider { t.Helper() return &capturingProvider{mockProvider: mockProvider{ @@ -182,3 +194,172 @@ func TestResponsesWithConversation_StreamingAppendsTurn(t *testing.T) { t.Fatalf("second item = %#v, want streamed assistant output", input[1]) } } + +func TestResponsesWithConversation_StreamingAppendFailureSuppressesCompletion(t *testing.T) { + provider := conversationTestProvider(t) + provider.streamData = strings.Join([]string{ + `data: {"type":"response.created","response":{"id":"resp_failed_persist"}}`, + "", + `data: {"type":"response.completed","response":{"id":"resp_failed_persist","output":[{"id":"msg_failed_persist","type":"message","role":"assistant","content":[{"type":"output_text","text":"not saved"}]}]}}`, + "", + "data: [DONE]", + "", + }, "\n") + store := &appendFailingConversationStore{ + MemoryStore: conversationstore.NewMemoryStore(), + err: errors.New("append unavailable"), + } + srv := New(provider, &Config{ConversationStore: store}) + convID := createTestConversation(t, srv, `{}`) + + rec := postResponses(t, srv, `{"model":"gpt-5-mini","input":"hello","conversation":"`+convID+`","stream":true}`) + if strings.Contains(rec.Body.String(), "response.completed") { + t.Fatalf("stream body = %s, must not report completion when persistence fails", rec.Body.String()) + } +} + +func TestResponsesWithConversation_PreservesReasoningFieldsOnReplay(t *testing.T) { + provider := conversationTestProvider(t) + var response core.ResponsesResponse + if err := json.Unmarshal([]byte(`{ + "id":"resp_reasoning","object":"response","model":"gpt-5-mini","status":"completed", + "output":[{"id":"rs_1","type":"reasoning","summary":[],"encrypted_content":"opaque"}] + }`), &response); err != nil { + t.Fatalf("decode reasoning response: %v", err) + } + provider.responsesResponse = &response + srv := New(provider, nil) + convID := createTestConversation(t, srv, `{}`) + + if rec := postResponses(t, srv, `{"model":"gpt-5-mini","input":"first","conversation":"`+convID+`"}`); rec.Code != http.StatusOK { + t.Fatalf("first response status = %d (%s)", rec.Code, rec.Body.String()) + } + if rec := postResponses(t, srv, `{"model":"gpt-5-mini","input":"second","conversation":"`+convID+`"}`); rec.Code != http.StatusOK { + t.Fatalf("second response status = %d (%s)", rec.Code, rec.Body.String()) + } + input, ok := provider.capturedResponsesReq.Input.([]any) + if !ok || len(input) != 3 { + t.Fatalf("replayed input = %#v, want first + reasoning + second", provider.capturedResponsesReq.Input) + } + reasoning, ok := input[1].(map[string]any) + if !ok || reasoning["type"] != "reasoning" || reasoning["encrypted_content"] != "opaque" { + t.Fatalf("reasoning item = %#v, want lossless replay", input[1]) + } + if _, ok := reasoning["summary"].([]any); !ok { + t.Fatalf("reasoning summary = %#v, want array", reasoning["summary"]) + } +} + +func TestMergeConversationInputPreservesLargeUnknownIntegers(t *testing.T) { + merged, err := mergeConversationInput([]json.RawMessage{ + json.RawMessage(`{"id":"future_1","type":"future_item","opaque_integer":9007199254740993}`), + }, nil) + if err != nil { + t.Fatalf("merge conversation input: %v", err) + } + encoded, err := json.Marshal(merged) + if err != nil { + t.Fatalf("marshal merged input: %v", err) + } + if !strings.Contains(string(encoded), `"opaque_integer":9007199254740993`) { + t.Fatalf("merged input = %s, want exact large integer", encoded) + } +} + +func TestResponsesWithConversation_RemapsReusedProviderItemIDs(t *testing.T) { + provider := conversationTestProvider(t) + srv := New(provider, nil) + convID := createTestConversation(t, srv, `{}`) + + returnedOutputIDs := make(map[string]struct{}, 3) + for _, input := range []string{"first", "second", "third"} { + rec := postResponses(t, srv, `{"model":"gpt-5-mini","input":"`+input+`","conversation":"`+convID+`"}`) + if rec.Code != http.StatusOK { + t.Fatalf("response for %q status = %d (%s)", input, rec.Code, rec.Body.String()) + } + var response core.ResponsesResponse + if err := json.Unmarshal(rec.Body.Bytes(), &response); err != nil || len(response.Output) != 1 { + t.Fatalf("decode response for %q: output=%v err=%v", input, response.Output, err) + } + returnedOutputIDs[response.Output[0].ID] = struct{}{} + } + + req := httptest.NewRequest(http.MethodGet, "/v1/conversations/"+convID+"/items?order=asc&limit=100", nil) + rec := httptest.NewRecorder() + srv.ServeHTTP(rec, req) + if rec.Code != http.StatusOK { + t.Fatalf("list status = %d (%s)", rec.Code, rec.Body.String()) + } + var list core.ConversationItemListResponse + if err := json.Unmarshal(rec.Body.Bytes(), &list); err != nil { + t.Fatalf("decode item list: %v", err) + } + if len(list.Data) != 6 { + t.Fatalf("items = %d, want three inputs and three outputs", len(list.Data)) + } + ids := make(map[string]struct{}, len(list.Data)) + for _, raw := range list.Data { + id := responseInputItemID(raw) + if id == "" { + t.Fatalf("item has no id: %s", raw) + } + if _, duplicate := ids[id]; duplicate { + t.Fatalf("duplicate persisted item id %q in %s", id, rec.Body.String()) + } + ids[id] = struct{}{} + } + for id := range returnedOutputIDs { + if _, persisted := ids[id]; !persisted { + t.Fatalf("response output id %q was not persisted: %s", id, rec.Body.String()) + } + itemReq := httptest.NewRequest(http.MethodGet, "/v1/conversations/"+convID+"/items/"+id, nil) + itemRec := httptest.NewRecorder() + srv.ServeHTTP(itemRec, itemReq) + if itemRec.Code != http.StatusOK { + t.Fatalf("retrieve returned output id %q status = %d (%s)", id, itemRec.Code, itemRec.Body.String()) + } + } +} + +func TestResponsesWithConversation_GeneratesMissingProviderOutputID(t *testing.T) { + provider := conversationTestProvider(t) + provider.responsesResponse.Output[0].ID = "" + srv := New(provider, nil) + convID := createTestConversation(t, srv, `{}`) + + rec := postResponses(t, srv, `{"model":"gpt-5-mini","input":"hello","conversation":"`+convID+`"}`) + if rec.Code != http.StatusOK { + t.Fatalf("response status = %d (%s)", rec.Code, rec.Body.String()) + } + var response core.ResponsesResponse + if err := json.Unmarshal(rec.Body.Bytes(), &response); err != nil || len(response.Output) != 1 { + t.Fatalf("decode response: output=%v err=%v", response.Output, err) + } + if response.Output[0].ID == "" { + t.Fatal("response output id is empty, want gateway-generated id") + } + itemReq := httptest.NewRequest(http.MethodGet, "/v1/conversations/"+convID+"/items/"+response.Output[0].ID, nil) + itemRec := httptest.NewRecorder() + srv.ServeHTTP(itemRec, itemReq) + if itemRec.Code != http.StatusOK { + t.Fatalf("retrieve generated output id status = %d (%s)", itemRec.Code, itemRec.Body.String()) + } +} + +func TestResponsesWithConversation_AppendFailureReturnsError(t *testing.T) { + provider := conversationTestProvider(t) + store := &appendFailingConversationStore{ + MemoryStore: conversationstore.NewMemoryStore(), + err: errors.New("append unavailable"), + } + srv := New(provider, &Config{ConversationStore: store}) + convID := createTestConversation(t, srv, `{}`) + + rec := postResponses(t, srv, `{"model":"gpt-5-mini","input":"hello","conversation":"`+convID+`"}`) + if rec.Code != http.StatusInternalServerError { + t.Fatalf("response status = %d (%s), want 500", rec.Code, rec.Body.String()) + } + if !strings.Contains(rec.Body.String(), "failed to append conversation turn") { + t.Fatalf("response body = %s, want append failure", rec.Body.String()) + } +} diff --git a/internal/server/http.go b/internal/server/http.go index f863d2e51..bc51bbd65 100644 --- a/internal/server/http.go +++ b/internal/server/http.go @@ -397,6 +397,10 @@ func New(provider core.RoutableProvider, cfg *Config) *Server { e.DELETE("/v1/responses/:id", handler.DeleteResponse) e.POST("/v1/responses", handler.Responses) e.POST("/v1/conversations", handler.CreateConversation) + e.POST("/v1/conversations/:id/items", handler.CreateConversationItems) + e.GET("/v1/conversations/:id/items", handler.ListConversationItems) + e.GET("/v1/conversations/:id/items/:item_id", handler.GetConversationItem) + e.DELETE("/v1/conversations/:id/items/:item_id", handler.DeleteConversationItem) e.GET("/v1/conversations/:id", handler.GetConversation) e.POST("/v1/conversations/:id", handler.UpdateConversation) e.DELETE("/v1/conversations/:id", handler.DeleteConversation) diff --git a/internal/server/native_conversation_items_service.go b/internal/server/native_conversation_items_service.go new file mode 100644 index 000000000..55a6b2e66 --- /dev/null +++ b/internal/server/native_conversation_items_service.go @@ -0,0 +1,151 @@ +package server + +import ( + "errors" + "fmt" + "net/http" + "strings" + + "github.com/labstack/echo/v5" + + "github.com/enterpilot/gomodel/internal/conversationstore" + "github.com/enterpilot/gomodel/internal/core" +) + +// CreateConversationItems handles POST /v1/conversations/{id}/items. +func (s *conversationService) CreateConversationItems(c *echo.Context) error { + ctx, _ := requestContextWithRequestID(c.Request()) + auditConversationEntry(c) + id, err := conversationIDFromRequest(c) + if err != nil { + return handleError(c, err) + } + body, err := requestBodyBytes(c) + if err != nil { + return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) + } + req, err := core.DecodeConversationItemCreateRequest(body) + if err != nil { + return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) + } + if len(req.Items) == 0 { + return handleError(c, core.NewInvalidRequestError("items is required", nil).WithParam("items")) + } + if len(req.Items) > core.MaxConversationInitialItems { + return handleError(c, core.NewInvalidRequestError( + fmt.Sprintf("items supports at most %d entries", core.MaxConversationInitialItems), nil, + ).WithParam("items")) + } + items, normalizeErr := normalizeConversationItems(req.Items) + if normalizeErr != nil { + return handleError(c, normalizeErr) + } + if err := s.conversationStore.AppendItems(ctx, id, items); err != nil { + switch { + case errors.Is(err, conversationstore.ErrNotFound): + return handleError(c, conversationNotFound(id)) + case errors.Is(err, conversationstore.ErrDuplicateItem): + return handleError(c, core.NewInvalidRequestError("duplicate conversation item id", err).WithParam("items")) + } + return handleError(c, core.NewProviderError("conversation_store", http.StatusInternalServerError, "failed to append conversation items", err)) + } + return c.JSON(http.StatusOK, conversationItemList(items, false, conversationItemIncludes(c))) +} + +// ListConversationItems handles GET /v1/conversations/{id}/items. +func (s *conversationService) ListConversationItems(c *echo.Context) error { + ctx, _ := requestContextWithRequestID(c.Request()) + auditConversationEntry(c) + id, err := conversationIDFromRequest(c) + if err != nil { + return handleError(c, err) + } + params, err := conversationItemListParamsFromRequest(c) + if err != nil { + return handleError(c, err) + } + stored, err := s.loadStoredConversation(ctx, id) + if err != nil { + return handleError(c, err) + } + page, pageErr := paginateConversationItems(stored.Items, params) + if pageErr != nil { + return handleError(c, pageErr) + } + return c.JSON(http.StatusOK, page) +} + +// GetConversationItem handles GET /v1/conversations/{id}/items/{item_id}. +func (s *conversationService) GetConversationItem(c *echo.Context) error { + ctx, _ := requestContextWithRequestID(c.Request()) + auditConversationEntry(c) + id, err := conversationIDFromRequest(c) + if err != nil { + return handleError(c, err) + } + itemID, err := conversationItemIDFromRequest(c) + if err != nil { + return handleError(c, err) + } + stored, err := s.loadStoredConversation(ctx, id) + if err != nil { + return handleError(c, err) + } + for _, item := range stored.Items { + if responseInputItemID(item) == itemID { + return c.JSONBlob(http.StatusOK, conversationItemForInclude(item, conversationItemIncludes(c))) + } + } + return handleError(c, conversationItemNotFound(itemID)) +} + +// DeleteConversationItem handles DELETE /v1/conversations/{id}/items/{item_id}. +func (s *conversationService) DeleteConversationItem(c *echo.Context) error { + ctx, _ := requestContextWithRequestID(c.Request()) + auditConversationEntry(c) + id, err := conversationIDFromRequest(c) + if err != nil { + return handleError(c, err) + } + itemID, err := conversationItemIDFromRequest(c) + if err != nil { + return handleError(c, err) + } + stored, err := s.conversationStore.DeleteItem(ctx, id, itemID) + if err != nil { + switch { + case errors.Is(err, conversationstore.ErrNotFound): + return handleError(c, conversationNotFound(id)) + case errors.Is(err, conversationstore.ErrItemNotFound): + return handleError(c, conversationItemNotFound(itemID)) + default: + return handleError(c, core.NewProviderError("conversation_store", http.StatusInternalServerError, "failed to delete conversation item", err)) + } + } + return c.JSON(http.StatusOK, stored.Conversation) +} + +func conversationItemListParamsFromRequest(c *echo.Context) (core.ConversationItemListParams, error) { + common, err := cursorListParamsFromRequest(c) + if err != nil { + return core.ConversationItemListParams{}, err + } + return core.ConversationItemListParams{ + After: common.after, + Include: common.include, + Limit: common.limit, + Order: common.order, + }, nil +} + +func conversationItemIDFromRequest(c *echo.Context) (string, error) { + id := strings.TrimSpace(c.Param("item_id")) + if id == "" { + return "", core.NewInvalidRequestError("conversation item id is required", nil) + } + return id, nil +} + +func conversationItemNotFound(id string) *core.GatewayError { + return core.NewNotFoundError("conversation item not found: " + id) +} diff --git a/internal/server/native_conversation_service.go b/internal/server/native_conversation_service.go index 4467308de..8d604d024 100644 --- a/internal/server/native_conversation_service.go +++ b/internal/server/native_conversation_service.go @@ -45,6 +45,10 @@ func (s *conversationService) CreateConversation(c *echo.Context) error { fmt.Sprintf("items supports at most %d entries", core.MaxConversationInitialItems), nil, ).WithParam("items")) } + items, normalizeErr := normalizeConversationItems(req.Items) + if normalizeErr != nil { + return handleError(c, normalizeErr) + } if verr := core.ValidateConversationMetadata(req.Metadata); verr != nil { return handleError(c, verr) } @@ -58,7 +62,7 @@ func (s *conversationService) CreateConversation(c *echo.Context) error { } stored := &conversationstore.StoredConversation{ Conversation: conversation, - Items: cloneRawConversationItems(req.Items), + Items: cloneRawConversationItems(items), UserPath: core.UserPathFromContext(ctx), RequestID: requestID, StoredAt: now, @@ -85,8 +89,8 @@ func (s *conversationService) GetConversation(c *echo.Context) error { return c.JSON(http.StatusOK, stored.Conversation) } -// UpdateConversation handles POST /v1/conversations/{id}. The metadata in the -// request replaces the conversation's metadata in full, matching OpenAI. +// UpdateConversation handles POST /v1/conversations/{id}. Supplied metadata is +// merged into the existing metadata, matching OpenAI. func (s *conversationService) UpdateConversation(c *echo.Context) error { ctx, _ := requestContextWithRequestID(c.Request()) auditConversationEntry(c) @@ -110,14 +114,15 @@ func (s *conversationService) UpdateConversation(c *echo.Context) error { return handleError(c, verr) } - stored, err := s.loadStoredConversation(ctx, id) + stored, err := s.conversationStore.MergeMetadata(ctx, id, *req.Metadata) if err != nil { - return handleError(c, err) - } - stored.Conversation.Metadata = normalizedConversationMetadata(*req.Metadata) - if err := s.conversationStore.Update(ctx, stored); err != nil { - if errors.Is(err, conversationstore.ErrNotFound) { + switch { + case errors.Is(err, conversationstore.ErrNotFound): return handleError(c, conversationNotFound(id)) + case errors.Is(err, conversationstore.ErrMetadataLimitExceeded): + return handleError(c, core.NewInvalidRequestError( + fmt.Sprintf("metadata supports at most %d key-value pairs after update", core.MaxConversationMetadataPairs), err, + ).WithParam("metadata").WithCode("metadata_max_properties_exceeded")) } return handleError(c, core.NewProviderError("conversation_store", http.StatusInternalServerError, "failed to update conversation", err)) } diff --git a/internal/server/native_response_service.go b/internal/server/native_response_service.go index b93f93232..f2c20a043 100644 --- a/internal/server/native_response_service.go +++ b/internal/server/native_response_service.go @@ -465,12 +465,37 @@ func responseRetrieveParamsFromRequest(c *echo.Context) (core.ResponseRetrievePa } func responseInputItemsParamsFromRequest(c *echo.Context) (core.ResponseInputItemsParams, error) { + common, err := cursorListParamsFromRequest(c) + if err != nil { + return core.ResponseInputItemsParams{}, err + } + return core.ResponseInputItemsParams{ + After: common.after, + Include: common.include, + Limit: common.limit, + Order: common.order, + }, nil +} + +type cursorListParams struct { + after string + include []string + limit int + order string +} + +const ( + defaultCursorListLimit = 20 + maxCursorListLimit = 100 +) + +func cursorListParamsFromRequest(c *echo.Context) (cursorListParams, error) { query := c.Request().URL.Query() - params := core.ResponseInputItemsParams{ - After: strings.TrimSpace(query.Get("after")), - Include: appendQueryArray(query["include"], query["include[]"]), - Limit: 20, - Order: "desc", + params := cursorListParams{ + after: strings.TrimSpace(query.Get("after")), + include: appendQueryArray(query["include"], query["include[]"]), + limit: defaultCursorListLimit, + order: "desc", } if raw := strings.TrimSpace(query.Get("limit")); raw != "" { limit, err := strconv.Atoi(raw) @@ -479,18 +504,18 @@ func responseInputItemsParamsFromRequest(c *echo.Context) (core.ResponseInputIte } switch { case limit <= 0: - params.Limit = 20 - case limit > 100: - params.Limit = 100 + params.limit = defaultCursorListLimit + case limit > maxCursorListLimit: + params.limit = maxCursorListLimit default: - params.Limit = limit + params.limit = limit } } if raw := strings.TrimSpace(query.Get("order")); raw != "" { if raw != "asc" && raw != "desc" { return params, core.NewInvalidRequestError("order must be asc or desc", nil).WithParam("order") } - params.Order = raw + params.order = raw } return params, nil } @@ -532,10 +557,10 @@ func paginateStoredResponseInputItems(items []json.RawMessage, params core.Respo limit := params.Limit if limit <= 0 { - limit = 20 + limit = defaultCursorListLimit } - if limit > 100 { - limit = 100 + if limit > maxCursorListLimit { + limit = maxCursorListLimit } remaining := max(count-start, 0) diff --git a/internal/server/raw_json_object.go b/internal/server/raw_json_object.go new file mode 100644 index 000000000..564ecd0e3 --- /dev/null +++ b/internal/server/raw_json_object.go @@ -0,0 +1,47 @@ +package server + +import ( + "bytes" + "fmt" + "strings" + + "github.com/goccy/go-json" +) + +// rawJSONObject keeps values encoded until a known field needs to be read or +// changed. This is important for forward-compatible API objects: decoding an +// unknown JSON number through any would round large integers through float64. +type rawJSONObject map[string]json.RawMessage + +func decodeRawJSONObject(raw json.RawMessage) (rawJSONObject, error) { + var object rawJSONObject + if err := json.Unmarshal(raw, &object); err != nil { + return nil, err + } + if object == nil { + return nil, fmt.Errorf("expected JSON object") + } + return object, nil +} + +func rawJSONString(object rawJSONObject, key string) string { + var value string + if err := json.Unmarshal(object[key], &value); err != nil { + return "" + } + return strings.TrimSpace(value) +} + +func setRawJSONValue(object rawJSONObject, key string, value any) error { + raw, err := json.Marshal(value) + if err != nil { + return err + } + object[key] = raw + return nil +} + +func rawJSONValuePresent(object rawJSONObject, key string) bool { + raw, exists := object[key] + return exists && len(bytes.TrimSpace(raw)) > 0 && !bytes.Equal(bytes.TrimSpace(raw), []byte("null")) +} diff --git a/internal/server/response_input_items.go b/internal/server/response_input_items.go index dd55ca7f5..651c8fa85 100644 --- a/internal/server/response_input_items.go +++ b/internal/server/response_input_items.go @@ -1,8 +1,8 @@ package server import ( + "bytes" "fmt" - "maps" "strings" "github.com/goccy/go-json" @@ -69,8 +69,8 @@ func normalizedResponseInputAny(responseID string, index int, item any) json.Raw } func normalizedResponseInputRaw(responseID string, index int, raw json.RawMessage) json.RawMessage { - var item map[string]any - if err := json.Unmarshal(raw, &item); err != nil { + item, err := decodeRawJSONObject(raw) + if err != nil { var decoded string text := strings.TrimSpace(string(raw)) if stringErr := json.Unmarshal(raw, &decoded); stringErr == nil { @@ -88,112 +88,79 @@ func normalizedResponseInputRaw(responseID string, index int, raw json.RawMessag }, }) } - if item == nil { - return nil - } - - itemType := strings.TrimSpace(stringFromMap(item, "type")) + itemType := rawJSONString(item, "type") if itemType == "" { itemType = "message" - item["type"] = itemType + _ = setRawJSONValue(item, "type", itemType) } switch itemType { case "message": - if strings.TrimSpace(stringFromMap(item, "role")) == "" { - item["role"] = "user" + if rawJSONString(item, "role") == "" { + _ = setRawJSONValue(item, "role", "user") + } + normalizedContent, contentErr := normalizeResponseInputContentRaw(item["content"]) + if contentErr == nil { + item["content"] = normalizedContent } - item["content"] = normalizeResponseInputContent(item["content"]) case "function_call", "function_call_output": // The decoded request has already normalized call_id/id aliases. default: // Unknown item types are preserved with an ID attached for pagination. } - if strings.TrimSpace(stringFromMap(item, "id")) == "" { - item["id"] = generatedResponseInputItemID(responseID, index, itemType, stringFromMap(item, "call_id")) + if rawJSONString(item, "id") == "" { + _ = setRawJSONValue(item, "id", generatedResponseInputItemID(responseID, index, itemType, rawJSONString(item, "call_id"))) } return mustRawJSON(item) } -func normalizeResponseInputContent(content any) any { - switch value := content.(type) { - case nil: - return []map[string]any{} - case string: - return []map[string]any{{"type": "input_text", "text": value}} - case []core.ContentPart: - items := make([]map[string]any, 0, len(value)) - for _, part := range value { - items = append(items, normalizeContentPart(part)) - } - return items - case []any: - items := make([]any, 0, len(value)) - for _, part := range value { - items = append(items, normalizeResponseInputContentPartAny(part)) - } - return items - default: - return content +func normalizeResponseInputContentRaw(raw json.RawMessage) (json.RawMessage, error) { + trimmed := bytes.TrimSpace(raw) + if len(trimmed) == 0 || bytes.Equal(trimmed, []byte("null")) { + return json.RawMessage("[]"), nil } -} - -func normalizeContentPart(part core.ContentPart) map[string]any { - switch strings.TrimSpace(part.Type) { - case "text", "input_text": - return map[string]any{"type": "input_text", "text": part.Text} - case "image_url", "input_image": - item := map[string]any{"type": "input_image"} - if part.ImageURL != nil { - item["image_url"] = part.ImageURL.URL - if detail := strings.TrimSpace(part.ImageURL.Detail); detail != "" { - item["detail"] = detail - } + if trimmed[0] == '"' { + var text string + if err := json.Unmarshal(trimmed, &text); err != nil { + return nil, err } - return item - case "input_audio": - item := map[string]any{"type": "input_audio"} - if part.InputAudio != nil { - item["input_audio"] = map[string]any{ - "data": part.InputAudio.Data, - "format": part.InputAudio.Format, - } - } - return item - default: - return map[string]any{"type": part.Type, "text": part.Text} + return json.Marshal([]map[string]any{{"type": "input_text", "text": text}}) } -} - -func normalizeResponseInputContentPartAny(part any) any { - item, ok := part.(map[string]any) - if !ok { - return part + if trimmed[0] != '[' { + return core.CloneRawJSON(trimmed), nil } - partType := strings.TrimSpace(stringFromMap(item, "type")) - switch partType { - case "text": - normalized := cloneAnyMap(item) - normalized["type"] = "input_text" - return normalized - case "image_url": - normalized := cloneAnyMap(item) - normalized["type"] = "input_image" - if image, ok := normalized["image_url"].(map[string]any); ok { - if url, ok := image["url"]; ok { - normalized["image_url"] = url - } - if detail, ok := image["detail"]; ok { - normalized["detail"] = detail + var parts []json.RawMessage + if err := json.Unmarshal(trimmed, &parts); err != nil { + return nil, err + } + for index, rawPart := range parts { + part, err := decodeRawJSONObject(rawPart) + if err != nil { + continue + } + switch rawJSONString(part, "type") { + case "text": + _ = setRawJSONValue(part, "type", "input_text") + case "image_url", "input_image": + _ = setRawJSONValue(part, "type", "input_image") + if image, imageErr := decodeRawJSONObject(part["image_url"]); imageErr == nil { + if url, exists := image["url"]; exists { + part["image_url"] = core.CloneRawJSON(url) + } + if detail, exists := image["detail"]; exists { + part["detail"] = core.CloneRawJSON(detail) + } } } - return normalized - default: - return item + parts[index], err = json.Marshal(part) + if err != nil { + return nil, err + } } + return json.Marshal(parts) } func generatedResponseInputItemID(responseID string, index int, itemType, callID string) string { @@ -219,17 +186,3 @@ func mustRawJSON(value any) json.RawMessage { } return raw } - -func stringFromMap(item map[string]any, key string) string { - value, _ := item[key].(string) - return strings.TrimSpace(value) -} - -func cloneAnyMap(src map[string]any) map[string]any { - if len(src) == 0 { - return nil - } - dst := make(map[string]any, len(src)) - maps.Copy(dst, src) - return dst -} diff --git a/internal/server/response_input_items_test.go b/internal/server/response_input_items_test.go index 986577d1a..b05abb69c 100644 --- a/internal/server/response_input_items_test.go +++ b/internal/server/response_input_items_test.go @@ -2,6 +2,7 @@ package server import ( "encoding/json" + "strings" "testing" "github.com/enterpilot/gomodel/internal/core" @@ -17,6 +18,23 @@ func TestNormalizedResponseInputItemsSkipsNilDefaultInput(t *testing.T) { } } +func TestNormalizedResponseInputRawPreservesLargeUnknownIntegers(t *testing.T) { + item := normalizedResponseInputRaw("resp_1", 0, json.RawMessage( + `{"type":"future_item","opaque_integer":9007199254740993,"nested":{"value":9007199254740995}}`, + )) + for _, want := range []string{ + `"opaque_integer":9007199254740993`, + `"value":9007199254740995`, + } { + if !strings.Contains(string(item), want) { + t.Fatalf("normalized item = %s, want %s", item, want) + } + } + if responseInputItemID(item) == "" { + t.Fatalf("normalized item = %s, want generated id", item) + } +} + func TestNormalizedResponseInputRawSkipsNullObject(t *testing.T) { item := normalizedResponseInputRaw("resp_1", 0, json.RawMessage("null")) if len(item) != 0 { diff --git a/internal/server/translated_inference_service.go b/internal/server/translated_inference_service.go index dbf968c15..cf1226eb4 100644 --- a/internal/server/translated_inference_service.go +++ b/internal/server/translated_inference_service.go @@ -282,7 +282,7 @@ func (s *translatedInferenceService) dispatchResponses(c *echo.Context, req *cor } stream := result.Stream if turn := conversationTurnFromContext(ctx); turn != nil { - stream = streaming.NewObservedSSEStream(stream, turn.streamObserver(ctx)) + stream = turn.persistingStream(ctx, stream) } return s.handleStreamingReadCloser( c, @@ -312,13 +312,17 @@ func (s *translatedInferenceService) dispatchResponses(c *echo.Context, req *cor result.Meta.ProviderName, ) - if err := s.storeResponseSnapshot(ctx, workflow, req, result.Response, result.Meta.ProviderType, result.Meta.ProviderName, requestID); err != nil { - s.recordResponseSnapshotStoreFailure(workflow, result.Response, result.Meta.ProviderType, result.Meta.ProviderName, requestID, err) - } if turn := conversationTurnFromContext(ctx); turn != nil { // Detach cancellation so a client disconnect after provider success // cannot lose the completed turn, mirroring the streaming observer. - turn.appendResponse(context.WithoutCancel(ctx), result.Response) + if err := turn.appendResponse(context.WithoutCancel(ctx), result.Response); err != nil { + return handleError(c, core.NewProviderError( + "conversation_store", http.StatusInternalServerError, "failed to append conversation turn", err, + )) + } + } + if err := s.storeResponseSnapshot(ctx, workflow, req, result.Response, result.Meta.ProviderType, result.Meta.ProviderName, requestID); err != nil { + s.recordResponseSnapshotStoreFailure(workflow, result.Response, result.Meta.ProviderType, result.Meta.ProviderName, requestID, err) } return c.JSON(http.StatusOK, result.Response) diff --git a/tests/contract/testdata/golden/xai/responses.golden.json b/tests/contract/testdata/golden/xai/responses.golden.json index 5333297c6..5a38dbb8e 100644 --- a/tests/contract/testdata/golden/xai/responses.golden.json +++ b/tests/contract/testdata/golden/xai/responses.golden.json @@ -13,6 +13,7 @@ ], "id": "rs_afd66348-34f4-4057-3c98-abfc8078703c", "status": "completed", + "summary": [], "type": "reasoning" }, { diff --git a/tests/e2e/conversations_test.go b/tests/e2e/conversations_test.go new file mode 100644 index 000000000..826e1d45b --- /dev/null +++ b/tests/e2e/conversations_test.go @@ -0,0 +1,112 @@ +//go:build e2e + +package e2e + +import ( + "bytes" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/require" +) + +type conversationAPIObject struct { + ID string `json:"id"` + Object string `json:"object"` + Metadata map[string]string `json:"metadata"` + Deleted bool `json:"deleted"` + Data []json.RawMessage `json:"data"` + FirstID string `json:"first_id"` + LastID string `json:"last_id"` + HasMore bool `json:"has_more"` +} + +// TestConversationsOfficialRoutes_E2E exercises every Conversations operation +// exposed by the official OpenAI SDK through the real router and auth stack. +func TestConversationsOfficialRoutes_E2E(t *testing.T) { + const key = "sk-e2e-conversations" + httpServer := httptest.NewServer(setupAuthServer(t, key)) + t.Cleanup(httpServer.Close) + + call := func(method, path, body string, target any) { + t.Helper() + var reader io.Reader + if body != "" { + reader = bytes.NewBufferString(body) + } + req, err := http.NewRequest(method, httpServer.URL+path, reader) + require.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+key) + if body != "" { + req.Header.Set("Content-Type", "application/json") + } + resp, err := http.DefaultClient.Do(req) + require.NoError(t, err) + defer resp.Body.Close() + payload, err := io.ReadAll(resp.Body) + require.NoError(t, err) + require.Equalf(t, http.StatusOK, resp.StatusCode, "%s %s: %s", method, path, payload) + if target != nil { + require.NoError(t, json.Unmarshal(payload, target)) + } + } + + var created conversationAPIObject + call(http.MethodPost, "/v1/conversations", `{ + "metadata":{"suite":"sdk","retained":"yes"}, + "items":[ + {"role":"developer","content":"Preserve exact JSON fields."}, + {"type":"message","role":"user","content":[{"type":"input_text","text":"first"}],"phase":"commentary"} + ] + }`, &created) + require.NotEmpty(t, created.ID) + require.Equal(t, "conversation", created.Object) + + var retrieved conversationAPIObject + call(http.MethodGet, "/v1/conversations/"+created.ID, "", &retrieved) + require.Equal(t, created.ID, retrieved.ID) + + var updated conversationAPIObject + call(http.MethodPost, "/v1/conversations/"+created.ID, `{"metadata":{"suite":"updated"}}`, &updated) + require.Equal(t, "updated", updated.Metadata["suite"]) + require.Equal(t, "yes", updated.Metadata["retained"]) + + var firstPage conversationAPIObject + call(http.MethodGet, "/v1/conversations/"+created.ID+"/items?order=asc&limit=1&include=reasoning.encrypted_content", "", &firstPage) + require.Len(t, firstPage.Data, 1) + require.True(t, firstPage.HasMore) + require.NotEmpty(t, firstPage.LastID) + + var secondPage conversationAPIObject + call(http.MethodGet, "/v1/conversations/"+created.ID+"/items?order=asc&limit=10&after="+firstPage.LastID, "", &secondPage) + require.Len(t, secondPage.Data, 1) + require.False(t, secondPage.HasMore) + + var added conversationAPIObject + call(http.MethodPost, "/v1/conversations/"+created.ID+"/items?include=reasoning.encrypted_content", `{ + "items":[ + {"type":"function_call","call_id":"call_e2e","name":"lookup","arguments":{"nested":true}}, + {"type":"function_call_output","call_id":"call_e2e","output":[{"type":"input_text","text":"done"}]}, + {"type":"reasoning","summary":[],"encrypted_content":"opaque-e2e"} + ] + }`, &added) + require.Len(t, added.Data, 3) + require.NotEmpty(t, added.LastID) + + var item map[string]any + call(http.MethodGet, "/v1/conversations/"+created.ID+"/items/"+added.LastID+"?include=reasoning.encrypted_content", "", &item) + require.Equal(t, "reasoning", item["type"]) + require.Equal(t, "opaque-e2e", item["encrypted_content"]) + + var afterItemDelete conversationAPIObject + call(http.MethodDelete, "/v1/conversations/"+created.ID+"/items/"+added.LastID, "", &afterItemDelete) + require.Equal(t, created.ID, afterItemDelete.ID) + + var deleted conversationAPIObject + call(http.MethodDelete, "/v1/conversations/"+created.ID, "", &deleted) + require.Equal(t, "conversation.deleted", deleted.Object) + require.True(t, deleted.Deleted) +} diff --git a/tests/integration/conversations_test.go b/tests/integration/conversations_test.go new file mode 100644 index 000000000..b59771a27 --- /dev/null +++ b/tests/integration/conversations_test.go @@ -0,0 +1,331 @@ +//go:build integration + +package integration + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "slices" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" + "go.mongodb.org/mongo-driver/v2/bson" +) + +const conversationIntegrationKey = "sk-conversation-integration" + +type conversationHTTPResult struct { + status int + body []byte + err error +} + +func TestConversationsComplexStateTransitions_PostgreSQL(t *testing.T) { + testConversationsComplexStateTransitions(t, "postgresql") +} + +func TestConversationsComplexStateTransitions_MongoDB(t *testing.T) { + testConversationsComplexStateTransitions(t, "mongodb") +} + +func testConversationsComplexStateTransitions(t *testing.T, dbType string) { + t.Helper() + fixture := SetupTestServer(t, TestServerConfig{ + DBType: dbType, + MasterKey: conversationIntegrationKey, + }) + t.Cleanup(func() { fixture.Shutdown(t) }) + + created := conversationJSONRequest(t, fixture, http.MethodPost, "/v1/conversations", map[string]any{ + "metadata": map[string]string{"suite": "complex-" + dbType, "unicode": "Zażółć 🧪"}, + "items": []any{ + map[string]any{"id": "dev_seed", "type": "message", "role": "developer", "content": "Preserve exact state."}, + map[string]any{"id": "image_seed", "type": "message", "role": "user", "content": []any{ + map[string]any{"type": "input_text", "text": "turn zero"}, + map[string]any{"type": "input_image", "image_url": "data:image/png;base64,cWE=", "detail": "high", "x_vendor": map[string]any{"deep": []any{1, map[string]any{"two": 2}}}}, + }}, + map[string]any{"id": "reasoning_seed", "type": "reasoning", "summary": []any{map[string]any{"type": "summary_text", "text": "kept summary"}}, "encrypted_content": "opaque-persisted"}, + map[string]any{"id": "future_seed", "type": "future_tool_v9", "payload": map[string]any{"nil": nil, "float": 1.25, "lambda": "λ"}, "x_extension": []any{"a", map[string]any{"b": true}}}, + }, + }, http.StatusOK) + conversationID := created["id"].(string) + + // Mix metadata patches and item appends at the same instant. Both operations + // touch the same logical conversation but different persisted fields. + const metadataWriters = 8 + const itemWriters = 16 + start := make(chan struct{}) + results := make(chan conversationHTTPResult, metadataWriters+itemWriters) + var wg sync.WaitGroup + for index := range metadataWriters { + wg.Go(func() { + <-start + results <- rawConversationRequest(fixture, http.MethodPost, "/v1/conversations/"+conversationID, + map[string]any{"metadata": map[string]string{fmt.Sprintf("parallel_%d", index): fmt.Sprintf("value-%d", index)}}) + }) + } + for index := range itemWriters { + wg.Go(func() { + <-start + results <- rawConversationRequest(fixture, http.MethodPost, "/v1/conversations/"+conversationID+"/items", map[string]any{ + "items": []any{map[string]any{"id": fmt.Sprintf("parallel_item_%02d", index), "type": "message", "role": "user", "content": fmt.Sprintf("parallel %d", index), "sequence": index}}, + }) + }) + } + close(start) + wg.Wait() + close(results) + for result := range results { + require.NoError(t, result.err) + require.Equal(t, http.StatusOK, result.status, string(result.body)) + } + + updated := conversationJSONRequest(t, fixture, http.MethodGet, "/v1/conversations/"+conversationID, nil, http.StatusOK) + metadata := updated["metadata"].(map[string]any) + require.Len(t, metadata, metadataWriters+2) + for index := range metadataWriters { + require.Equal(t, fmt.Sprintf("value-%d", index), metadata[fmt.Sprintf("parallel_%d", index)]) + } + + // Race identical explicit IDs. The store-level compare-and-append must leave + // exactly one addressable item even when every handler read the same snapshot. + const collisionWriters = 12 + start = make(chan struct{}) + collisions := make(chan conversationHTTPResult, collisionWriters) + for range collisionWriters { + wg.Go(func() { + <-start + collisions <- rawConversationRequest(fixture, http.MethodPost, "/v1/conversations/"+conversationID+"/items", map[string]any{ + "items": []any{map[string]any{"id": "shared_collision", "type": "message", "role": "user", "content": "same id"}}, + }) + }) + } + close(start) + wg.Wait() + close(collisions) + statusCounts := map[int]int{} + for result := range collisions { + require.NoError(t, result.err) + statusCounts[result.status]++ + } + require.Equal(t, 1, statusCounts[http.StatusOK], statusCounts) + require.Equal(t, collisionWriters-1, statusCounts[http.StatusBadRequest], statusCounts) + + // Delete adjacent items concurrently. PostgreSQL must calculate each JSONB + // array index after acquiring the latest row version, or the second writer + // can remove the item that shifted into a stale index. + conversationJSONRequest(t, fixture, http.MethodPost, "/v1/conversations/"+conversationID+"/items", map[string]any{ + "items": []any{ + map[string]any{"id": "delete_race_a", "type": "future_race", "marker": "a"}, + map[string]any{"id": "delete_race_b", "type": "future_race", "marker": "b"}, + map[string]any{"id": "delete_race_guard", "type": "future_race", "marker": "guard"}, + }, + }, http.StatusOK) + start = make(chan struct{}) + concurrentDeletes := make(chan conversationHTTPResult, 2) + for _, itemID := range []string{"delete_race_a", "delete_race_b"} { + wg.Go(func() { + <-start + concurrentDeletes <- rawConversationRequest(fixture, http.MethodDelete, + "/v1/conversations/"+conversationID+"/items/"+itemID, nil) + }) + } + close(start) + wg.Wait() + close(concurrentDeletes) + for result := range concurrentDeletes { + require.NoError(t, result.err) + require.Equal(t, http.StatusOK, result.status, string(result.body)) + } + conversationJSONRequest(t, fixture, http.MethodGet, + "/v1/conversations/"+conversationID+"/items/delete_race_a", nil, http.StatusNotFound) + conversationJSONRequest(t, fixture, http.MethodGet, + "/v1/conversations/"+conversationID+"/items/delete_race_b", nil, http.StatusNotFound) + conversationJSONRequest(t, fixture, http.MethodGet, + "/v1/conversations/"+conversationID+"/items/delete_race_guard", nil, http.StatusOK) + + ascending := collectConversationItemIDs(t, fixture, conversationID, "asc", 7) + require.Len(t, ascending, 4+itemWriters+1+1) + require.Len(t, uniqueStrings(ascending), len(ascending)) + descending := collectConversationItemIDs(t, fixture, conversationID, "desc", 100) + reversed := slices.Clone(ascending) + slices.Reverse(reversed) + require.Equal(t, reversed, descending) + + // Optional-field redaction must be a view, not a destructive rewrite. + withoutInclude := conversationJSONRequest(t, fixture, http.MethodGet, + "/v1/conversations/"+conversationID+"/items/reasoning_seed", nil, http.StatusOK) + require.NotContains(t, withoutInclude, "encrypted_content") + withInclude := conversationJSONRequest(t, fixture, http.MethodGet, + "/v1/conversations/"+conversationID+"/items/reasoning_seed?include=reasoning.encrypted_content", nil, http.StatusOK) + require.Equal(t, "opaque-persisted", withInclude["encrypted_content"]) + + // Delete one old item while another writer appends. MongoDB uses a CAS and + // PostgreSQL uses JSONB row updates; neither operation may undo the other. + start = make(chan struct{}) + deleteAppend := make(chan conversationHTTPResult, 2) + wg.Add(2) + go func() { + defer wg.Done() + <-start + deleteAppend <- rawConversationRequest(fixture, http.MethodDelete, + "/v1/conversations/"+conversationID+"/items/future_seed", nil) + }() + go func() { + defer wg.Done() + <-start + deleteAppend <- rawConversationRequest(fixture, http.MethodPost, + "/v1/conversations/"+conversationID+"/items", map[string]any{ + "items": []any{map[string]any{"id": "post_delete_append", "type": "future_after_delete", "payload": map[string]any{"kept": true}}}, + }) + }() + close(start) + wg.Wait() + close(deleteAppend) + for result := range deleteAppend { + require.NoError(t, result.err) + require.Equal(t, http.StatusOK, result.status, string(result.body)) + } + conversationJSONRequest(t, fixture, http.MethodGet, + "/v1/conversations/"+conversationID+"/items/future_seed", nil, http.StatusNotFound) + conversationJSONRequest(t, fixture, http.MethodGet, + "/v1/conversations/"+conversationID+"/items/post_delete_append", nil, http.StatusOK) + + // Three Responses turns exercise replay through a real provider adapter. The + // integration mock intentionally reuses response and output IDs each turn; + // persisted conversation IDs must still be unique and all turns retained. + for turn := range 3 { + conversationJSONRequest(t, fixture, http.MethodPost, "/v1/responses", map[string]any{ + "model": "gpt-4", + "conversation": map[string]any{"id": conversationID}, + "input": []any{map[string]any{ + "type": "message", "role": "user", + "content": []any{map[string]any{"type": "input_text", "text": fmt.Sprintf("complex turn %d", turn)}}, + "phase": "commentary", + }}, + "reasoning": map[string]any{"effort": "medium", "summary": "detailed"}, + "text": map[string]any{"verbosity": "low"}, + "parallel_tool_calls": false, + "store": false, + }, http.StatusOK) + } + afterTurns := collectConversationItemIDs(t, fixture, conversationID, "asc", 100) + require.Len(t, afterTurns, len(ascending)+6) + require.Len(t, uniqueStrings(afterTurns), len(afterTurns)) + + assertConversationDatabaseState(t, fixture, conversationID, metadataWriters+2, len(afterTurns)) +} + +func rawConversationRequest(fixture *TestServerFixture, method, path string, payload any) conversationHTTPResult { + var body io.Reader + if payload != nil { + encoded, err := json.Marshal(payload) + if err != nil { + return conversationHTTPResult{err: err} + } + body = bytes.NewReader(encoded) + } + req, err := http.NewRequest(method, fixture.ServerURL+path, body) + if err != nil { + return conversationHTTPResult{err: err} + } + req.Header.Set("Authorization", "Bearer "+conversationIntegrationKey) + if payload != nil { + req.Header.Set("Content-Type", "application/json") + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + return conversationHTTPResult{err: err} + } + defer resp.Body.Close() + responseBody, err := io.ReadAll(resp.Body) + return conversationHTTPResult{status: resp.StatusCode, body: responseBody, err: err} +} + +func conversationJSONRequest(t *testing.T, fixture *TestServerFixture, method, path string, payload any, wantStatus int) map[string]any { + t.Helper() + result := rawConversationRequest(fixture, method, path, payload) + require.NoError(t, result.err) + require.Equal(t, wantStatus, result.status, string(result.body)) + var decoded map[string]any + require.NoError(t, json.Unmarshal(result.body, &decoded), string(result.body)) + return decoded +} + +func collectConversationItemIDs(t *testing.T, fixture *TestServerFixture, conversationID, order string, limit int) []string { + t.Helper() + var result []string + after := "" + for { + path := fmt.Sprintf("/v1/conversations/%s/items?order=%s&limit=%d", conversationID, order, limit) + if after != "" { + path += "&after=" + after + } + page := conversationJSONRequest(t, fixture, http.MethodGet, path, nil, http.StatusOK) + data := page["data"].([]any) + for _, raw := range data { + result = append(result, raw.(map[string]any)["id"].(string)) + } + if !page["has_more"].(bool) { + return result + } + after = page["last_id"].(string) + } +} + +func uniqueStrings(values []string) map[string]struct{} { + unique := make(map[string]struct{}, len(values)) + for _, value := range values { + unique[value] = struct{}{} + } + return unique +} + +func assertConversationDatabaseState(t *testing.T, fixture *TestServerFixture, conversationID string, metadataCount, itemCount int) { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + switch fixture.DBType { + case "postgresql": + var data string + var items []byte + err := fixture.PgPool.QueryRow(ctx, + "SELECT data, items::text FROM conversation_snapshots WHERE id = $1", conversationID, + ).Scan(&data, &items) + require.NoError(t, err) + require.Contains(t, string(items), "opaque-persisted") + require.NotContains(t, string(items), "future_seed") + var decodedItems []any + require.NoError(t, json.Unmarshal(items, &decodedItems)) + require.Len(t, decodedItems, itemCount) + var decodedData map[string]any + require.NoError(t, json.Unmarshal([]byte(data), &decodedData)) + conversation := decodedData["conversation"].(map[string]any) + require.Len(t, conversation["metadata"].(map[string]any), metadataCount) + case "mongodb": + var doc struct { + Data string `bson:"data"` + Items []string `bson:"items"` + } + err := fixture.MongoDb.Collection("conversation_snapshots").FindOne(ctx, bson.M{"_id": conversationID}).Decode(&doc) + require.NoError(t, err) + require.Len(t, doc.Items, itemCount) + encodedItems, err := json.Marshal(doc.Items) + require.NoError(t, err) + require.Contains(t, string(encodedItems), "opaque-persisted") + require.NotContains(t, string(encodedItems), "future_seed") + var decodedData map[string]any + require.NoError(t, json.Unmarshal([]byte(doc.Data), &decodedData)) + conversation := decodedData["conversation"].(map[string]any) + require.Len(t, conversation["metadata"].(map[string]any), metadataCount) + default: + t.Fatalf("unsupported database type %q", fixture.DBType) + } +} diff --git a/tests/integration/setup_test.go b/tests/integration/setup_test.go index 7a0f1d369..177e7dd72 100644 --- a/tests/integration/setup_test.go +++ b/tests/integration/setup_test.go @@ -191,6 +191,8 @@ func resetPostgreSQLStorage(t *testing.T) { "model_overrides", "aliases", "batches", + "response_snapshots", + "conversation_snapshots", } for _, table := range tables { _, err := pool.Exec(ctx, fmt.Sprintf("DROP TABLE IF EXISTS %s CASCADE", table)) @@ -218,6 +220,8 @@ func resetMongoDBStorage(t *testing.T) { "model_overrides", "aliases", "batches", + "response_snapshots", + "conversation_snapshots", } for _, collection := range collections { err := db.Collection(collection).Drop(ctx)