From c237c2140bd2b547da25bf81eaab9fb283304bed Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Wed, 29 Apr 2026 23:05:23 +0200 Subject: [PATCH 01/15] refactor(llmclient): extract circuit breaker into peer file internal/llmclient/client.go was 848 lines; the circuitBreaker struct, its state constants, and its methods were a self-contained block at the tail. Move them to circuit_breaker.go (same package, no API change). Pure relocation: client.go drops to 720 lines. Co-Authored-By: Claude Opus 4.7 --- internal/llmclient/circuit_breaker.go | 135 ++++++++++++++++++++++++++ internal/llmclient/client.go | 129 ------------------------ 2 files changed, 135 insertions(+), 129 deletions(-) create mode 100644 internal/llmclient/circuit_breaker.go diff --git a/internal/llmclient/circuit_breaker.go b/internal/llmclient/circuit_breaker.go new file mode 100644 index 000000000..4f251dba0 --- /dev/null +++ b/internal/llmclient/circuit_breaker.go @@ -0,0 +1,135 @@ +package llmclient + +import ( + "sync" + "time" +) + +// circuitBreaker implements a circuit breaker pattern with half-open state protection +type circuitBreaker struct { + mu sync.Mutex + state circuitState + failures int + successes int + failureThreshold int + successThreshold int + timeout time.Duration + lastFailure time.Time + halfOpenAllowed bool // Controls single-request probe in half-open state +} + +type circuitState int + +const ( + circuitClosed circuitState = iota + circuitOpen + circuitHalfOpen +) + +func newCircuitBreaker(failureThreshold, successThreshold int, timeout time.Duration) *circuitBreaker { + return &circuitBreaker{ + state: circuitClosed, + failureThreshold: failureThreshold, + successThreshold: successThreshold, + timeout: timeout, + halfOpenAllowed: true, + } +} + +// acquire checks if a request should be allowed through the circuit breaker. +// The second return value reports whether the caller is the single half-open probe. +func (cb *circuitBreaker) acquire() (bool, bool) { + cb.mu.Lock() + defer cb.mu.Unlock() + + switch cb.state { + case circuitClosed: + return true, false + case circuitOpen: + // Check if timeout has passed + if time.Since(cb.lastFailure) > cb.timeout { + cb.state = circuitHalfOpen + cb.successes = 0 + cb.halfOpenAllowed = true // Allow the first probe request + } else { + return false, false + } + // Fall through to half-open handling + fallthrough + case circuitHalfOpen: + // Only allow one request through at a time in half-open state + // This prevents thundering herd when transitioning from open + if cb.halfOpenAllowed { + cb.halfOpenAllowed = false + return true, true + } + return false, false + } + return true, false +} + +// Allow reports whether any request may proceed. +func (cb *circuitBreaker) Allow() bool { + allowed, _ := cb.acquire() + return allowed +} + +// RecordSuccess records a successful request +func (cb *circuitBreaker) RecordSuccess() { + cb.mu.Lock() + defer cb.mu.Unlock() + + switch cb.state { + case circuitHalfOpen: + cb.successes++ + cb.halfOpenAllowed = true // Allow next probe request + if cb.successes >= cb.successThreshold { + cb.state = circuitClosed + cb.failures = 0 + } + case circuitClosed: + cb.failures = 0 + } +} + +// RecordFailure records a failed request +func (cb *circuitBreaker) RecordFailure() { + cb.mu.Lock() + defer cb.mu.Unlock() + + cb.failures++ + cb.lastFailure = time.Now() + + switch cb.state { + case circuitClosed: + if cb.failures >= cb.failureThreshold { + cb.state = circuitOpen + } + case circuitHalfOpen: + cb.state = circuitOpen + cb.successes = 0 + cb.halfOpenAllowed = true // Reset for next timeout period + } +} + +// State returns the current circuit state (for testing/monitoring) +func (cb *circuitBreaker) State() string { + cb.mu.Lock() + defer cb.mu.Unlock() + + switch cb.state { + case circuitClosed: + return "closed" + case circuitOpen: + return "open" + case circuitHalfOpen: + return "half-open" + } + return "unknown" +} + +func (cb *circuitBreaker) IsHalfOpen() bool { + cb.mu.Lock() + defer cb.mu.Unlock() + return cb.state == circuitHalfOpen +} diff --git a/internal/llmclient/client.go b/internal/llmclient/client.go index 313da843e..3bda4ff71 100644 --- a/internal/llmclient/client.go +++ b/internal/llmclient/client.go @@ -717,132 +717,3 @@ func isClientTimeoutGatewayError(err error) bool { } return isTimeoutError(gatewayErr) } - -// circuitBreaker implements a circuit breaker pattern with half-open state protection -type circuitBreaker struct { - mu sync.Mutex - state circuitState - failures int - successes int - failureThreshold int - successThreshold int - timeout time.Duration - lastFailure time.Time - halfOpenAllowed bool // Controls single-request probe in half-open state -} - -type circuitState int - -const ( - circuitClosed circuitState = iota - circuitOpen - circuitHalfOpen -) - -func newCircuitBreaker(failureThreshold, successThreshold int, timeout time.Duration) *circuitBreaker { - return &circuitBreaker{ - state: circuitClosed, - failureThreshold: failureThreshold, - successThreshold: successThreshold, - timeout: timeout, - halfOpenAllowed: true, - } -} - -// acquire checks if a request should be allowed through the circuit breaker. -// The second return value reports whether the caller is the single half-open probe. -func (cb *circuitBreaker) acquire() (bool, bool) { - cb.mu.Lock() - defer cb.mu.Unlock() - - switch cb.state { - case circuitClosed: - return true, false - case circuitOpen: - // Check if timeout has passed - if time.Since(cb.lastFailure) > cb.timeout { - cb.state = circuitHalfOpen - cb.successes = 0 - cb.halfOpenAllowed = true // Allow the first probe request - } else { - return false, false - } - // Fall through to half-open handling - fallthrough - case circuitHalfOpen: - // Only allow one request through at a time in half-open state - // This prevents thundering herd when transitioning from open - if cb.halfOpenAllowed { - cb.halfOpenAllowed = false - return true, true - } - return false, false - } - return true, false -} - -// Allow reports whether any request may proceed. -func (cb *circuitBreaker) Allow() bool { - allowed, _ := cb.acquire() - return allowed -} - -// RecordSuccess records a successful request -func (cb *circuitBreaker) RecordSuccess() { - cb.mu.Lock() - defer cb.mu.Unlock() - - switch cb.state { - case circuitHalfOpen: - cb.successes++ - cb.halfOpenAllowed = true // Allow next probe request - if cb.successes >= cb.successThreshold { - cb.state = circuitClosed - cb.failures = 0 - } - case circuitClosed: - cb.failures = 0 - } -} - -// RecordFailure records a failed request -func (cb *circuitBreaker) RecordFailure() { - cb.mu.Lock() - defer cb.mu.Unlock() - - cb.failures++ - cb.lastFailure = time.Now() - - switch cb.state { - case circuitClosed: - if cb.failures >= cb.failureThreshold { - cb.state = circuitOpen - } - case circuitHalfOpen: - cb.state = circuitOpen - cb.successes = 0 - cb.halfOpenAllowed = true // Reset for next timeout period - } -} - -// State returns the current circuit state (for testing/monitoring) -func (cb *circuitBreaker) State() string { - cb.mu.Lock() - defer cb.mu.Unlock() - - switch cb.state { - case circuitClosed: - return "closed" - case circuitOpen: - return "open" - case circuitHalfOpen: - return "half-open" - } - return "unknown" -} - -func (cb *circuitBreaker) IsHalfOpen() bool { - cb.mu.Lock() - defer cb.mu.Unlock() - return cb.state == circuitHalfOpen -} From ed67fdc114bb20f2ca2d1c9da296522e508469cf Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Wed, 29 Apr 2026 23:10:36 +0200 Subject: [PATCH 02/15] refactor(anthropic): split provider into per-concern peer files internal/providers/anthropic/anthropic.go was 1582 lines mixing provider plumbing, chat completion, chat streaming, responses API (non-streaming + streaming), and the full batch lifecycle. Split into peer files in the same package: chat.go (85) ChatCompletion, convertFromAnthropicResponse chat_stream.go (362) StreamChatCompletion, streamConverter responses.go (382) Responses, StreamResponses, responsesStreamConverter, responses-usage helpers batch.go (366) Create/Get/List/Cancel + result parsing, parseOptionalUnix, mapAnthropicBatchResponse anthropic.go (448) keeps Provider struct + lifecycle helpers, ListModels, content extractors, shared stream/usage helpers, and the no-op Embeddings stub. Pure relocation; tests pass. Co-Authored-By: Claude Opus 4.7 --- internal/providers/anthropic/anthropic.go | 1250 +------------------ internal/providers/anthropic/batch.go | 366 ++++++ internal/providers/anthropic/chat.go | 85 ++ internal/providers/anthropic/chat_stream.go | 362 ++++++ internal/providers/anthropic/responses.go | 382 ++++++ 5 files changed, 1253 insertions(+), 1192 deletions(-) create mode 100644 internal/providers/anthropic/batch.go create mode 100644 internal/providers/anthropic/chat.go create mode 100644 internal/providers/anthropic/chat_stream.go create mode 100644 internal/providers/anthropic/responses.go diff --git a/internal/providers/anthropic/anthropic.go b/internal/providers/anthropic/anthropic.go index 9619e8944..949cc0a02 100644 --- a/internal/providers/anthropic/anthropic.go +++ b/internal/providers/anthropic/anthropic.go @@ -2,23 +2,17 @@ package anthropic import ( - "bufio" "bytes" "context" "encoding/json" - "fmt" "io" "log/slog" "maps" "net/http" - "net/url" - "strconv" "strings" "sync" "time" - "github.com/google/uuid" - "gomodel/internal/core" "gomodel/internal/llmclient" "gomodel/internal/providers" @@ -223,536 +217,6 @@ func normalizeEffort(effort string) string { } } -// convertFromAnthropicResponse converts Anthropic response to core.ChatResponse -func convertFromAnthropicResponse(resp *anthropicResponse) *core.ChatResponse { - content := extractTextContent(resp.Content) - thinking := extractThinkingContent(resp.Content) - toolCalls := extractToolCalls(resp.Content) - - finishReason := normalizeAnthropicStopReason(resp.StopReason) - if finishReason == "" { - finishReason = "stop" - } - - usage := core.Usage{ - PromptTokens: resp.Usage.InputTokens, - CompletionTokens: resp.Usage.OutputTokens, - TotalTokens: resp.Usage.InputTokens + resp.Usage.OutputTokens, - } - - rawUsage := buildAnthropicRawUsage(resp.Usage) - if len(rawUsage) > 0 { - usage.RawUsage = rawUsage - } - - msg := core.ResponseMessage{ - Role: "assistant", - Content: content, - ToolCalls: toolCalls, - } - - // Surface thinking content as reasoning_content (OpenAI-compatible format). - if thinking != "" { - raw, err := json.Marshal(thinking) - if err == nil { - msg.ExtraFields = core.UnknownJSONFieldsFromMap(map[string]json.RawMessage{ - "reasoning_content": raw, - }) - } - } - - return &core.ChatResponse{ - ID: resp.ID, - Object: "chat.completion", - Model: resp.Model, - Created: time.Now().Unix(), - Choices: []core.Choice{ - { - Index: 0, - Message: msg, - FinishReason: finishReason, - }, - }, - Usage: usage, - } -} - -// ChatCompletion sends a chat completion request to Anthropic -func (p *Provider) ChatCompletion(ctx context.Context, req *core.ChatRequest) (*core.ChatResponse, error) { - anthropicReq, err := convertToAnthropicRequest(req) - if err != nil { - return nil, err - } - - var anthropicResp anthropicResponse - err = p.client.Do(ctx, llmclient.Request{ - Method: http.MethodPost, - Endpoint: "/messages", - Body: anthropicReq, - }, &anthropicResp) - if err != nil { - return nil, err - } - - return convertFromAnthropicResponse(&anthropicResp), nil -} - -// StreamChatCompletion returns a raw response body for streaming (caller must close) -func (p *Provider) StreamChatCompletion(ctx context.Context, req *core.ChatRequest) (io.ReadCloser, error) { - anthropicReq, err := convertToAnthropicRequest(req) - if err != nil { - return nil, err - } - anthropicReq.Stream = true - - stream, err := p.client.DoStream(ctx, llmclient.Request{ - Method: http.MethodPost, - Endpoint: "/messages", - Body: anthropicReq, - }) - if err != nil { - return nil, err - } - - // Return a reader that converts Anthropic SSE format to OpenAI format - return newStreamConverter(stream, req.Model), nil -} - -// streamConverter wraps an Anthropic stream and converts it to OpenAI format -type streamConverter struct { - reader *bufio.Reader - body io.ReadCloser - model string - msgID string - nextToolCallIndex int - toolCalls map[int]*streamToolCallState - thinkingBlocks map[int]bool // tracks which content block indices are thinking blocks - usage anthropicUsage - hasUsage bool - buffer streaming.StreamBuffer - closed bool - emittedToolCalls bool -} - -type streamToolCallState struct { - ID string - Name string - Arguments strings.Builder - Index int - Started bool - PlaceholderObject bool -} - -func newStreamConverter(body io.ReadCloser, model string) *streamConverter { - return &streamConverter{ - reader: bufio.NewReader(body), - body: body, - model: model, - toolCalls: make(map[int]*streamToolCallState), - thinkingBlocks: make(map[int]bool), - buffer: streaming.NewStreamBuffer(1024), - } -} - -func malformedAnthropicStreamError(err error) error { - return core.NewProviderError("anthropic", http.StatusBadGateway, "failed to decode anthropic stream event: "+err.Error(), err) -} - -func consumeAnthropicSSELine(p []byte, line []byte, body io.ReadCloser, buffer *streaming.StreamBuffer, convert func(*anthropicStreamEvent) string) (n int, handled bool, err error) { - line = bytes.TrimSpace(line) - if len(line) == 0 || bytes.HasPrefix(line, []byte("event:")) { - return 0, false, nil - } - if !bytes.HasPrefix(line, []byte("data:")) { - return 0, false, nil - } - - data := bytes.TrimSpace(bytes.TrimPrefix(line, []byte("data:"))) - - var event anthropicStreamEvent - if err := json.Unmarshal(data, &event); err != nil { - _ = body.Close() //nolint:errcheck - return 0, false, malformedAnthropicStreamError(err) - } - - chunk := convert(&event) - if chunk == "" { - return 0, false, nil - } - - buffer.AppendString(chunk) - return buffer.Read(p), true, nil -} - -func mergeAnthropicUsage(dst *anthropicUsage, src *anthropicUsage) bool { - if dst == nil || src == nil { - return false - } - - merged := false - if src.InputTokens != 0 { - dst.InputTokens = src.InputTokens - merged = true - } - if src.OutputTokens != 0 { - dst.OutputTokens = src.OutputTokens - merged = true - } - if src.CacheCreationInputTokens != 0 { - dst.CacheCreationInputTokens = src.CacheCreationInputTokens - merged = true - } - if src.CacheReadInputTokens != 0 { - dst.CacheReadInputTokens = src.CacheReadInputTokens - merged = true - } - - return merged -} - -func anthropicChatUsagePayload(usage *anthropicUsage) map[string]any { - if usage == nil { - return nil - } - - payload := map[string]any{ - "prompt_tokens": usage.InputTokens, - "completion_tokens": usage.OutputTokens, - "total_tokens": usage.InputTokens + usage.OutputTokens, - } - if usage.CacheReadInputTokens > 0 { - payload["cache_read_input_tokens"] = usage.CacheReadInputTokens - } - if usage.CacheCreationInputTokens > 0 { - payload["cache_creation_input_tokens"] = usage.CacheCreationInputTokens - } - return payload -} - -func anthropicResponsesUsagePayload(usage *anthropicUsage) map[string]any { - if usage == nil { - return nil - } - - payload := map[string]any{ - "input_tokens": usage.InputTokens, - "output_tokens": usage.OutputTokens, - "total_tokens": usage.InputTokens + usage.OutputTokens, - } - if usage.CacheReadInputTokens > 0 { - payload["cache_read_input_tokens"] = usage.CacheReadInputTokens - } - if usage.CacheCreationInputTokens > 0 { - payload["cache_creation_input_tokens"] = usage.CacheCreationInputTokens - } - return payload -} - -func (sc *streamConverter) Read(p []byte) (n int, err error) { - // If we have buffered data, return it first - if sc.buffer.Len() > 0 { - return sc.buffer.Read(p), nil - } - - if sc.closed { - sc.releaseBuffer() - return 0, io.EOF - } - - // Read the next SSE event from Anthropic - for { - line, err := sc.reader.ReadBytes('\n') - if err != nil { - if err == io.EOF { - // Send final [DONE] message - sc.buffer.AppendString("data: [DONE]\n\n") - n = sc.buffer.Read(p) - sc.closed = true - _ = sc.body.Close() //nolint:errcheck - return n, nil - } - return 0, err - } - - n, handled, err := consumeAnthropicSSELine(p, line, sc.body, &sc.buffer, sc.convertEvent) - if err != nil { - sc.closed = true - sc.releaseBuffer() - return 0, err - } - if handled { - if n == 0 { - continue - } - return n, nil - } - } -} - -func (sc *streamConverter) Close() error { - if sc.closed { - sc.releaseBuffer() - return nil - } - sc.closed = true - sc.releaseBuffer() - return sc.body.Close() -} - -func (sc *streamConverter) releaseBuffer() { - sc.buffer.Release() -} - -func (sc *streamConverter) mapStreamStopReason(reason string) string { - // Preserve raw "tool_use" when the upstream stream never produced any - // tool call deltas. This avoids claiming OpenAI-style tool calls for a - // malformed or partial Anthropic stream. - if reason == "tool_use" && !sc.emittedToolCalls { - return reason - } - return normalizeAnthropicStopReason(reason) -} - -func extractInitialToolArguments(input json.RawMessage) string { - if len(input) == 0 { - return "" - } - - trimmed := strings.TrimSpace(string(input)) - if trimmed == "" || trimmed == "null" { - return "" - } - - var parsed any - if err := json.Unmarshal(input, &parsed); err != nil { - return trimmed - } - - canonical, err := json.Marshal(parsed) - if err != nil { - return trimmed - } - - return string(canonical) -} - -func normalizeAnthropicStopReason(stopReason string) string { - switch stopReason { - case "tool_use": - return "tool_calls" - case "end_turn", "stop_sequence": - return "stop" - case "max_tokens", "model_context_window_exceeded": - return "length" - default: - return stopReason - } -} - -func (sc *streamConverter) formatChatChunk(delta map[string]any, finishReason any, usage *anthropicUsage) string { - chunk := map[string]any{ - "id": sc.msgID, - "object": "chat.completion.chunk", - "created": time.Now().Unix(), - "model": sc.model, - "provider": "anthropic", - "choices": []map[string]any{ - { - "index": 0, - "delta": delta, - "finish_reason": finishReason, - }, - }, - } - if usage != nil { - chunk["usage"] = anthropicChatUsagePayload(usage) - } - - jsonData, err := json.Marshal(chunk) - if err != nil { - slog.Error("failed to marshal chat completion chunk", "error", err, "msg_id", sc.msgID) - return "" - } - - return fmt.Sprintf("data: %s\n\n", jsonData) -} - -func (sc *streamConverter) convertEvent(event *anthropicStreamEvent) string { - switch event.Type { - case "message_start": - role := "" - if event.Message != nil { - sc.msgID = event.Message.ID - if mergeAnthropicUsage(&sc.usage, &event.Message.Usage) { - sc.hasUsage = true - } - role = strings.TrimSpace(event.Message.Role) - } - if mergeAnthropicUsage(&sc.usage, event.Usage) { - sc.hasUsage = true - } - if event.Message != nil { - if role == "" { - role = "assistant" - } - return sc.formatChatChunk(map[string]any{ - "role": role, - }, nil, nil) - } - return "" - - case "content_block_start": - if event.ContentBlock != nil && event.ContentBlock.Type == "thinking" { - sc.thinkingBlocks[event.Index] = true - return "" - } - if event.ContentBlock != nil && event.ContentBlock.Type == "tool_use" { - state := &streamToolCallState{ - ID: event.ContentBlock.ID, - Name: event.ContentBlock.Name, - Index: sc.nextToolCallIndex, - } - sc.nextToolCallIndex++ - - initialArguments := extractInitialToolArguments(event.ContentBlock.Input) - state.PlaceholderObject = initialArguments == "{}" - if state.PlaceholderObject { - sc.toolCalls[event.Index] = state - return "" - } - if initialArguments != "" { - _, _ = state.Arguments.WriteString(initialArguments) - } - state.Started = true - sc.toolCalls[event.Index] = state - sc.emittedToolCalls = true - - return sc.formatChatChunk(map[string]any{ - "tool_calls": []map[string]any{ - { - "index": state.Index, - "id": state.ID, - "type": "function", - "function": map[string]any{ - "name": state.Name, - "arguments": initialArguments, - }, - }, - }, - }, nil, nil) - } - return "" - - case "content_block_delta": - if event.Delta == nil { - return "" - } - - switch event.Delta.Type { - case "thinking_delta": - if sc.thinkingBlocks[event.Index] && event.Delta.Thinking != "" { - return sc.formatChatChunk(map[string]any{ - "reasoning_content": event.Delta.Thinking, - }, nil, nil) - } - case "signature_delta": - // Signature deltas are internal to Anthropic's thinking protocol; - // no OpenAI-compatible equivalent to emit. - return "" - case "text_delta": - if event.Delta.Text != "" { - return sc.formatChatChunk(map[string]any{ - "content": event.Delta.Text, - }, nil, nil) - } - case "input_json_delta": - if event.Delta.PartialJSON == "" { - return "" - } - state := sc.toolCalls[event.Index] - if state == nil { - return "" - } - if state.PlaceholderObject { - state.Arguments = strings.Builder{} - state.PlaceholderObject = false - } - _, _ = state.Arguments.WriteString(event.Delta.PartialJSON) - if !state.Started { - state.Started = true - sc.emittedToolCalls = true - return sc.formatChatChunk(map[string]any{ - "tool_calls": []map[string]any{ - { - "index": state.Index, - "id": state.ID, - "type": "function", - "function": map[string]any{ - "name": state.Name, - "arguments": event.Delta.PartialJSON, - }, - }, - }, - }, nil, nil) - } - sc.emittedToolCalls = true - return sc.formatChatChunk(map[string]any{ - "tool_calls": []map[string]any{ - { - "index": state.Index, - "function": map[string]any{ - "arguments": event.Delta.PartialJSON, - }, - }, - }, - }, nil, nil) - } - - case "content_block_stop": - state := sc.toolCalls[event.Index] - if state != nil && !state.Started && state.PlaceholderObject { - state.Started = true - sc.emittedToolCalls = true - return sc.formatChatChunk(map[string]any{ - "tool_calls": []map[string]any{ - { - "index": state.Index, - "id": state.ID, - "type": "function", - "function": map[string]any{ - "name": state.Name, - "arguments": "{}", - }, - }, - }, - }, nil, nil) - } - return "" - - case "message_delta": - if mergeAnthropicUsage(&sc.usage, event.Usage) { - sc.hasUsage = true - } - // Emit chunk if we have stop_reason or usage data - if (event.Delta != nil && event.Delta.StopReason != "") || event.Usage != nil { - var finishReason any - if event.Delta != nil && event.Delta.StopReason != "" { - finishReason = sc.mapStreamStopReason(event.Delta.StopReason) - } - var usage *anthropicUsage - if sc.hasUsage { - usage = &sc.usage - } - return sc.formatChatChunk(map[string]any{}, finishReason, usage) - } - - case "message_stop": - return "" - } - - return "" -} - // ListModels retrieves the list of available models from Anthropic's /v1/models endpoint func (p *Provider) ListModels(ctx context.Context) (*core.ModelsResponse, error) { var anthropicResp anthropicModelsResponse @@ -870,32 +334,6 @@ func extractToolCalls(blocks []anthropicContent) []core.ToolCall { return out } -// convertAnthropicResponseToResponses converts an Anthropic response to ResponsesResponse -func convertAnthropicResponseToResponses(resp *anthropicResponse, model string) *core.ResponsesResponse { - content := extractTextContent(resp.Content) - toolCalls := extractToolCalls(resp.Content) - - msg := core.Message{ - Content: content, - ToolCalls: toolCalls, - } - output := providers.BuildResponsesOutputItems(core.ResponseMessage{ - Role: "assistant", - Content: msg.Content, - ToolCalls: msg.ToolCalls, - }) - - return &core.ResponsesResponse{ - ID: resp.ID, - Object: "response", - CreatedAt: time.Now().Unix(), - Model: model, - Status: "completed", - Output: output, - Usage: buildAnthropicResponsesUsage(resp.Usage), - } -} - // buildAnthropicRawUsage extracts cache fields from anthropicUsage into a RawData map. func buildAnthropicRawUsage(u anthropicUsage) map[string]any { raw := make(map[string]any) @@ -911,385 +349,96 @@ func buildAnthropicRawUsage(u anthropicUsage) map[string]any { return raw } -// buildAnthropicResponsesUsage creates a ResponsesUsage from anthropicUsage, including RawUsage. -func buildAnthropicResponsesUsage(u anthropicUsage) *core.ResponsesUsage { - usage := &core.ResponsesUsage{ - InputTokens: u.InputTokens, - OutputTokens: u.OutputTokens, - TotalTokens: u.InputTokens + u.OutputTokens, - } - rawUsage := buildAnthropicRawUsage(u) - if len(rawUsage) > 0 { - usage.RawUsage = rawUsage - } - return usage +func malformedAnthropicStreamError(err error) error { + return core.NewProviderError("anthropic", http.StatusBadGateway, "failed to decode anthropic stream event: "+err.Error(), err) } -// Responses sends a Responses API request to Anthropic (converted to messages format) -func (p *Provider) Responses(ctx context.Context, req *core.ResponsesRequest) (*core.ResponsesResponse, error) { - anthropicReq, err := convertResponsesRequestToAnthropic(req) - if err != nil { - return nil, err +func consumeAnthropicSSELine(p []byte, line []byte, body io.ReadCloser, buffer *streaming.StreamBuffer, convert func(*anthropicStreamEvent) string) (n int, handled bool, err error) { + line = bytes.TrimSpace(line) + if len(line) == 0 || bytes.HasPrefix(line, []byte("event:")) { + return 0, false, nil } - - var anthropicResp anthropicResponse - err = p.client.Do(ctx, llmclient.Request{ - Method: http.MethodPost, - Endpoint: "/messages", - Body: anthropicReq, - }, &anthropicResp) - if err != nil { - return nil, err + if !bytes.HasPrefix(line, []byte("data:")) { + return 0, false, nil } - return convertAnthropicResponseToResponses(&anthropicResp, req.Model), nil -} - -func parseOptionalUnix(ts string) *int64 { - ts = strings.TrimSpace(ts) - if ts == "" { - return nil - } - t, err := time.Parse(time.RFC3339, ts) - if err != nil { - return nil - } - u := t.Unix() - return &u -} + data := bytes.TrimSpace(bytes.TrimPrefix(line, []byte("data:"))) -func mapAnthropicBatchResponse(resp *anthropicBatchResponse) *core.BatchResponse { - if resp == nil { - return nil + var event anthropicStreamEvent + if err := json.Unmarshal(data, &event); err != nil { + _ = body.Close() //nolint:errcheck + return 0, false, malformedAnthropicStreamError(err) } - total := resp.RequestCounts.Processing + resp.RequestCounts.Succeeded + resp.RequestCounts.Errored + resp.RequestCounts.Canceled + resp.RequestCounts.Expired - failed := resp.RequestCounts.Errored + resp.RequestCounts.Canceled + resp.RequestCounts.Expired - - status := "in_progress" - switch resp.ProcessingStatus { - case "canceling": - status = "cancelling" - case "ended": - switch { - case resp.RequestCounts.Canceled > 0 && resp.RequestCounts.Succeeded == 0 && resp.RequestCounts.Errored == 0: - status = "cancelled" - case resp.RequestCounts.Errored > 0 && resp.RequestCounts.Succeeded == 0: - status = "failed" - default: - status = "completed" - } + chunk := convert(&event) + if chunk == "" { + return 0, false, nil } - return &core.BatchResponse{ - ID: resp.ID, - Object: "batch", - Status: status, - CreatedAt: parseCreatedAt(resp.CreatedAt), - CompletedAt: parseOptionalUnix(resp.EndedAt), - CancellingAt: parseOptionalUnix(resp.CancelInitiatedAt), - RequestCounts: core.BatchRequestCounts{ - Total: total, - Completed: resp.RequestCounts.Succeeded, - Failed: failed, - }, - } + buffer.AppendString(chunk) + return buffer.Read(p), true, nil } -func buildAnthropicBatchCreateRequest(req *core.BatchRequest) (*anthropicBatchCreateRequest, map[string]string, error) { - const maxAnthropicBatchRequests = 10000 - - if req == nil { - return nil, nil, core.NewInvalidRequestError("request is required for anthropic batch processing", nil) - } - if len(req.Requests) == 0 { - return nil, nil, core.NewInvalidRequestError("requests is required for anthropic batch processing", nil) - } - if len(req.Requests) > maxAnthropicBatchRequests { - return nil, nil, core.NewInvalidRequestError("too many requests for anthropic batch processing", nil) - } - - out := &anthropicBatchCreateRequest{ - Requests: make([]anthropicBatchRequest, 0, len(req.Requests)), +func mergeAnthropicUsage(dst *anthropicUsage, src *anthropicUsage) bool { + if dst == nil || src == nil { + return false } - endpointByCustomID := make(map[string]string, len(req.Requests)) - seenCustomIDs := make(map[string]int, len(req.Requests)) - - for i, item := range req.Requests { - decoded, err := core.DecodeKnownBatchItemRequest(req.Endpoint, item) - if err != nil { - return nil, nil, core.NewInvalidRequestError(fmt.Sprintf("batch item %d: %s", i, err.Error()), err) - } - params, err := convertDecodedBatchItemToAnthropic(decoded) - if err != nil { - return nil, nil, prefixAnthropicBatchItemError(i, err) - } - - customID := strings.TrimSpace(item.CustomID) - if customID == "" { - customID = fmt.Sprintf("req-%d", i) - } - if previousIndex, exists := seenCustomIDs[customID]; exists { - return nil, nil, core.NewInvalidRequestError( - fmt.Sprintf("batch item %d: duplicate custom_id %q (already used by batch item %d)", i, customID, previousIndex), - nil, - ) - } - seenCustomIDs[customID] = i - out.Requests = append(out.Requests, anthropicBatchRequest{ - CustomID: customID, - Params: *params, - }) - endpointByCustomID[customID] = decoded.Endpoint + merged := false + if src.InputTokens != 0 { + dst.InputTokens = src.InputTokens + merged = true } - - return out, endpointByCustomID, nil -} - -func (p *Provider) createBatch(ctx context.Context, req *core.BatchRequest) (*core.BatchResponse, map[string]string, error) { - anthropicReq, endpointByCustomID, err := buildAnthropicBatchCreateRequest(req) - if err != nil { - return nil, nil, err + if src.OutputTokens != 0 { + dst.OutputTokens = src.OutputTokens + merged = true } - - var resp anthropicBatchResponse - err = p.client.Do(ctx, llmclient.Request{ - Method: http.MethodPost, - Endpoint: "/messages/batches", - Body: anthropicReq, - }, &resp) - if err != nil { - return nil, nil, err + if src.CacheCreationInputTokens != 0 { + dst.CacheCreationInputTokens = src.CacheCreationInputTokens + merged = true } - - mapped := mapAnthropicBatchResponse(&resp) - if mapped == nil { - return nil, nil, core.NewProviderError("anthropic", http.StatusBadGateway, "failed to map anthropic batch response", nil) + if src.CacheReadInputTokens != 0 { + dst.CacheReadInputTokens = src.CacheReadInputTokens + merged = true } - mapped.ProviderBatchID = mapped.ID - p.setBatchResultEndpoints(mapped.ProviderBatchID, endpointByCustomID) - return mapped, cloneBatchResultEndpoints(endpointByCustomID), nil -} - -// CreateBatch creates an Anthropic native message batch. -func (p *Provider) CreateBatch(ctx context.Context, req *core.BatchRequest) (*core.BatchResponse, error) { - mapped, _, err := p.createBatch(ctx, req) - return mapped, err -} - -// CreateBatchWithHints creates an Anthropic native message batch and returns -// persisted per-item endpoint hints for later result shaping. -func (p *Provider) CreateBatchWithHints(ctx context.Context, req *core.BatchRequest) (*core.BatchResponse, map[string]string, error) { - return p.createBatch(ctx, req) -} -// GetBatch retrieves an Anthropic native message batch. -func (p *Provider) GetBatch(ctx context.Context, id string) (*core.BatchResponse, error) { - var resp anthropicBatchResponse - err := p.client.Do(ctx, llmclient.Request{ - Method: http.MethodGet, - Endpoint: "/messages/batches/" + url.PathEscape(id), - }, &resp) - if err != nil { - return nil, err - } - mapped := mapAnthropicBatchResponse(&resp) - if mapped == nil { - return nil, core.NewProviderError("anthropic", http.StatusBadGateway, "failed to map anthropic batch response", nil) - } - mapped.ProviderBatchID = mapped.ID - return mapped, nil + return merged } -// ListBatches lists Anthropic native message batches. -func (p *Provider) ListBatches(ctx context.Context, limit int, after string) (*core.BatchListResponse, error) { - values := url.Values{} - if limit > 0 { - values.Set("limit", strconv.Itoa(limit)) - } - // Anthropic uses before_id for reverse-chronological pagination. - // Gateway `after` is mapped directly to before_id for provider-native paging. - if after != "" { - values.Set("before_id", after) - } - endpoint := "/messages/batches" - if encoded := values.Encode(); encoded != "" { - endpoint += "?" + encoded - } - - var resp anthropicBatchListResponse - err := p.client.Do(ctx, llmclient.Request{ - Method: http.MethodGet, - Endpoint: endpoint, - }, &resp) - if err != nil { - return nil, err +func extractInitialToolArguments(input json.RawMessage) string { + if len(input) == 0 { + return "" } - data := make([]core.BatchResponse, 0, len(resp.Data)) - for _, row := range resp.Data { - mapped := mapAnthropicBatchResponse(&row) - if mapped == nil { - continue - } - mapped.ProviderBatchID = mapped.ID - data = append(data, *mapped) + trimmed := strings.TrimSpace(string(input)) + if trimmed == "" || trimmed == "null" { + return "" } - return &core.BatchListResponse{ - Object: "list", - Data: data, - HasMore: resp.HasMore, - FirstID: resp.FirstID, - LastID: resp.LastID, - }, nil -} - -// CancelBatch cancels an Anthropic native message batch. -func (p *Provider) CancelBatch(ctx context.Context, id string) (*core.BatchResponse, error) { - var resp anthropicBatchResponse - err := p.client.Do(ctx, llmclient.Request{ - Method: http.MethodPost, - Endpoint: "/messages/batches/" + url.PathEscape(id) + "/cancel", - }, &resp) - if err != nil { - return nil, err - } - mapped := mapAnthropicBatchResponse(&resp) - if mapped == nil { - return nil, core.NewProviderError("anthropic", http.StatusBadGateway, "failed to map anthropic batch response", nil) + var parsed any + if err := json.Unmarshal(input, &parsed); err != nil { + return trimmed } - mapped.ProviderBatchID = mapped.ID - return mapped, nil -} -func (p *Provider) getBatchResults(ctx context.Context, id string, endpointByCustomID map[string]string) (*core.BatchResultsResponse, error) { - resp, err := p.client.DoPassthrough(ctx, llmclient.Request{ - Method: http.MethodGet, - Endpoint: "/messages/batches/" + url.PathEscape(id) + "/results", - }) + canonical, err := json.Marshal(parsed) if err != nil { - return nil, err - } - defer func() { _ = resp.Body.Close() }() - - if resp.StatusCode != http.StatusOK { - body, readErr := io.ReadAll(resp.Body) - if readErr != nil { - body = []byte("failed to read error response") - } - return nil, core.ParseProviderError("anthropic", resp.StatusCode, body, nil) - } - - scanner := bufio.NewScanner(resp.Body) - // Allow larger result lines than Scanner's default 64K. - scanner.Buffer(make([]byte, 0, 64*1024), 4*1024*1024) - if endpointByCustomID == nil { - endpointByCustomID = p.getBatchResultEndpoints(id) - } else { - endpointByCustomID = cloneBatchResultEndpoints(endpointByCustomID) - } - - results := make([]core.BatchResultItem, 0) - index := 0 - for scanner.Scan() { - line := bytes.TrimSpace(scanner.Bytes()) - if len(line) == 0 { - continue - } - - var row anthropicBatchResultLine - if err := json.Unmarshal(line, &row); err != nil { - slog.Warn( - "failed to decode anthropic batch result line", - "error", err, - "batch_id", id, - "line_index", index, - "line_bytes", len(line), - ) - continue - } - itemEndpoint := "/v1/chat/completions" - if endpointByCustomID != nil { - if endpoint := strings.TrimSpace(endpointByCustomID[row.CustomID]); endpoint != "" { - itemEndpoint = endpoint - } - } - - item := core.BatchResultItem{ - Index: index, - CustomID: row.CustomID, - URL: itemEndpoint, - Provider: "anthropic", - } - switch row.Result.Type { - case "succeeded": - item.StatusCode = http.StatusOK - if len(row.Result.Message) > 0 { - var anthropicPayload anthropicResponse - if err := json.Unmarshal(row.Result.Message, &anthropicPayload); err == nil { - switch itemEndpoint { - case "/v1/responses": - mapped := convertAnthropicResponseToResponses(&anthropicPayload, anthropicPayload.Model) - item.Response = mapped - item.Model = mapped.Model - default: - mapped := convertFromAnthropicResponse(&anthropicPayload) - item.Response = mapped - item.Model = mapped.Model - } - } else { - item.Response = string(row.Result.Message) - } - } - default: - item.StatusCode = http.StatusBadRequest - errType := row.Result.Type - errMsg := "batch item failed" - if row.Result.Error != nil { - if row.Result.Error.Type != "" { - errType = row.Result.Error.Type - } - if row.Result.Error.Message != "" { - errMsg = row.Result.Error.Message - } - } - item.Error = &core.BatchError{ - Type: errType, - Message: errMsg, - } - } - - results = append(results, item) - index++ - } - if err := scanner.Err(); err != nil { - return nil, core.NewProviderError("anthropic", http.StatusBadGateway, "failed to parse anthropic batch results", err) + return trimmed } - return &core.BatchResultsResponse{ - Object: "list", - BatchID: id, - Data: results, - }, nil -} - -// GetBatchResults retrieves Anthropic native message batch results. -func (p *Provider) GetBatchResults(ctx context.Context, id string) (*core.BatchResultsResponse, error) { - return p.getBatchResults(ctx, id, nil) -} - -// GetBatchResultsWithHints retrieves Anthropic native batch results using -// persisted per-item endpoint hints instead of transient in-memory state. -func (p *Provider) GetBatchResultsWithHints(ctx context.Context, id string, endpointByCustomID map[string]string) (*core.BatchResultsResponse, error) { - return p.getBatchResults(ctx, id, endpointByCustomID) + return string(canonical) } -// ClearBatchResultHints clears transient per-batch endpoint hints once they -// have been persisted by the gateway. -func (p *Provider) ClearBatchResultHints(batchID string) { - p.clearBatchResultEndpoints(batchID) +func normalizeAnthropicStopReason(stopReason string) string { + switch stopReason { + case "tool_use": + return "tool_calls" + case "end_turn", "stop_sequence": + return "stop" + case "max_tokens", "model_context_window_exceeded": + return "length" + default: + return stopReason + } } // Embeddings returns an error because Anthropic does not natively support embeddings. @@ -1297,286 +446,3 @@ func (p *Provider) ClearBatchResultHints(batchID string) { func (p *Provider) Embeddings(_ context.Context, _ *core.EmbeddingRequest) (*core.EmbeddingResponse, error) { return nil, core.NewInvalidRequestError("anthropic does not support embeddings — consider using Voyage AI", nil) } - -// StreamResponses returns a raw response body for streaming Responses API (caller must close) -func (p *Provider) StreamResponses(ctx context.Context, req *core.ResponsesRequest) (io.ReadCloser, error) { - anthropicReq, err := convertResponsesRequestToAnthropic(req) - if err != nil { - return nil, err - } - anthropicReq.Stream = true - - stream, err := p.client.DoStream(ctx, llmclient.Request{ - Method: http.MethodPost, - Endpoint: "/messages", - Body: anthropicReq, - }) - if err != nil { - return nil, err - } - - // Return a reader that converts Anthropic SSE format to Responses API format - return newResponsesStreamConverter(stream, req.Model), nil -} - -// responsesStreamConverter wraps an Anthropic stream and converts it to Responses API format -type responsesStreamConverter struct { - reader *bufio.Reader - body io.ReadCloser - model string - responseID string - output *providers.ResponsesOutputEventState - nextOutputIndex int - toolCalls map[int]*providers.ResponsesOutputToolCallState - thinkingBlocks map[int]bool // tracks which content block indices are thinking blocks - buffer streaming.StreamBuffer - closed bool - sentDone bool - usage anthropicUsage - hasUsage bool -} - -func newResponsesStreamConverter(body io.ReadCloser, model string) *responsesStreamConverter { - responseID := "resp_" + uuid.New().String() - return &responsesStreamConverter{ - reader: bufio.NewReader(body), - body: body, - model: model, - responseID: responseID, - output: providers.NewResponsesOutputEventState(responseID), - toolCalls: make(map[int]*providers.ResponsesOutputToolCallState), - thinkingBlocks: make(map[int]bool), - buffer: streaming.NewStreamBuffer(1024), - } -} - -func (sc *responsesStreamConverter) Read(p []byte) (n int, err error) { - if sc.closed { - sc.releaseBuffer() - return 0, io.EOF - } - - // If we have buffered data, return it first - if sc.buffer.Len() > 0 { - return sc.buffer.Read(p), nil - } - - // Read the next SSE event from Anthropic - for { - line, err := sc.reader.ReadBytes('\n') - if err != nil { - if err == io.EOF { - // Send final done event and [DONE] message - if !sc.sentDone { - sc.sentDone = true - prefix := sc.output.CompleteAssistantOutput(0) - responseData := map[string]any{ - "id": sc.responseID, - "object": "response", - "status": "completed", - "model": sc.model, - "provider": "anthropic", - "created_at": time.Now().Unix(), - } - // Include merged usage data captured across message_start/message_delta. - if sc.hasUsage { - responseData["usage"] = anthropicResponsesUsagePayload(&sc.usage) - } - doneEvent := map[string]any{ - "type": "response.completed", - "response": responseData, - } - jsonData, marshalErr := json.Marshal(doneEvent) - if marshalErr != nil { - slog.Error("failed to marshal response.completed event", "error", marshalErr, "response_id", sc.responseID) - sc.closed = true - sc.releaseBuffer() - _ = sc.body.Close() //nolint:errcheck - return 0, io.EOF - } - sc.buffer.AppendString(prefix) - sc.buffer.AppendString("event: response.completed\ndata: ") - sc.buffer.AppendBytes(jsonData) - sc.buffer.AppendString("\n\ndata: [DONE]\n\n") - return sc.buffer.Read(p), nil - } - sc.closed = true - sc.releaseBuffer() - _ = sc.body.Close() //nolint:errcheck - return 0, io.EOF - } - return 0, err - } - - n, handled, err := consumeAnthropicSSELine(p, line, sc.body, &sc.buffer, sc.convertEvent) - if err != nil { - sc.closed = true - sc.releaseBuffer() - return 0, err - } - if handled { - if n == 0 { - continue - } - return n, nil - } - } -} - -func (sc *responsesStreamConverter) Close() error { - if sc.closed { - sc.releaseBuffer() - return nil - } - sc.closed = true - sc.releaseBuffer() - return sc.body.Close() -} - -func (sc *responsesStreamConverter) releaseBuffer() { - sc.buffer.Release() -} - -func (sc *responsesStreamConverter) reserveAssistantMessageOutput() { - if sc.output.AssistantReserved() { - return - } - sc.output.ReserveAssistant() - sc.nextOutputIndex++ -} - -func (sc *responsesStreamConverter) newResponsesToolCallState(contentBlock *anthropicContent) *providers.ResponsesOutputToolCallState { - callID := providers.ResponsesFunctionCallCallID(contentBlock.ID) - state := &providers.ResponsesOutputToolCallState{ - CallID: callID, - Name: contentBlock.Name, - OutputIndex: sc.nextOutputIndex, - } - sc.nextOutputIndex++ - - initialArguments := extractInitialToolArguments(contentBlock.Input) - state.PlaceholderObject = initialArguments == "{}" - if initialArguments != "" && !state.PlaceholderObject { - _, _ = state.Arguments.WriteString(initialArguments) - } - - return state -} - -func (sc *responsesStreamConverter) convertEvent(event *anthropicStreamEvent) string { - switch event.Type { - case "message_start": - if event.Message != nil { - if mergeAnthropicUsage(&sc.usage, &event.Message.Usage) { - sc.hasUsage = true - } - } - if mergeAnthropicUsage(&sc.usage, event.Usage) { - sc.hasUsage = true - } - // Send response.created event - createdEvent := map[string]any{ - "type": "response.created", - "response": map[string]any{ - "id": sc.responseID, - "object": "response", - "status": "in_progress", - "model": sc.model, - "provider": "anthropic", - "created_at": time.Now().Unix(), - }, - } - jsonData, err := json.Marshal(createdEvent) - if err != nil { - slog.Error("failed to marshal response.created event", "error", err, "response_id", sc.responseID) - return "" - } - return fmt.Sprintf("event: response.created\ndata: %s\n\n", jsonData) - - case "content_block_start": - if event.ContentBlock != nil && event.ContentBlock.Type == "thinking" { - sc.thinkingBlocks[event.Index] = true - return "" - } - if event.ContentBlock != nil && event.ContentBlock.Type == "tool_use" { - if sc.output.AssistantStarted() && !sc.output.AssistantDone() { - prefix := sc.output.CompleteAssistantOutput(0) - state := sc.newResponsesToolCallState(event.ContentBlock) - sc.toolCalls[event.Index] = state - return prefix + sc.output.StartToolCall(state, true) - } - state := sc.newResponsesToolCallState(event.ContentBlock) - sc.toolCalls[event.Index] = state - return sc.output.StartToolCall(state, true) - } - return "" - - case "content_block_delta": - if event.Delta == nil { - return "" - } - - switch event.Delta.Type { - case "thinking_delta", "signature_delta": - // Thinking and signature deltas are part of Anthropic's extended thinking; - // the Responses API format does not have a direct equivalent, so skip them. - return "" - case "text_delta": - if event.Delta.Text != "" { - sc.reserveAssistantMessageOutput() - prefix := sc.output.StartAssistantOutput(0) - sc.output.AppendAssistantText(event.Delta.Text) - deltaEvent := map[string]any{ - "type": "response.output_text.delta", - "delta": event.Delta.Text, - } - jsonData, err := json.Marshal(deltaEvent) - if err != nil { - slog.Error("failed to marshal content delta event", "error", err, "response_id", sc.responseID) - return "" - } - return prefix + fmt.Sprintf("event: response.output_text.delta\ndata: %s\n\n", jsonData) - } - case "input_json_delta": - if event.Delta.PartialJSON == "" { - return "" - } - state := sc.toolCalls[event.Index] - if state == nil { - return "" - } - if state.PlaceholderObject { - state.Arguments = strings.Builder{} - state.PlaceholderObject = false - } - _, _ = state.Arguments.WriteString(event.Delta.PartialJSON) - return sc.output.WriteEvent("response.function_call_arguments.delta", map[string]any{ - "type": "response.function_call_arguments.delta", - "item_id": state.ItemID, - "output_index": state.OutputIndex, - "delta": event.Delta.PartialJSON, - }) - } - return "" - - case "content_block_stop": - state := sc.toolCalls[event.Index] - return sc.output.CompleteToolCall(state, true) - - case "message_delta": - // Capture usage data for inclusion in response.completed - if mergeAnthropicUsage(&sc.usage, event.Usage) { - sc.hasUsage = true - } - if !sc.output.AssistantReserved() && len(sc.toolCalls) == 0 { - sc.reserveAssistantMessageOutput() - } - return "" - - case "message_stop": - // Will be handled in Read() when we get EOF - return "" - } - - return "" -} diff --git a/internal/providers/anthropic/batch.go b/internal/providers/anthropic/batch.go new file mode 100644 index 000000000..1ccd43c0e --- /dev/null +++ b/internal/providers/anthropic/batch.go @@ -0,0 +1,366 @@ +package anthropic + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "log/slog" + "net/http" + "net/url" + "strconv" + "strings" + "time" + + "gomodel/internal/core" + "gomodel/internal/llmclient" +) + +func parseOptionalUnix(ts string) *int64 { + ts = strings.TrimSpace(ts) + if ts == "" { + return nil + } + t, err := time.Parse(time.RFC3339, ts) + if err != nil { + return nil + } + u := t.Unix() + return &u +} + +func mapAnthropicBatchResponse(resp *anthropicBatchResponse) *core.BatchResponse { + if resp == nil { + return nil + } + + total := resp.RequestCounts.Processing + resp.RequestCounts.Succeeded + resp.RequestCounts.Errored + resp.RequestCounts.Canceled + resp.RequestCounts.Expired + failed := resp.RequestCounts.Errored + resp.RequestCounts.Canceled + resp.RequestCounts.Expired + + status := "in_progress" + switch resp.ProcessingStatus { + case "canceling": + status = "cancelling" + case "ended": + switch { + case resp.RequestCounts.Canceled > 0 && resp.RequestCounts.Succeeded == 0 && resp.RequestCounts.Errored == 0: + status = "cancelled" + case resp.RequestCounts.Errored > 0 && resp.RequestCounts.Succeeded == 0: + status = "failed" + default: + status = "completed" + } + } + + return &core.BatchResponse{ + ID: resp.ID, + Object: "batch", + Status: status, + CreatedAt: parseCreatedAt(resp.CreatedAt), + CompletedAt: parseOptionalUnix(resp.EndedAt), + CancellingAt: parseOptionalUnix(resp.CancelInitiatedAt), + RequestCounts: core.BatchRequestCounts{ + Total: total, + Completed: resp.RequestCounts.Succeeded, + Failed: failed, + }, + } +} + +func buildAnthropicBatchCreateRequest(req *core.BatchRequest) (*anthropicBatchCreateRequest, map[string]string, error) { + const maxAnthropicBatchRequests = 10000 + + if req == nil { + return nil, nil, core.NewInvalidRequestError("request is required for anthropic batch processing", nil) + } + if len(req.Requests) == 0 { + return nil, nil, core.NewInvalidRequestError("requests is required for anthropic batch processing", nil) + } + if len(req.Requests) > maxAnthropicBatchRequests { + return nil, nil, core.NewInvalidRequestError("too many requests for anthropic batch processing", nil) + } + + out := &anthropicBatchCreateRequest{ + Requests: make([]anthropicBatchRequest, 0, len(req.Requests)), + } + endpointByCustomID := make(map[string]string, len(req.Requests)) + seenCustomIDs := make(map[string]int, len(req.Requests)) + + for i, item := range req.Requests { + decoded, err := core.DecodeKnownBatchItemRequest(req.Endpoint, item) + if err != nil { + return nil, nil, core.NewInvalidRequestError(fmt.Sprintf("batch item %d: %s", i, err.Error()), err) + } + + params, err := convertDecodedBatchItemToAnthropic(decoded) + if err != nil { + return nil, nil, prefixAnthropicBatchItemError(i, err) + } + + customID := strings.TrimSpace(item.CustomID) + if customID == "" { + customID = fmt.Sprintf("req-%d", i) + } + if previousIndex, exists := seenCustomIDs[customID]; exists { + return nil, nil, core.NewInvalidRequestError( + fmt.Sprintf("batch item %d: duplicate custom_id %q (already used by batch item %d)", i, customID, previousIndex), + nil, + ) + } + seenCustomIDs[customID] = i + out.Requests = append(out.Requests, anthropicBatchRequest{ + CustomID: customID, + Params: *params, + }) + endpointByCustomID[customID] = decoded.Endpoint + } + + return out, endpointByCustomID, nil +} + +func (p *Provider) createBatch(ctx context.Context, req *core.BatchRequest) (*core.BatchResponse, map[string]string, error) { + anthropicReq, endpointByCustomID, err := buildAnthropicBatchCreateRequest(req) + if err != nil { + return nil, nil, err + } + + var resp anthropicBatchResponse + err = p.client.Do(ctx, llmclient.Request{ + Method: http.MethodPost, + Endpoint: "/messages/batches", + Body: anthropicReq, + }, &resp) + if err != nil { + return nil, nil, err + } + + mapped := mapAnthropicBatchResponse(&resp) + if mapped == nil { + return nil, nil, core.NewProviderError("anthropic", http.StatusBadGateway, "failed to map anthropic batch response", nil) + } + mapped.ProviderBatchID = mapped.ID + p.setBatchResultEndpoints(mapped.ProviderBatchID, endpointByCustomID) + return mapped, cloneBatchResultEndpoints(endpointByCustomID), nil +} + +// CreateBatch creates an Anthropic native message batch. +func (p *Provider) CreateBatch(ctx context.Context, req *core.BatchRequest) (*core.BatchResponse, error) { + mapped, _, err := p.createBatch(ctx, req) + return mapped, err +} + +// CreateBatchWithHints creates an Anthropic native message batch and returns +// persisted per-item endpoint hints for later result shaping. +func (p *Provider) CreateBatchWithHints(ctx context.Context, req *core.BatchRequest) (*core.BatchResponse, map[string]string, error) { + return p.createBatch(ctx, req) +} + +// GetBatch retrieves an Anthropic native message batch. +func (p *Provider) GetBatch(ctx context.Context, id string) (*core.BatchResponse, error) { + var resp anthropicBatchResponse + err := p.client.Do(ctx, llmclient.Request{ + Method: http.MethodGet, + Endpoint: "/messages/batches/" + url.PathEscape(id), + }, &resp) + if err != nil { + return nil, err + } + mapped := mapAnthropicBatchResponse(&resp) + if mapped == nil { + return nil, core.NewProviderError("anthropic", http.StatusBadGateway, "failed to map anthropic batch response", nil) + } + mapped.ProviderBatchID = mapped.ID + return mapped, nil +} + +// ListBatches lists Anthropic native message batches. +func (p *Provider) ListBatches(ctx context.Context, limit int, after string) (*core.BatchListResponse, error) { + values := url.Values{} + if limit > 0 { + values.Set("limit", strconv.Itoa(limit)) + } + // Anthropic uses before_id for reverse-chronological pagination. + // Gateway `after` is mapped directly to before_id for provider-native paging. + if after != "" { + values.Set("before_id", after) + } + endpoint := "/messages/batches" + if encoded := values.Encode(); encoded != "" { + endpoint += "?" + encoded + } + + var resp anthropicBatchListResponse + err := p.client.Do(ctx, llmclient.Request{ + Method: http.MethodGet, + Endpoint: endpoint, + }, &resp) + if err != nil { + return nil, err + } + + data := make([]core.BatchResponse, 0, len(resp.Data)) + for _, row := range resp.Data { + mapped := mapAnthropicBatchResponse(&row) + if mapped == nil { + continue + } + mapped.ProviderBatchID = mapped.ID + data = append(data, *mapped) + } + + return &core.BatchListResponse{ + Object: "list", + Data: data, + HasMore: resp.HasMore, + FirstID: resp.FirstID, + LastID: resp.LastID, + }, nil +} + +// CancelBatch cancels an Anthropic native message batch. +func (p *Provider) CancelBatch(ctx context.Context, id string) (*core.BatchResponse, error) { + var resp anthropicBatchResponse + err := p.client.Do(ctx, llmclient.Request{ + Method: http.MethodPost, + Endpoint: "/messages/batches/" + url.PathEscape(id) + "/cancel", + }, &resp) + if err != nil { + return nil, err + } + mapped := mapAnthropicBatchResponse(&resp) + if mapped == nil { + return nil, core.NewProviderError("anthropic", http.StatusBadGateway, "failed to map anthropic batch response", nil) + } + mapped.ProviderBatchID = mapped.ID + return mapped, nil +} + +func (p *Provider) getBatchResults(ctx context.Context, id string, endpointByCustomID map[string]string) (*core.BatchResultsResponse, error) { + resp, err := p.client.DoPassthrough(ctx, llmclient.Request{ + Method: http.MethodGet, + Endpoint: "/messages/batches/" + url.PathEscape(id) + "/results", + }) + if err != nil { + return nil, err + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode != http.StatusOK { + body, readErr := io.ReadAll(resp.Body) + if readErr != nil { + body = []byte("failed to read error response") + } + return nil, core.ParseProviderError("anthropic", resp.StatusCode, body, nil) + } + + scanner := bufio.NewScanner(resp.Body) + // Allow larger result lines than Scanner's default 64K. + scanner.Buffer(make([]byte, 0, 64*1024), 4*1024*1024) + if endpointByCustomID == nil { + endpointByCustomID = p.getBatchResultEndpoints(id) + } else { + endpointByCustomID = cloneBatchResultEndpoints(endpointByCustomID) + } + + results := make([]core.BatchResultItem, 0) + index := 0 + for scanner.Scan() { + line := bytes.TrimSpace(scanner.Bytes()) + if len(line) == 0 { + continue + } + + var row anthropicBatchResultLine + if err := json.Unmarshal(line, &row); err != nil { + slog.Warn( + "failed to decode anthropic batch result line", + "error", err, + "batch_id", id, + "line_index", index, + "line_bytes", len(line), + ) + continue + } + itemEndpoint := "/v1/chat/completions" + if endpointByCustomID != nil { + if endpoint := strings.TrimSpace(endpointByCustomID[row.CustomID]); endpoint != "" { + itemEndpoint = endpoint + } + } + + item := core.BatchResultItem{ + Index: index, + CustomID: row.CustomID, + URL: itemEndpoint, + Provider: "anthropic", + } + switch row.Result.Type { + case "succeeded": + item.StatusCode = http.StatusOK + if len(row.Result.Message) > 0 { + var anthropicPayload anthropicResponse + if err := json.Unmarshal(row.Result.Message, &anthropicPayload); err == nil { + switch itemEndpoint { + case "/v1/responses": + mapped := convertAnthropicResponseToResponses(&anthropicPayload, anthropicPayload.Model) + item.Response = mapped + item.Model = mapped.Model + default: + mapped := convertFromAnthropicResponse(&anthropicPayload) + item.Response = mapped + item.Model = mapped.Model + } + } else { + item.Response = string(row.Result.Message) + } + } + default: + item.StatusCode = http.StatusBadRequest + errType := row.Result.Type + errMsg := "batch item failed" + if row.Result.Error != nil { + if row.Result.Error.Type != "" { + errType = row.Result.Error.Type + } + if row.Result.Error.Message != "" { + errMsg = row.Result.Error.Message + } + } + item.Error = &core.BatchError{ + Type: errType, + Message: errMsg, + } + } + + results = append(results, item) + index++ + } + if err := scanner.Err(); err != nil { + return nil, core.NewProviderError("anthropic", http.StatusBadGateway, "failed to parse anthropic batch results", err) + } + + return &core.BatchResultsResponse{ + Object: "list", + BatchID: id, + Data: results, + }, nil +} + +// GetBatchResults retrieves Anthropic native message batch results. +func (p *Provider) GetBatchResults(ctx context.Context, id string) (*core.BatchResultsResponse, error) { + return p.getBatchResults(ctx, id, nil) +} + +// GetBatchResultsWithHints retrieves Anthropic native batch results using +// persisted per-item endpoint hints instead of transient in-memory state. +func (p *Provider) GetBatchResultsWithHints(ctx context.Context, id string, endpointByCustomID map[string]string) (*core.BatchResultsResponse, error) { + return p.getBatchResults(ctx, id, endpointByCustomID) +} + +// ClearBatchResultHints clears transient per-batch endpoint hints once they +// have been persisted by the gateway. +func (p *Provider) ClearBatchResultHints(batchID string) { + p.clearBatchResultEndpoints(batchID) +} diff --git a/internal/providers/anthropic/chat.go b/internal/providers/anthropic/chat.go new file mode 100644 index 000000000..16a78e347 --- /dev/null +++ b/internal/providers/anthropic/chat.go @@ -0,0 +1,85 @@ +package anthropic + +import ( + "context" + "encoding/json" + "net/http" + "time" + + "gomodel/internal/core" + "gomodel/internal/llmclient" +) + +// convertFromAnthropicResponse converts Anthropic response to core.ChatResponse +func convertFromAnthropicResponse(resp *anthropicResponse) *core.ChatResponse { + content := extractTextContent(resp.Content) + thinking := extractThinkingContent(resp.Content) + toolCalls := extractToolCalls(resp.Content) + + finishReason := normalizeAnthropicStopReason(resp.StopReason) + if finishReason == "" { + finishReason = "stop" + } + + usage := core.Usage{ + PromptTokens: resp.Usage.InputTokens, + CompletionTokens: resp.Usage.OutputTokens, + TotalTokens: resp.Usage.InputTokens + resp.Usage.OutputTokens, + } + + rawUsage := buildAnthropicRawUsage(resp.Usage) + if len(rawUsage) > 0 { + usage.RawUsage = rawUsage + } + + msg := core.ResponseMessage{ + Role: "assistant", + Content: content, + ToolCalls: toolCalls, + } + + // Surface thinking content as reasoning_content (OpenAI-compatible format). + if thinking != "" { + raw, err := json.Marshal(thinking) + if err == nil { + msg.ExtraFields = core.UnknownJSONFieldsFromMap(map[string]json.RawMessage{ + "reasoning_content": raw, + }) + } + } + + return &core.ChatResponse{ + ID: resp.ID, + Object: "chat.completion", + Model: resp.Model, + Created: time.Now().Unix(), + Choices: []core.Choice{ + { + Index: 0, + Message: msg, + FinishReason: finishReason, + }, + }, + Usage: usage, + } +} + +// ChatCompletion sends a chat completion request to Anthropic +func (p *Provider) ChatCompletion(ctx context.Context, req *core.ChatRequest) (*core.ChatResponse, error) { + anthropicReq, err := convertToAnthropicRequest(req) + if err != nil { + return nil, err + } + + var anthropicResp anthropicResponse + err = p.client.Do(ctx, llmclient.Request{ + Method: http.MethodPost, + Endpoint: "/messages", + Body: anthropicReq, + }, &anthropicResp) + if err != nil { + return nil, err + } + + return convertFromAnthropicResponse(&anthropicResp), nil +} diff --git a/internal/providers/anthropic/chat_stream.go b/internal/providers/anthropic/chat_stream.go new file mode 100644 index 000000000..a3883806a --- /dev/null +++ b/internal/providers/anthropic/chat_stream.go @@ -0,0 +1,362 @@ +package anthropic + +import ( + "bufio" + "context" + "encoding/json" + "fmt" + "io" + "log/slog" + "net/http" + "strings" + "time" + + "gomodel/internal/core" + "gomodel/internal/llmclient" + "gomodel/internal/streaming" +) + +// StreamChatCompletion returns a raw response body for streaming (caller must close) +func (p *Provider) StreamChatCompletion(ctx context.Context, req *core.ChatRequest) (io.ReadCloser, error) { + anthropicReq, err := convertToAnthropicRequest(req) + if err != nil { + return nil, err + } + anthropicReq.Stream = true + + stream, err := p.client.DoStream(ctx, llmclient.Request{ + Method: http.MethodPost, + Endpoint: "/messages", + Body: anthropicReq, + }) + if err != nil { + return nil, err + } + + // Return a reader that converts Anthropic SSE format to OpenAI format + return newStreamConverter(stream, req.Model), nil +} + +// streamConverter wraps an Anthropic stream and converts it to OpenAI format +type streamConverter struct { + reader *bufio.Reader + body io.ReadCloser + model string + msgID string + nextToolCallIndex int + toolCalls map[int]*streamToolCallState + thinkingBlocks map[int]bool // tracks which content block indices are thinking blocks + usage anthropicUsage + hasUsage bool + buffer streaming.StreamBuffer + closed bool + emittedToolCalls bool +} + +type streamToolCallState struct { + ID string + Name string + Arguments strings.Builder + Index int + Started bool + PlaceholderObject bool +} + +func newStreamConverter(body io.ReadCloser, model string) *streamConverter { + return &streamConverter{ + reader: bufio.NewReader(body), + body: body, + model: model, + toolCalls: make(map[int]*streamToolCallState), + thinkingBlocks: make(map[int]bool), + buffer: streaming.NewStreamBuffer(1024), + } +} + +func anthropicChatUsagePayload(usage *anthropicUsage) map[string]any { + if usage == nil { + return nil + } + + payload := map[string]any{ + "prompt_tokens": usage.InputTokens, + "completion_tokens": usage.OutputTokens, + "total_tokens": usage.InputTokens + usage.OutputTokens, + } + if usage.CacheReadInputTokens > 0 { + payload["cache_read_input_tokens"] = usage.CacheReadInputTokens + } + if usage.CacheCreationInputTokens > 0 { + payload["cache_creation_input_tokens"] = usage.CacheCreationInputTokens + } + return payload +} + +func (sc *streamConverter) Read(p []byte) (n int, err error) { + // If we have buffered data, return it first + if sc.buffer.Len() > 0 { + return sc.buffer.Read(p), nil + } + + if sc.closed { + sc.releaseBuffer() + return 0, io.EOF + } + + // Read the next SSE event from Anthropic + for { + line, err := sc.reader.ReadBytes('\n') + if err != nil { + if err == io.EOF { + // Send final [DONE] message + sc.buffer.AppendString("data: [DONE]\n\n") + n = sc.buffer.Read(p) + sc.closed = true + _ = sc.body.Close() //nolint:errcheck + return n, nil + } + return 0, err + } + + n, handled, err := consumeAnthropicSSELine(p, line, sc.body, &sc.buffer, sc.convertEvent) + if err != nil { + sc.closed = true + sc.releaseBuffer() + return 0, err + } + if handled { + if n == 0 { + continue + } + return n, nil + } + } +} + +func (sc *streamConverter) Close() error { + if sc.closed { + sc.releaseBuffer() + return nil + } + sc.closed = true + sc.releaseBuffer() + return sc.body.Close() +} + +func (sc *streamConverter) releaseBuffer() { + sc.buffer.Release() +} + +func (sc *streamConverter) mapStreamStopReason(reason string) string { + // Preserve raw "tool_use" when the upstream stream never produced any + // tool call deltas. This avoids claiming OpenAI-style tool calls for a + // malformed or partial Anthropic stream. + if reason == "tool_use" && !sc.emittedToolCalls { + return reason + } + return normalizeAnthropicStopReason(reason) +} + +func (sc *streamConverter) formatChatChunk(delta map[string]any, finishReason any, usage *anthropicUsage) string { + chunk := map[string]any{ + "id": sc.msgID, + "object": "chat.completion.chunk", + "created": time.Now().Unix(), + "model": sc.model, + "provider": "anthropic", + "choices": []map[string]any{ + { + "index": 0, + "delta": delta, + "finish_reason": finishReason, + }, + }, + } + if usage != nil { + chunk["usage"] = anthropicChatUsagePayload(usage) + } + + jsonData, err := json.Marshal(chunk) + if err != nil { + slog.Error("failed to marshal chat completion chunk", "error", err, "msg_id", sc.msgID) + return "" + } + + return fmt.Sprintf("data: %s\n\n", jsonData) +} + +func (sc *streamConverter) convertEvent(event *anthropicStreamEvent) string { + switch event.Type { + case "message_start": + role := "" + if event.Message != nil { + sc.msgID = event.Message.ID + if mergeAnthropicUsage(&sc.usage, &event.Message.Usage) { + sc.hasUsage = true + } + role = strings.TrimSpace(event.Message.Role) + } + if mergeAnthropicUsage(&sc.usage, event.Usage) { + sc.hasUsage = true + } + if event.Message != nil { + if role == "" { + role = "assistant" + } + return sc.formatChatChunk(map[string]any{ + "role": role, + }, nil, nil) + } + return "" + + case "content_block_start": + if event.ContentBlock != nil && event.ContentBlock.Type == "thinking" { + sc.thinkingBlocks[event.Index] = true + return "" + } + if event.ContentBlock != nil && event.ContentBlock.Type == "tool_use" { + state := &streamToolCallState{ + ID: event.ContentBlock.ID, + Name: event.ContentBlock.Name, + Index: sc.nextToolCallIndex, + } + sc.nextToolCallIndex++ + + initialArguments := extractInitialToolArguments(event.ContentBlock.Input) + state.PlaceholderObject = initialArguments == "{}" + if state.PlaceholderObject { + sc.toolCalls[event.Index] = state + return "" + } + if initialArguments != "" { + _, _ = state.Arguments.WriteString(initialArguments) + } + state.Started = true + sc.toolCalls[event.Index] = state + sc.emittedToolCalls = true + + return sc.formatChatChunk(map[string]any{ + "tool_calls": []map[string]any{ + { + "index": state.Index, + "id": state.ID, + "type": "function", + "function": map[string]any{ + "name": state.Name, + "arguments": initialArguments, + }, + }, + }, + }, nil, nil) + } + return "" + + case "content_block_delta": + if event.Delta == nil { + return "" + } + + switch event.Delta.Type { + case "thinking_delta": + if sc.thinkingBlocks[event.Index] && event.Delta.Thinking != "" { + return sc.formatChatChunk(map[string]any{ + "reasoning_content": event.Delta.Thinking, + }, nil, nil) + } + case "signature_delta": + // Signature deltas are internal to Anthropic's thinking protocol; + // no OpenAI-compatible equivalent to emit. + return "" + case "text_delta": + if event.Delta.Text != "" { + return sc.formatChatChunk(map[string]any{ + "content": event.Delta.Text, + }, nil, nil) + } + case "input_json_delta": + if event.Delta.PartialJSON == "" { + return "" + } + state := sc.toolCalls[event.Index] + if state == nil { + return "" + } + if state.PlaceholderObject { + state.Arguments = strings.Builder{} + state.PlaceholderObject = false + } + _, _ = state.Arguments.WriteString(event.Delta.PartialJSON) + if !state.Started { + state.Started = true + sc.emittedToolCalls = true + return sc.formatChatChunk(map[string]any{ + "tool_calls": []map[string]any{ + { + "index": state.Index, + "id": state.ID, + "type": "function", + "function": map[string]any{ + "name": state.Name, + "arguments": event.Delta.PartialJSON, + }, + }, + }, + }, nil, nil) + } + sc.emittedToolCalls = true + return sc.formatChatChunk(map[string]any{ + "tool_calls": []map[string]any{ + { + "index": state.Index, + "function": map[string]any{ + "arguments": event.Delta.PartialJSON, + }, + }, + }, + }, nil, nil) + } + + case "content_block_stop": + state := sc.toolCalls[event.Index] + if state != nil && !state.Started && state.PlaceholderObject { + state.Started = true + sc.emittedToolCalls = true + return sc.formatChatChunk(map[string]any{ + "tool_calls": []map[string]any{ + { + "index": state.Index, + "id": state.ID, + "type": "function", + "function": map[string]any{ + "name": state.Name, + "arguments": "{}", + }, + }, + }, + }, nil, nil) + } + return "" + + case "message_delta": + if mergeAnthropicUsage(&sc.usage, event.Usage) { + sc.hasUsage = true + } + // Emit chunk if we have stop_reason or usage data + if (event.Delta != nil && event.Delta.StopReason != "") || event.Usage != nil { + var finishReason any + if event.Delta != nil && event.Delta.StopReason != "" { + finishReason = sc.mapStreamStopReason(event.Delta.StopReason) + } + var usage *anthropicUsage + if sc.hasUsage { + usage = &sc.usage + } + return sc.formatChatChunk(map[string]any{}, finishReason, usage) + } + + case "message_stop": + return "" + } + + return "" +} diff --git a/internal/providers/anthropic/responses.go b/internal/providers/anthropic/responses.go new file mode 100644 index 000000000..5557bbfa1 --- /dev/null +++ b/internal/providers/anthropic/responses.go @@ -0,0 +1,382 @@ +package anthropic + +import ( + "bufio" + "context" + "encoding/json" + "fmt" + "io" + "log/slog" + "net/http" + "strings" + "time" + + "github.com/google/uuid" + + "gomodel/internal/core" + "gomodel/internal/llmclient" + "gomodel/internal/providers" + "gomodel/internal/streaming" +) + +// convertAnthropicResponseToResponses converts an Anthropic response to ResponsesResponse +func convertAnthropicResponseToResponses(resp *anthropicResponse, model string) *core.ResponsesResponse { + content := extractTextContent(resp.Content) + toolCalls := extractToolCalls(resp.Content) + + msg := core.Message{ + Content: content, + ToolCalls: toolCalls, + } + output := providers.BuildResponsesOutputItems(core.ResponseMessage{ + Role: "assistant", + Content: msg.Content, + ToolCalls: msg.ToolCalls, + }) + + return &core.ResponsesResponse{ + ID: resp.ID, + Object: "response", + CreatedAt: time.Now().Unix(), + Model: model, + Status: "completed", + Output: output, + Usage: buildAnthropicResponsesUsage(resp.Usage), + } +} + +// buildAnthropicResponsesUsage creates a ResponsesUsage from anthropicUsage, including RawUsage. +func buildAnthropicResponsesUsage(u anthropicUsage) *core.ResponsesUsage { + usage := &core.ResponsesUsage{ + InputTokens: u.InputTokens, + OutputTokens: u.OutputTokens, + TotalTokens: u.InputTokens + u.OutputTokens, + } + rawUsage := buildAnthropicRawUsage(u) + if len(rawUsage) > 0 { + usage.RawUsage = rawUsage + } + return usage +} + +func anthropicResponsesUsagePayload(usage *anthropicUsage) map[string]any { + if usage == nil { + return nil + } + + payload := map[string]any{ + "input_tokens": usage.InputTokens, + "output_tokens": usage.OutputTokens, + "total_tokens": usage.InputTokens + usage.OutputTokens, + } + if usage.CacheReadInputTokens > 0 { + payload["cache_read_input_tokens"] = usage.CacheReadInputTokens + } + if usage.CacheCreationInputTokens > 0 { + payload["cache_creation_input_tokens"] = usage.CacheCreationInputTokens + } + return payload +} + +// Responses sends a Responses API request to Anthropic (converted to messages format) +func (p *Provider) Responses(ctx context.Context, req *core.ResponsesRequest) (*core.ResponsesResponse, error) { + anthropicReq, err := convertResponsesRequestToAnthropic(req) + if err != nil { + return nil, err + } + + var anthropicResp anthropicResponse + err = p.client.Do(ctx, llmclient.Request{ + Method: http.MethodPost, + Endpoint: "/messages", + Body: anthropicReq, + }, &anthropicResp) + if err != nil { + return nil, err + } + + return convertAnthropicResponseToResponses(&anthropicResp, req.Model), nil +} + +// StreamResponses returns a raw response body for streaming Responses API (caller must close) +func (p *Provider) StreamResponses(ctx context.Context, req *core.ResponsesRequest) (io.ReadCloser, error) { + anthropicReq, err := convertResponsesRequestToAnthropic(req) + if err != nil { + return nil, err + } + anthropicReq.Stream = true + + stream, err := p.client.DoStream(ctx, llmclient.Request{ + Method: http.MethodPost, + Endpoint: "/messages", + Body: anthropicReq, + }) + if err != nil { + return nil, err + } + + // Return a reader that converts Anthropic SSE format to Responses API format + return newResponsesStreamConverter(stream, req.Model), nil +} + +// responsesStreamConverter wraps an Anthropic stream and converts it to Responses API format +type responsesStreamConverter struct { + reader *bufio.Reader + body io.ReadCloser + model string + responseID string + output *providers.ResponsesOutputEventState + nextOutputIndex int + toolCalls map[int]*providers.ResponsesOutputToolCallState + thinkingBlocks map[int]bool // tracks which content block indices are thinking blocks + buffer streaming.StreamBuffer + closed bool + sentDone bool + usage anthropicUsage + hasUsage bool +} + +func newResponsesStreamConverter(body io.ReadCloser, model string) *responsesStreamConverter { + responseID := "resp_" + uuid.New().String() + return &responsesStreamConverter{ + reader: bufio.NewReader(body), + body: body, + model: model, + responseID: responseID, + output: providers.NewResponsesOutputEventState(responseID), + toolCalls: make(map[int]*providers.ResponsesOutputToolCallState), + thinkingBlocks: make(map[int]bool), + buffer: streaming.NewStreamBuffer(1024), + } +} + +func (sc *responsesStreamConverter) Read(p []byte) (n int, err error) { + if sc.closed { + sc.releaseBuffer() + return 0, io.EOF + } + + // If we have buffered data, return it first + if sc.buffer.Len() > 0 { + return sc.buffer.Read(p), nil + } + + // Read the next SSE event from Anthropic + for { + line, err := sc.reader.ReadBytes('\n') + if err != nil { + if err == io.EOF { + // Send final done event and [DONE] message + if !sc.sentDone { + sc.sentDone = true + prefix := sc.output.CompleteAssistantOutput(0) + responseData := map[string]any{ + "id": sc.responseID, + "object": "response", + "status": "completed", + "model": sc.model, + "provider": "anthropic", + "created_at": time.Now().Unix(), + } + // Include merged usage data captured across message_start/message_delta. + if sc.hasUsage { + responseData["usage"] = anthropicResponsesUsagePayload(&sc.usage) + } + doneEvent := map[string]any{ + "type": "response.completed", + "response": responseData, + } + jsonData, marshalErr := json.Marshal(doneEvent) + if marshalErr != nil { + slog.Error("failed to marshal response.completed event", "error", marshalErr, "response_id", sc.responseID) + sc.closed = true + sc.releaseBuffer() + _ = sc.body.Close() //nolint:errcheck + return 0, io.EOF + } + sc.buffer.AppendString(prefix) + sc.buffer.AppendString("event: response.completed\ndata: ") + sc.buffer.AppendBytes(jsonData) + sc.buffer.AppendString("\n\ndata: [DONE]\n\n") + return sc.buffer.Read(p), nil + } + sc.closed = true + sc.releaseBuffer() + _ = sc.body.Close() //nolint:errcheck + return 0, io.EOF + } + return 0, err + } + + n, handled, err := consumeAnthropicSSELine(p, line, sc.body, &sc.buffer, sc.convertEvent) + if err != nil { + sc.closed = true + sc.releaseBuffer() + return 0, err + } + if handled { + if n == 0 { + continue + } + return n, nil + } + } +} + +func (sc *responsesStreamConverter) Close() error { + if sc.closed { + sc.releaseBuffer() + return nil + } + sc.closed = true + sc.releaseBuffer() + return sc.body.Close() +} + +func (sc *responsesStreamConverter) releaseBuffer() { + sc.buffer.Release() +} + +func (sc *responsesStreamConverter) reserveAssistantMessageOutput() { + if sc.output.AssistantReserved() { + return + } + sc.output.ReserveAssistant() + sc.nextOutputIndex++ +} + +func (sc *responsesStreamConverter) newResponsesToolCallState(contentBlock *anthropicContent) *providers.ResponsesOutputToolCallState { + callID := providers.ResponsesFunctionCallCallID(contentBlock.ID) + state := &providers.ResponsesOutputToolCallState{ + CallID: callID, + Name: contentBlock.Name, + OutputIndex: sc.nextOutputIndex, + } + sc.nextOutputIndex++ + + initialArguments := extractInitialToolArguments(contentBlock.Input) + state.PlaceholderObject = initialArguments == "{}" + if initialArguments != "" && !state.PlaceholderObject { + _, _ = state.Arguments.WriteString(initialArguments) + } + + return state +} + +func (sc *responsesStreamConverter) convertEvent(event *anthropicStreamEvent) string { + switch event.Type { + case "message_start": + if event.Message != nil { + if mergeAnthropicUsage(&sc.usage, &event.Message.Usage) { + sc.hasUsage = true + } + } + if mergeAnthropicUsage(&sc.usage, event.Usage) { + sc.hasUsage = true + } + // Send response.created event + createdEvent := map[string]any{ + "type": "response.created", + "response": map[string]any{ + "id": sc.responseID, + "object": "response", + "status": "in_progress", + "model": sc.model, + "provider": "anthropic", + "created_at": time.Now().Unix(), + }, + } + jsonData, err := json.Marshal(createdEvent) + if err != nil { + slog.Error("failed to marshal response.created event", "error", err, "response_id", sc.responseID) + return "" + } + return fmt.Sprintf("event: response.created\ndata: %s\n\n", jsonData) + + case "content_block_start": + if event.ContentBlock != nil && event.ContentBlock.Type == "thinking" { + sc.thinkingBlocks[event.Index] = true + return "" + } + if event.ContentBlock != nil && event.ContentBlock.Type == "tool_use" { + if sc.output.AssistantStarted() && !sc.output.AssistantDone() { + prefix := sc.output.CompleteAssistantOutput(0) + state := sc.newResponsesToolCallState(event.ContentBlock) + sc.toolCalls[event.Index] = state + return prefix + sc.output.StartToolCall(state, true) + } + state := sc.newResponsesToolCallState(event.ContentBlock) + sc.toolCalls[event.Index] = state + return sc.output.StartToolCall(state, true) + } + return "" + + case "content_block_delta": + if event.Delta == nil { + return "" + } + + switch event.Delta.Type { + case "thinking_delta", "signature_delta": + // Thinking and signature deltas are part of Anthropic's extended thinking; + // the Responses API format does not have a direct equivalent, so skip them. + return "" + case "text_delta": + if event.Delta.Text != "" { + sc.reserveAssistantMessageOutput() + prefix := sc.output.StartAssistantOutput(0) + sc.output.AppendAssistantText(event.Delta.Text) + deltaEvent := map[string]any{ + "type": "response.output_text.delta", + "delta": event.Delta.Text, + } + jsonData, err := json.Marshal(deltaEvent) + if err != nil { + slog.Error("failed to marshal content delta event", "error", err, "response_id", sc.responseID) + return "" + } + return prefix + fmt.Sprintf("event: response.output_text.delta\ndata: %s\n\n", jsonData) + } + case "input_json_delta": + if event.Delta.PartialJSON == "" { + return "" + } + state := sc.toolCalls[event.Index] + if state == nil { + return "" + } + if state.PlaceholderObject { + state.Arguments = strings.Builder{} + state.PlaceholderObject = false + } + _, _ = state.Arguments.WriteString(event.Delta.PartialJSON) + return sc.output.WriteEvent("response.function_call_arguments.delta", map[string]any{ + "type": "response.function_call_arguments.delta", + "item_id": state.ItemID, + "output_index": state.OutputIndex, + "delta": event.Delta.PartialJSON, + }) + } + return "" + + case "content_block_stop": + state := sc.toolCalls[event.Index] + return sc.output.CompleteToolCall(state, true) + + case "message_delta": + // Capture usage data for inclusion in response.completed + if mergeAnthropicUsage(&sc.usage, event.Usage) { + sc.hasUsage = true + } + if !sc.output.AssistantReserved() && len(sc.toolCalls) == 0 { + sc.reserveAssistantMessageOutput() + } + return "" + + case "message_stop": + // Will be handled in Read() when we get EOF + return "" + } + + return "" +} From 630f8b9b1989f74f539eb88dbcced9028130f4d3 Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Wed, 29 Apr 2026 23:14:34 +0200 Subject: [PATCH 03/15] refactor(providers): split responses adapter into per-direction files MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit internal/providers/responses_adapter.go was 1005 lines covering four distinct directions of Responses↔Chat translation. Split into peer files in the same package: responses_adapter.go (179) ChatProvider iface, request->chat, tool/tool-choice normalization, ResponsesViaChat / StreamResponsesViaChat responses_input.go (297) Input items -> Chat messages, assistant-message merging, stringify/firstNonEmpty helpers responses_content.go (367) Content normalization, image/audio part conversion, ExtractContentFromInput, extras-extraction helpers responses_output.go (185) Output items + ConvertChatResponseToResponses, ResponsesFunctionCall{CallID,ItemID} Pure relocation; tests pass. Co-Authored-By: Claude Opus 4.7 --- internal/providers/responses_adapter.go | 826 ------------------------ internal/providers/responses_content.go | 367 +++++++++++ internal/providers/responses_input.go | 297 +++++++++ internal/providers/responses_output.go | 185 ++++++ 4 files changed, 849 insertions(+), 826 deletions(-) create mode 100644 internal/providers/responses_content.go create mode 100644 internal/providers/responses_input.go create mode 100644 internal/providers/responses_output.go diff --git a/internal/providers/responses_adapter.go b/internal/providers/responses_adapter.go index f4e3f1d50..c50faf537 100644 --- a/internal/providers/responses_adapter.go +++ b/internal/providers/responses_adapter.go @@ -2,14 +2,10 @@ package providers import ( "context" - "encoding/json" - "fmt" "io" "maps" "strings" - "github.com/google/uuid" - "gomodel/internal/core" ) @@ -146,828 +142,6 @@ func cloneStringAnyMap(src map[string]any) map[string]any { return dst } -// ConvertResponsesInputToMessages converts a Responses API input payload into Chat API messages. -func ConvertResponsesInputToMessages(input any) ([]core.Message, error) { - switch in := input.(type) { - case string: - return []core.Message{{Role: "user", Content: in}}, nil - case []map[string]any: - items := make([]any, 0, len(in)) - for _, item := range in { - items = append(items, item) - } - return convertResponsesInputItems(items) - case []any: - return convertResponsesInputItems(in) - case []core.ResponsesInputElement: - items := make([]any, 0, len(in)) - for _, item := range in { - items = append(items, item) - } - return convertResponsesInputItems(items) - case nil: - return nil, core.NewInvalidRequestError("invalid responses input: unsupported type", nil) - default: - return nil, core.NewInvalidRequestError("invalid responses input: unsupported type", nil) - } -} - -func convertResponsesInputItems(items []any) ([]core.Message, error) { - messages := make([]core.Message, 0, len(items)) - var pendingAssistant *core.Message - - flushPendingAssistant := func() { - if pendingAssistant == nil { - return - } - messages = append(messages, *pendingAssistant) - pendingAssistant = nil - } - - for i, item := range items { - msg, itemType, err := convertResponsesInputItem(item, i) - if err != nil { - return nil, err - } - - if msg.Role == "assistant" { - if itemType == "message" { - flushPendingAssistant() - } - if pendingAssistant == nil { - assistant := cloneResponsesMessage(msg) - pendingAssistant = &assistant - } else if canMergeAssistantMessages(*pendingAssistant, msg) { - mergeAssistantMessage(pendingAssistant, msg) - } else { - flushPendingAssistant() - assistant := cloneResponsesMessage(msg) - pendingAssistant = &assistant - } - continue - } - - flushPendingAssistant() - messages = append(messages, msg) - } - - flushPendingAssistant() - return messages, nil -} - -func convertResponsesInputItem(item any, index int) (core.Message, string, error) { - switch typed := item.(type) { - case core.ResponsesInputElement: - return convertResponsesInputElement(typed, index) - case map[string]any: - return convertResponsesInputMap(typed, index) - default: - return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: expected object", index), nil) - } -} - -func convertResponsesInputElement(item core.ResponsesInputElement, index int) (core.Message, string, error) { - switch item.Type { - case "function_call": - name := strings.TrimSpace(item.Name) - if name == "" { - return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: function_call name is required", index), nil) - } - callID := ResponsesFunctionCallCallID(item.CallID) - return core.Message{ - Role: "assistant", - Content: "", - ContentNull: true, - ToolCalls: []core.ToolCall{ - { - ID: callID, - Type: "function", - ExtraFields: core.CloneUnknownJSONFields(item.ExtraFields), - Function: core.FunctionCall{ - Name: name, - Arguments: item.Arguments, - }, - }, - }, - }, "function_call", nil - case "function_call_output": - callID := strings.TrimSpace(item.CallID) - if callID == "" { - return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: function_call_output call_id is required", index), nil) - } - content, err := stringifyResponsesInputValueWithError(item.Output) - if err != nil { - return core.Message{}, "", core.NewInvalidRequestError( - fmt.Sprintf("invalid responses input item at index %d: function_call_output.output must be JSON-serializable", index), - err, - ) - } - return core.Message{ - Role: "tool", - ToolCallID: callID, - Content: content, - ExtraFields: core.CloneUnknownJSONFields(item.ExtraFields), - }, "function_call_output", nil - default: // message (type="" or "message") - role := strings.TrimSpace(item.Role) - if role == "" { - return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: role is required", index), nil) - } - content, ok := ConvertResponsesContentToChatContent(item.Content) - if !ok { - return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: unsupported content", index), nil) - } - return core.Message{ - Role: role, - Content: content, - ExtraFields: core.CloneUnknownJSONFields(item.ExtraFields), - }, "message", nil - } -} - -func convertResponsesInputMap(item map[string]any, index int) (core.Message, string, error) { - itemType, _ := item["type"].(string) - switch itemType { - case "function_call": - name, _ := item["name"].(string) - callID := firstNonEmptyString(item, "call_id", "id") - if strings.TrimSpace(name) == "" { - return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: function_call name is required", index), nil) - } - callID = ResponsesFunctionCallCallID(callID) - return core.Message{ - Role: "assistant", - Content: "", - ContentNull: true, - ToolCalls: []core.ToolCall{ - { - ID: callID, - Type: "function", - ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownKeys(item, "type", "call_id", "id", "name", "arguments", "status")), - Function: core.FunctionCall{ - Name: name, - Arguments: stringifyResponsesInputValue(item["arguments"]), - }, - }, - }, - }, "function_call", nil - case "function_call_output": - callID := firstNonEmptyString(item, "call_id") - if callID == "" { - return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: function_call_output call_id is required", index), nil) - } - content, err := stringifyResponsesInputValueWithError(item["output"]) - if err != nil { - return core.Message{}, "", core.NewInvalidRequestError( - fmt.Sprintf("invalid responses input item at index %d: function_call_output.output must be JSON-serializable", index), - err, - ) - } - return core.Message{ - Role: "tool", - ToolCallID: callID, - Content: content, - ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownKeys(item, "type", "call_id", "status", "output")), - }, "function_call_output", nil - } - - role, _ := item["role"].(string) - role = strings.TrimSpace(role) - if role == "" { - return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: role is required", index), nil) - } - - content, ok := ConvertResponsesContentToChatContent(item["content"]) - if !ok { - return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: unsupported content", index), nil) - } - return core.Message{ - Role: role, - Content: content, - ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownKeys(item, "type", "role", "status", "content")), - }, "message", nil -} - -func cloneResponsesMessage(msg core.Message) core.Message { - cloned := msg - if len(msg.ToolCalls) > 0 { - cloned.ToolCalls = make([]core.ToolCall, len(msg.ToolCalls)) - for i, call := range msg.ToolCalls { - cloned.ToolCalls[i] = cloneResponsesToolCall(call) - } - } - if parts, ok := msg.Content.([]core.ContentPart); ok { - clonedParts := make([]core.ContentPart, len(parts)) - for i, part := range parts { - clonedParts[i] = cloneResponsesContentPart(part) - } - cloned.Content = clonedParts - } - cloned.ExtraFields = core.CloneUnknownJSONFields(msg.ExtraFields) - return cloned -} - -func canMergeAssistantMessages(current, next core.Message) bool { - if !current.ExtraFields.IsEmpty() || !next.ExtraFields.IsEmpty() { - return false - } - if !core.HasStructuredContent(current.Content) && !core.HasStructuredContent(next.Content) { - return true - } - return isAssistantToolCallOnlyMessage(next) -} - -func mergeAssistantMessage(dst *core.Message, src core.Message) { - if text := core.ExtractTextContent(src.Content); text != "" { - existing := core.ExtractTextContent(dst.Content) - dst.Content = existing + text - dst.ContentNull = false - } - if len(src.ToolCalls) > 0 { - dst.ToolCalls = append(dst.ToolCalls, src.ToolCalls...) - if core.ExtractTextContent(dst.Content) == "" { - dst.ContentNull = dst.ContentNull || src.ContentNull - } - } -} - -func isAssistantToolCallOnlyMessage(msg core.Message) bool { - if msg.Role != "assistant" || len(msg.ToolCalls) == 0 { - return false - } - if core.HasStructuredContent(msg.Content) { - return false - } - return core.ExtractTextContent(msg.Content) == "" -} - -// ConvertResponsesContentToChatContent maps Responses input content to Chat content. -// Text-only arrays are flattened to strings for broader provider compatibility. -// Any non-text part preserves the array form so multimodal payloads survive routing. -func ConvertResponsesContentToChatContent(content any) (any, bool) { - switch c := content.(type) { - case string: - return c, true - case []map[string]any: - items := make([]any, 0, len(c)) - for _, item := range c { - items = append(items, item) - } - return convertResponsesContentParts(items) - case []any: - return convertResponsesContentParts(c) - case []core.ContentPart: - parts := make([]core.ContentPart, 0, len(c)) - for _, part := range c { - normalized, ok := normalizeTypedResponsesContentPart(part) - if !ok { - return nil, false - } - parts = append(parts, normalized) - } - return finalizeResponsesChatContent(parts) - case core.ContentPart: - normalized, ok := normalizeTypedResponsesContentPart(c) - if !ok { - return nil, false - } - return finalizeResponsesChatContent([]core.ContentPart{normalized}) - default: - return nil, false - } -} - -func convertResponsesContentParts(parts []any) (any, bool) { - typedParts := make([]core.ContentPart, 0, len(parts)) - - for _, part := range parts { - partMap, ok := part.(map[string]any) - if !ok { - return nil, false - } - - partType, _ := partMap["type"].(string) - switch partType { - case "text", "input_text", "output_text": - text, ok := partMap["text"].(string) - if !ok || text == "" { - return nil, false - } - typedParts = append(typedParts, core.ContentPart{ - Type: "text", - Text: text, - ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownKeys(partMap, "type", "text")), - }) - case "image_url", "input_image": - imageURL, ok := normalizeResponsesImageURLForChat(partMap["image_url"]) - if !ok { - return nil, false - } - typedParts = append(typedParts, core.ContentPart{ - Type: "image_url", - ImageURL: imageURL, - ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownKeys(partMap, "type", "image_url")), - }) - case "input_audio": - inputAudio, ok := normalizeResponsesInputAudioForChat(partMap["input_audio"]) - if !ok { - return nil, false - } - typedParts = append(typedParts, core.ContentPart{ - Type: "input_audio", - InputAudio: inputAudio, - ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownKeys(partMap, "type", "input_audio")), - }) - default: - if nested, ok := partMap["content"]; ok { - text := ExtractContentFromInput(nested) - if text == "" { - return nil, false - } - typedParts = append(typedParts, core.ContentPart{Type: "text", Text: text}) - continue - } - return nil, false - } - } - - if len(typedParts) == 0 { - return nil, false - } - return finalizeResponsesChatContent(typedParts) -} - -func normalizeTypedResponsesContentPart(part core.ContentPart) (core.ContentPart, bool) { - switch part.Type { - case "text", "input_text", "output_text": - if part.Text == "" { - return core.ContentPart{}, false - } - return core.ContentPart{ - Type: "text", - Text: part.Text, - ExtraFields: core.CloneUnknownJSONFields(part.ExtraFields), - }, true - case "image_url", "input_image": - if part.ImageURL == nil { - return core.ContentPart{}, false - } - url := strings.TrimSpace(part.ImageURL.URL) - if url == "" { - return core.ContentPart{}, false - } - return core.ContentPart{ - Type: "image_url", - ImageURL: &core.ImageURLContent{ - URL: url, - Detail: strings.TrimSpace(part.ImageURL.Detail), - MediaType: strings.TrimSpace(part.ImageURL.MediaType), - ExtraFields: core.CloneUnknownJSONFields(part.ImageURL.ExtraFields), - }, - ExtraFields: core.CloneUnknownJSONFields(part.ExtraFields), - }, true - case "input_audio": - if part.InputAudio == nil { - return core.ContentPart{}, false - } - data := strings.TrimSpace(part.InputAudio.Data) - format := strings.TrimSpace(part.InputAudio.Format) - if data == "" || format == "" { - return core.ContentPart{}, false - } - return core.ContentPart{ - Type: "input_audio", - InputAudio: &core.InputAudioContent{ - Data: data, - Format: format, - ExtraFields: core.CloneUnknownJSONFields(part.InputAudio.ExtraFields), - }, - ExtraFields: core.CloneUnknownJSONFields(part.ExtraFields), - }, true - default: - return core.ContentPart{}, false - } -} - -func finalizeResponsesChatContent(parts []core.ContentPart) (any, bool) { - if len(parts) == 0 { - return nil, false - } - - if !canFlattenResponsesPartsToText(parts) { - return parts, true - } - - texts := make([]string, 0, len(parts)) - for _, part := range parts { - texts = append(texts, part.Text) - } - return strings.Join(texts, " "), true -} - -func canFlattenResponsesPartsToText(parts []core.ContentPart) bool { - for _, part := range parts { - if part.Type != "text" { - return false - } - if !part.ExtraFields.IsEmpty() { - return false - } - } - return true -} - -func normalizeResponsesImageURLForChat(value any) (*core.ImageURLContent, bool) { - switch v := value.(type) { - case string: - url := strings.TrimSpace(v) - if url == "" { - return nil, false - } - return &core.ImageURLContent{URL: url}, true - case map[string]string: - url := strings.TrimSpace(v["url"]) - if url == "" { - return nil, false - } - return &core.ImageURLContent{ - URL: url, - Detail: strings.TrimSpace(v["detail"]), - MediaType: strings.TrimSpace(v["media_type"]), - ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownStringKeys(v, "url", "detail", "media_type")), - }, true - case map[string]any: - url, _ := v["url"].(string) - url = strings.TrimSpace(url) - if url == "" { - return nil, false - } - detail, _ := v["detail"].(string) - mediaType, _ := v["media_type"].(string) - return &core.ImageURLContent{ - URL: url, - Detail: strings.TrimSpace(detail), - MediaType: strings.TrimSpace(mediaType), - ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownKeys(v, "url", "detail", "media_type")), - }, true - default: - return nil, false - } -} - -func normalizeResponsesInputAudioForChat(value any) (*core.InputAudioContent, bool) { - switch v := value.(type) { - case map[string]string: - data := strings.TrimSpace(v["data"]) - format := strings.TrimSpace(v["format"]) - if data == "" || format == "" { - return nil, false - } - return &core.InputAudioContent{ - Data: data, - Format: format, - ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownStringKeys(v, "data", "format")), - }, true - case map[string]any: - data, _ := v["data"].(string) - format, _ := v["format"].(string) - data = strings.TrimSpace(data) - format = strings.TrimSpace(format) - if data == "" || format == "" { - return nil, false - } - return &core.InputAudioContent{ - Data: data, - Format: format, - ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownKeys(v, "data", "format")), - }, true - default: - return nil, false - } -} - -func cloneResponsesToolCall(call core.ToolCall) core.ToolCall { - cloned := call - cloned.ExtraFields = core.CloneUnknownJSONFields(call.ExtraFields) - cloned.Function.ExtraFields = core.CloneUnknownJSONFields(call.Function.ExtraFields) - return cloned -} - -func cloneResponsesContentPart(part core.ContentPart) core.ContentPart { - cloned := part - cloned.ExtraFields = core.CloneUnknownJSONFields(part.ExtraFields) - if part.ImageURL != nil { - image := *part.ImageURL - image.ExtraFields = core.CloneUnknownJSONFields(part.ImageURL.ExtraFields) - cloned.ImageURL = &image - } - if part.InputAudio != nil { - audio := *part.InputAudio - audio.ExtraFields = core.CloneUnknownJSONFields(part.InputAudio.ExtraFields) - cloned.InputAudio = &audio - } - return cloned -} - -func rawJSONMapFromUnknownKeys(src map[string]any, knownKeys ...string) map[string]json.RawMessage { - if len(src) == 0 { - return nil - } - known := make(map[string]struct{}, len(knownKeys)) - for _, key := range knownKeys { - known[key] = struct{}{} - } - - var extras map[string]json.RawMessage - for key, value := range src { - if _, ok := known[key]; ok { - continue - } - raw, err := json.Marshal(value) - if err != nil { - continue - } - if extras == nil { - extras = make(map[string]json.RawMessage) - } - extras[key] = raw - } - return extras -} - -func rawJSONMapFromUnknownStringKeys(src map[string]string, knownKeys ...string) map[string]json.RawMessage { - if len(src) == 0 { - return nil - } - - converted := make(map[string]any, len(src)) - for key, value := range src { - converted[key] = value - } - return rawJSONMapFromUnknownKeys(converted, knownKeys...) -} - -func firstNonEmptyString(item map[string]any, keys ...string) string { - for _, key := range keys { - value, _ := item[key].(string) - if strings.TrimSpace(value) != "" { - return value - } - } - return "" -} - -func stringifyResponsesInputValue(value any) string { - encoded, err := stringifyResponsesInputValueWithError(value) - if err != nil { - return "" - } - return encoded -} - -func stringifyResponsesInputValueWithError(value any) (string, error) { - switch v := value.(type) { - case nil: - return "", nil - case string: - return v, nil - default: - encoded, err := json.Marshal(v) - if err != nil { - return "", err - } - return string(encoded), nil - } -} - -// ExtractContentFromInput extracts text content from responses input. -func ExtractContentFromInput(content any) string { - switch c := content.(type) { - case string: - return c - case []core.ContentPart: - texts := make([]string, 0, len(c)) - for _, part := range c { - if part.Text != "" { - texts = append(texts, part.Text) - } - } - return strings.Join(texts, " ") - case []map[string]any: - return extractTextFromMapSlice(c) - case []any: - texts := make([]string, 0, len(c)) - for _, part := range c { - if partMap, ok := part.(map[string]any); ok { - if text := extractTextFromInputMap(partMap); text != "" { - texts = append(texts, text) - } - } - } - return strings.Join(texts, " ") - default: - return "" - } -} - -func extractTextFromMapSlice(parts []map[string]any) string { - texts := make([]string, 0, len(parts)) - for _, part := range parts { - if text := extractTextFromInputMap(part); text != "" { - texts = append(texts, text) - } - } - return strings.Join(texts, " ") -} - -func extractTextFromInputMap(part map[string]any) string { - texts := make([]string, 0, 2) - if text, ok := part["text"].(string); ok && text != "" { - texts = append(texts, text) - } - if nested, ok := part["content"]; ok { - if text := ExtractContentFromInput(nested); text != "" { - texts = append(texts, text) - } - } - return strings.Join(texts, " ") -} - -// ResponsesFunctionCallCallID returns the call id if present or generates one. -func ResponsesFunctionCallCallID(callID string) string { - if strings.TrimSpace(callID) != "" { - return callID - } - return "call_" + uuid.New().String() -} - -// ResponsesFunctionCallItemID returns a stable function-call item id. -func ResponsesFunctionCallItemID(callID string) string { - normalizedCallID := strings.TrimSpace(callID) - if normalizedCallID == "" { - normalizedCallID = "call_" + uuid.New().String() - } - return "fc_" + normalizedCallID -} - -func buildResponsesMessageContent(content any) []core.ResponsesContentItem { - switch c := content.(type) { - case string: - return []core.ResponsesContentItem{ - { - Type: "output_text", - Text: c, - Annotations: []json.RawMessage{}, - }, - } - case []core.ContentPart: - return buildResponsesContentItemsFromParts(c) - case []any: - parts, ok := core.NormalizeContentParts(c) - if !ok { - return nil - } - return buildResponsesContentItemsFromParts(parts) - default: - text := core.ExtractTextContent(content) - if text == "" { - return nil - } - return []core.ResponsesContentItem{ - { - Type: "output_text", - Text: text, - Annotations: []json.RawMessage{}, - }, - } - } -} - -func buildResponsesContentItemsFromParts(parts []core.ContentPart) []core.ResponsesContentItem { - items := make([]core.ResponsesContentItem, 0, len(parts)) - for _, part := range parts { - switch part.Type { - case "text": - items = append(items, core.ResponsesContentItem{ - Type: "output_text", - Text: part.Text, - Annotations: []json.RawMessage{}, - }) - case "image_url": - if part.ImageURL == nil { - continue - } - url := strings.TrimSpace(part.ImageURL.URL) - if url == "" { - continue - } - items = append(items, core.ResponsesContentItem{ - Type: "input_image", - ImageURL: &core.ImageURLContent{ - URL: url, - Detail: strings.TrimSpace(part.ImageURL.Detail), - MediaType: strings.TrimSpace(part.ImageURL.MediaType), - ExtraFields: core.CloneUnknownJSONFields(part.ImageURL.ExtraFields), - }, - }) - case "input_audio": - if part.InputAudio == nil { - continue - } - data := strings.TrimSpace(part.InputAudio.Data) - format := strings.TrimSpace(part.InputAudio.Format) - if data == "" || format == "" { - continue - } - items = append(items, core.ResponsesContentItem{ - Type: "input_audio", - InputAudio: &core.InputAudioContent{ - Data: data, - Format: format, - ExtraFields: core.CloneUnknownJSONFields(part.InputAudio.ExtraFields), - }, - }) - } - } - return items -} - -// BuildResponsesOutputItems converts a response message into Responses API output items. -func BuildResponsesOutputItems(msg core.ResponseMessage) []core.ResponsesOutputItem { - output := make([]core.ResponsesOutputItem, 0, len(msg.ToolCalls)+1) - contentItems := buildResponsesMessageContent(msg.Content) - if len(contentItems) > 0 || len(msg.ToolCalls) == 0 { - if len(contentItems) == 0 { - contentItems = []core.ResponsesContentItem{ - { - Type: "output_text", - Text: "", - Annotations: []json.RawMessage{}, - }, - } - } - output = append(output, core.ResponsesOutputItem{ - ID: "msg_" + uuid.New().String(), - Type: "message", - Role: "assistant", - Status: "completed", - Content: contentItems, - }) - } - for _, toolCall := range msg.ToolCalls { - callID := ResponsesFunctionCallCallID(toolCall.ID) - output = append(output, core.ResponsesOutputItem{ - ID: ResponsesFunctionCallItemID(callID), - Type: "function_call", - Status: "completed", - CallID: callID, - Name: toolCall.Function.Name, - Arguments: toolCall.Function.Arguments, - }) - } - return output -} - -// ConvertChatResponseToResponses converts a ChatResponse to a ResponsesResponse. -func ConvertChatResponseToResponses(resp *core.ChatResponse) *core.ResponsesResponse { - output := []core.ResponsesOutputItem{ - { - ID: "msg_" + uuid.New().String(), - Type: "message", - Role: "assistant", - Status: "completed", - Content: []core.ResponsesContentItem{ - { - Type: "output_text", - Text: "", - Annotations: []json.RawMessage{}, - }, - }, - }, - } - if len(resp.Choices) > 0 { - output = BuildResponsesOutputItems(resp.Choices[0].Message) - } - - return &core.ResponsesResponse{ - ID: resp.ID, - Object: "response", - CreatedAt: resp.Created, - Model: resp.Model, - Provider: resp.Provider, - Status: "completed", - Output: output, - Usage: &core.ResponsesUsage{ - InputTokens: resp.Usage.PromptTokens, - OutputTokens: resp.Usage.CompletionTokens, - TotalTokens: resp.Usage.TotalTokens, - PromptTokensDetails: resp.Usage.PromptTokensDetails, - CompletionTokensDetails: resp.Usage.CompletionTokensDetails, - RawUsage: resp.Usage.RawUsage, - }, - } -} - // ResponsesViaChat implements the Responses API by converting to/from Chat format. func ResponsesViaChat(ctx context.Context, p ChatProvider, req *core.ResponsesRequest) (*core.ResponsesResponse, error) { chatReq, err := ConvertResponsesRequestToChat(req) diff --git a/internal/providers/responses_content.go b/internal/providers/responses_content.go new file mode 100644 index 000000000..71df14fc8 --- /dev/null +++ b/internal/providers/responses_content.go @@ -0,0 +1,367 @@ +package providers + +import ( + "encoding/json" + "strings" + + "gomodel/internal/core" +) + +// ConvertResponsesContentToChatContent maps Responses input content to Chat content. +// Text-only arrays are flattened to strings for broader provider compatibility. +// Any non-text part preserves the array form so multimodal payloads survive routing. +func ConvertResponsesContentToChatContent(content any) (any, bool) { + switch c := content.(type) { + case string: + return c, true + case []map[string]any: + items := make([]any, 0, len(c)) + for _, item := range c { + items = append(items, item) + } + return convertResponsesContentParts(items) + case []any: + return convertResponsesContentParts(c) + case []core.ContentPart: + parts := make([]core.ContentPart, 0, len(c)) + for _, part := range c { + normalized, ok := normalizeTypedResponsesContentPart(part) + if !ok { + return nil, false + } + parts = append(parts, normalized) + } + return finalizeResponsesChatContent(parts) + case core.ContentPart: + normalized, ok := normalizeTypedResponsesContentPart(c) + if !ok { + return nil, false + } + return finalizeResponsesChatContent([]core.ContentPart{normalized}) + default: + return nil, false + } +} + +func convertResponsesContentParts(parts []any) (any, bool) { + typedParts := make([]core.ContentPart, 0, len(parts)) + + for _, part := range parts { + partMap, ok := part.(map[string]any) + if !ok { + return nil, false + } + + partType, _ := partMap["type"].(string) + switch partType { + case "text", "input_text", "output_text": + text, ok := partMap["text"].(string) + if !ok || text == "" { + return nil, false + } + typedParts = append(typedParts, core.ContentPart{ + Type: "text", + Text: text, + ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownKeys(partMap, "type", "text")), + }) + case "image_url", "input_image": + imageURL, ok := normalizeResponsesImageURLForChat(partMap["image_url"]) + if !ok { + return nil, false + } + typedParts = append(typedParts, core.ContentPart{ + Type: "image_url", + ImageURL: imageURL, + ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownKeys(partMap, "type", "image_url")), + }) + case "input_audio": + inputAudio, ok := normalizeResponsesInputAudioForChat(partMap["input_audio"]) + if !ok { + return nil, false + } + typedParts = append(typedParts, core.ContentPart{ + Type: "input_audio", + InputAudio: inputAudio, + ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownKeys(partMap, "type", "input_audio")), + }) + default: + if nested, ok := partMap["content"]; ok { + text := ExtractContentFromInput(nested) + if text == "" { + return nil, false + } + typedParts = append(typedParts, core.ContentPart{Type: "text", Text: text}) + continue + } + return nil, false + } + } + + if len(typedParts) == 0 { + return nil, false + } + return finalizeResponsesChatContent(typedParts) +} + +func normalizeTypedResponsesContentPart(part core.ContentPart) (core.ContentPart, bool) { + switch part.Type { + case "text", "input_text", "output_text": + if part.Text == "" { + return core.ContentPart{}, false + } + return core.ContentPart{ + Type: "text", + Text: part.Text, + ExtraFields: core.CloneUnknownJSONFields(part.ExtraFields), + }, true + case "image_url", "input_image": + if part.ImageURL == nil { + return core.ContentPart{}, false + } + url := strings.TrimSpace(part.ImageURL.URL) + if url == "" { + return core.ContentPart{}, false + } + return core.ContentPart{ + Type: "image_url", + ImageURL: &core.ImageURLContent{ + URL: url, + Detail: strings.TrimSpace(part.ImageURL.Detail), + MediaType: strings.TrimSpace(part.ImageURL.MediaType), + ExtraFields: core.CloneUnknownJSONFields(part.ImageURL.ExtraFields), + }, + ExtraFields: core.CloneUnknownJSONFields(part.ExtraFields), + }, true + case "input_audio": + if part.InputAudio == nil { + return core.ContentPart{}, false + } + data := strings.TrimSpace(part.InputAudio.Data) + format := strings.TrimSpace(part.InputAudio.Format) + if data == "" || format == "" { + return core.ContentPart{}, false + } + return core.ContentPart{ + Type: "input_audio", + InputAudio: &core.InputAudioContent{ + Data: data, + Format: format, + ExtraFields: core.CloneUnknownJSONFields(part.InputAudio.ExtraFields), + }, + ExtraFields: core.CloneUnknownJSONFields(part.ExtraFields), + }, true + default: + return core.ContentPart{}, false + } +} + +func finalizeResponsesChatContent(parts []core.ContentPart) (any, bool) { + if len(parts) == 0 { + return nil, false + } + + if !canFlattenResponsesPartsToText(parts) { + return parts, true + } + + texts := make([]string, 0, len(parts)) + for _, part := range parts { + texts = append(texts, part.Text) + } + return strings.Join(texts, " "), true +} + +func canFlattenResponsesPartsToText(parts []core.ContentPart) bool { + for _, part := range parts { + if part.Type != "text" { + return false + } + if !part.ExtraFields.IsEmpty() { + return false + } + } + return true +} + +func normalizeResponsesImageURLForChat(value any) (*core.ImageURLContent, bool) { + switch v := value.(type) { + case string: + url := strings.TrimSpace(v) + if url == "" { + return nil, false + } + return &core.ImageURLContent{URL: url}, true + case map[string]string: + url := strings.TrimSpace(v["url"]) + if url == "" { + return nil, false + } + return &core.ImageURLContent{ + URL: url, + Detail: strings.TrimSpace(v["detail"]), + MediaType: strings.TrimSpace(v["media_type"]), + ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownStringKeys(v, "url", "detail", "media_type")), + }, true + case map[string]any: + url, _ := v["url"].(string) + url = strings.TrimSpace(url) + if url == "" { + return nil, false + } + detail, _ := v["detail"].(string) + mediaType, _ := v["media_type"].(string) + return &core.ImageURLContent{ + URL: url, + Detail: strings.TrimSpace(detail), + MediaType: strings.TrimSpace(mediaType), + ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownKeys(v, "url", "detail", "media_type")), + }, true + default: + return nil, false + } +} + +func normalizeResponsesInputAudioForChat(value any) (*core.InputAudioContent, bool) { + switch v := value.(type) { + case map[string]string: + data := strings.TrimSpace(v["data"]) + format := strings.TrimSpace(v["format"]) + if data == "" || format == "" { + return nil, false + } + return &core.InputAudioContent{ + Data: data, + Format: format, + ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownStringKeys(v, "data", "format")), + }, true + case map[string]any: + data, _ := v["data"].(string) + format, _ := v["format"].(string) + data = strings.TrimSpace(data) + format = strings.TrimSpace(format) + if data == "" || format == "" { + return nil, false + } + return &core.InputAudioContent{ + Data: data, + Format: format, + ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownKeys(v, "data", "format")), + }, true + default: + return nil, false + } +} + +func cloneResponsesToolCall(call core.ToolCall) core.ToolCall { + cloned := call + cloned.ExtraFields = core.CloneUnknownJSONFields(call.ExtraFields) + cloned.Function.ExtraFields = core.CloneUnknownJSONFields(call.Function.ExtraFields) + return cloned +} + +func cloneResponsesContentPart(part core.ContentPart) core.ContentPart { + cloned := part + cloned.ExtraFields = core.CloneUnknownJSONFields(part.ExtraFields) + if part.ImageURL != nil { + image := *part.ImageURL + image.ExtraFields = core.CloneUnknownJSONFields(part.ImageURL.ExtraFields) + cloned.ImageURL = &image + } + if part.InputAudio != nil { + audio := *part.InputAudio + audio.ExtraFields = core.CloneUnknownJSONFields(part.InputAudio.ExtraFields) + cloned.InputAudio = &audio + } + return cloned +} + +func rawJSONMapFromUnknownKeys(src map[string]any, knownKeys ...string) map[string]json.RawMessage { + if len(src) == 0 { + return nil + } + known := make(map[string]struct{}, len(knownKeys)) + for _, key := range knownKeys { + known[key] = struct{}{} + } + + var extras map[string]json.RawMessage + for key, value := range src { + if _, ok := known[key]; ok { + continue + } + raw, err := json.Marshal(value) + if err != nil { + continue + } + if extras == nil { + extras = make(map[string]json.RawMessage) + } + extras[key] = raw + } + return extras +} + +func rawJSONMapFromUnknownStringKeys(src map[string]string, knownKeys ...string) map[string]json.RawMessage { + if len(src) == 0 { + return nil + } + + converted := make(map[string]any, len(src)) + for key, value := range src { + converted[key] = value + } + return rawJSONMapFromUnknownKeys(converted, knownKeys...) +} + +// ExtractContentFromInput extracts text content from responses input. +func ExtractContentFromInput(content any) string { + switch c := content.(type) { + case string: + return c + case []core.ContentPart: + texts := make([]string, 0, len(c)) + for _, part := range c { + if part.Text != "" { + texts = append(texts, part.Text) + } + } + return strings.Join(texts, " ") + case []map[string]any: + return extractTextFromMapSlice(c) + case []any: + texts := make([]string, 0, len(c)) + for _, part := range c { + if partMap, ok := part.(map[string]any); ok { + if text := extractTextFromInputMap(partMap); text != "" { + texts = append(texts, text) + } + } + } + return strings.Join(texts, " ") + default: + return "" + } +} + +func extractTextFromMapSlice(parts []map[string]any) string { + texts := make([]string, 0, len(parts)) + for _, part := range parts { + if text := extractTextFromInputMap(part); text != "" { + texts = append(texts, text) + } + } + return strings.Join(texts, " ") +} + +func extractTextFromInputMap(part map[string]any) string { + texts := make([]string, 0, 2) + if text, ok := part["text"].(string); ok && text != "" { + texts = append(texts, text) + } + if nested, ok := part["content"]; ok { + if text := ExtractContentFromInput(nested); text != "" { + texts = append(texts, text) + } + } + return strings.Join(texts, " ") +} diff --git a/internal/providers/responses_input.go b/internal/providers/responses_input.go new file mode 100644 index 000000000..cc5abd2dd --- /dev/null +++ b/internal/providers/responses_input.go @@ -0,0 +1,297 @@ +package providers + +import ( + "encoding/json" + "fmt" + "strings" + + "gomodel/internal/core" +) + +// ConvertResponsesInputToMessages converts a Responses API input payload into Chat API messages. +func ConvertResponsesInputToMessages(input any) ([]core.Message, error) { + switch in := input.(type) { + case string: + return []core.Message{{Role: "user", Content: in}}, nil + case []map[string]any: + items := make([]any, 0, len(in)) + for _, item := range in { + items = append(items, item) + } + return convertResponsesInputItems(items) + case []any: + return convertResponsesInputItems(in) + case []core.ResponsesInputElement: + items := make([]any, 0, len(in)) + for _, item := range in { + items = append(items, item) + } + return convertResponsesInputItems(items) + case nil: + return nil, core.NewInvalidRequestError("invalid responses input: unsupported type", nil) + default: + return nil, core.NewInvalidRequestError("invalid responses input: unsupported type", nil) + } +} + +func convertResponsesInputItems(items []any) ([]core.Message, error) { + messages := make([]core.Message, 0, len(items)) + var pendingAssistant *core.Message + + flushPendingAssistant := func() { + if pendingAssistant == nil { + return + } + messages = append(messages, *pendingAssistant) + pendingAssistant = nil + } + + for i, item := range items { + msg, itemType, err := convertResponsesInputItem(item, i) + if err != nil { + return nil, err + } + + if msg.Role == "assistant" { + if itemType == "message" { + flushPendingAssistant() + } + if pendingAssistant == nil { + assistant := cloneResponsesMessage(msg) + pendingAssistant = &assistant + } else if canMergeAssistantMessages(*pendingAssistant, msg) { + mergeAssistantMessage(pendingAssistant, msg) + } else { + flushPendingAssistant() + assistant := cloneResponsesMessage(msg) + pendingAssistant = &assistant + } + continue + } + + flushPendingAssistant() + messages = append(messages, msg) + } + + flushPendingAssistant() + return messages, nil +} + +func convertResponsesInputItem(item any, index int) (core.Message, string, error) { + switch typed := item.(type) { + case core.ResponsesInputElement: + return convertResponsesInputElement(typed, index) + case map[string]any: + return convertResponsesInputMap(typed, index) + default: + return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: expected object", index), nil) + } +} + +func convertResponsesInputElement(item core.ResponsesInputElement, index int) (core.Message, string, error) { + switch item.Type { + case "function_call": + name := strings.TrimSpace(item.Name) + if name == "" { + return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: function_call name is required", index), nil) + } + callID := ResponsesFunctionCallCallID(item.CallID) + return core.Message{ + Role: "assistant", + Content: "", + ContentNull: true, + ToolCalls: []core.ToolCall{ + { + ID: callID, + Type: "function", + ExtraFields: core.CloneUnknownJSONFields(item.ExtraFields), + Function: core.FunctionCall{ + Name: name, + Arguments: item.Arguments, + }, + }, + }, + }, "function_call", nil + case "function_call_output": + callID := strings.TrimSpace(item.CallID) + if callID == "" { + return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: function_call_output call_id is required", index), nil) + } + content, err := stringifyResponsesInputValueWithError(item.Output) + if err != nil { + return core.Message{}, "", core.NewInvalidRequestError( + fmt.Sprintf("invalid responses input item at index %d: function_call_output.output must be JSON-serializable", index), + err, + ) + } + return core.Message{ + Role: "tool", + ToolCallID: callID, + Content: content, + ExtraFields: core.CloneUnknownJSONFields(item.ExtraFields), + }, "function_call_output", nil + default: // message (type="" or "message") + role := strings.TrimSpace(item.Role) + if role == "" { + return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: role is required", index), nil) + } + content, ok := ConvertResponsesContentToChatContent(item.Content) + if !ok { + return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: unsupported content", index), nil) + } + return core.Message{ + Role: role, + Content: content, + ExtraFields: core.CloneUnknownJSONFields(item.ExtraFields), + }, "message", nil + } +} + +func convertResponsesInputMap(item map[string]any, index int) (core.Message, string, error) { + itemType, _ := item["type"].(string) + switch itemType { + case "function_call": + name, _ := item["name"].(string) + callID := firstNonEmptyString(item, "call_id", "id") + if strings.TrimSpace(name) == "" { + return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: function_call name is required", index), nil) + } + callID = ResponsesFunctionCallCallID(callID) + return core.Message{ + Role: "assistant", + Content: "", + ContentNull: true, + ToolCalls: []core.ToolCall{ + { + ID: callID, + Type: "function", + ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownKeys(item, "type", "call_id", "id", "name", "arguments", "status")), + Function: core.FunctionCall{ + Name: name, + Arguments: stringifyResponsesInputValue(item["arguments"]), + }, + }, + }, + }, "function_call", nil + case "function_call_output": + callID := firstNonEmptyString(item, "call_id") + if callID == "" { + return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: function_call_output call_id is required", index), nil) + } + content, err := stringifyResponsesInputValueWithError(item["output"]) + if err != nil { + return core.Message{}, "", core.NewInvalidRequestError( + fmt.Sprintf("invalid responses input item at index %d: function_call_output.output must be JSON-serializable", index), + err, + ) + } + return core.Message{ + Role: "tool", + ToolCallID: callID, + Content: content, + ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownKeys(item, "type", "call_id", "status", "output")), + }, "function_call_output", nil + } + + role, _ := item["role"].(string) + role = strings.TrimSpace(role) + if role == "" { + return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: role is required", index), nil) + } + + content, ok := ConvertResponsesContentToChatContent(item["content"]) + if !ok { + return core.Message{}, "", core.NewInvalidRequestError(fmt.Sprintf("invalid responses input item at index %d: unsupported content", index), nil) + } + return core.Message{ + Role: role, + Content: content, + ExtraFields: core.UnknownJSONFieldsFromMap(rawJSONMapFromUnknownKeys(item, "type", "role", "status", "content")), + }, "message", nil +} + +func cloneResponsesMessage(msg core.Message) core.Message { + cloned := msg + if len(msg.ToolCalls) > 0 { + cloned.ToolCalls = make([]core.ToolCall, len(msg.ToolCalls)) + for i, call := range msg.ToolCalls { + cloned.ToolCalls[i] = cloneResponsesToolCall(call) + } + } + if parts, ok := msg.Content.([]core.ContentPart); ok { + clonedParts := make([]core.ContentPart, len(parts)) + for i, part := range parts { + clonedParts[i] = cloneResponsesContentPart(part) + } + cloned.Content = clonedParts + } + cloned.ExtraFields = core.CloneUnknownJSONFields(msg.ExtraFields) + return cloned +} + +func canMergeAssistantMessages(current, next core.Message) bool { + if !current.ExtraFields.IsEmpty() || !next.ExtraFields.IsEmpty() { + return false + } + if !core.HasStructuredContent(current.Content) && !core.HasStructuredContent(next.Content) { + return true + } + return isAssistantToolCallOnlyMessage(next) +} + +func mergeAssistantMessage(dst *core.Message, src core.Message) { + if text := core.ExtractTextContent(src.Content); text != "" { + existing := core.ExtractTextContent(dst.Content) + dst.Content = existing + text + dst.ContentNull = false + } + if len(src.ToolCalls) > 0 { + dst.ToolCalls = append(dst.ToolCalls, src.ToolCalls...) + if core.ExtractTextContent(dst.Content) == "" { + dst.ContentNull = dst.ContentNull || src.ContentNull + } + } +} + +func isAssistantToolCallOnlyMessage(msg core.Message) bool { + if msg.Role != "assistant" || len(msg.ToolCalls) == 0 { + return false + } + if core.HasStructuredContent(msg.Content) { + return false + } + return core.ExtractTextContent(msg.Content) == "" +} + +func firstNonEmptyString(item map[string]any, keys ...string) string { + for _, key := range keys { + value, _ := item[key].(string) + if strings.TrimSpace(value) != "" { + return value + } + } + return "" +} + +func stringifyResponsesInputValue(value any) string { + encoded, err := stringifyResponsesInputValueWithError(value) + if err != nil { + return "" + } + return encoded +} + +func stringifyResponsesInputValueWithError(value any) (string, error) { + switch v := value.(type) { + case nil: + return "", nil + case string: + return v, nil + default: + encoded, err := json.Marshal(v) + if err != nil { + return "", err + } + return string(encoded), nil + } +} diff --git a/internal/providers/responses_output.go b/internal/providers/responses_output.go new file mode 100644 index 000000000..284e63de0 --- /dev/null +++ b/internal/providers/responses_output.go @@ -0,0 +1,185 @@ +package providers + +import ( + "encoding/json" + "strings" + + "github.com/google/uuid" + + "gomodel/internal/core" +) + +// ResponsesFunctionCallCallID returns the call id if present or generates one. +func ResponsesFunctionCallCallID(callID string) string { + if strings.TrimSpace(callID) != "" { + return callID + } + return "call_" + uuid.New().String() +} + +// ResponsesFunctionCallItemID returns a stable function-call item id. +func ResponsesFunctionCallItemID(callID string) string { + normalizedCallID := strings.TrimSpace(callID) + if normalizedCallID == "" { + normalizedCallID = "call_" + uuid.New().String() + } + return "fc_" + normalizedCallID +} + +func buildResponsesMessageContent(content any) []core.ResponsesContentItem { + switch c := content.(type) { + case string: + return []core.ResponsesContentItem{ + { + Type: "output_text", + Text: c, + Annotations: []json.RawMessage{}, + }, + } + case []core.ContentPart: + return buildResponsesContentItemsFromParts(c) + case []any: + parts, ok := core.NormalizeContentParts(c) + if !ok { + return nil + } + return buildResponsesContentItemsFromParts(parts) + default: + text := core.ExtractTextContent(content) + if text == "" { + return nil + } + return []core.ResponsesContentItem{ + { + Type: "output_text", + Text: text, + Annotations: []json.RawMessage{}, + }, + } + } +} + +func buildResponsesContentItemsFromParts(parts []core.ContentPart) []core.ResponsesContentItem { + items := make([]core.ResponsesContentItem, 0, len(parts)) + for _, part := range parts { + switch part.Type { + case "text": + items = append(items, core.ResponsesContentItem{ + Type: "output_text", + Text: part.Text, + Annotations: []json.RawMessage{}, + }) + case "image_url": + if part.ImageURL == nil { + continue + } + url := strings.TrimSpace(part.ImageURL.URL) + if url == "" { + continue + } + items = append(items, core.ResponsesContentItem{ + Type: "input_image", + ImageURL: &core.ImageURLContent{ + URL: url, + Detail: strings.TrimSpace(part.ImageURL.Detail), + MediaType: strings.TrimSpace(part.ImageURL.MediaType), + ExtraFields: core.CloneUnknownJSONFields(part.ImageURL.ExtraFields), + }, + }) + case "input_audio": + if part.InputAudio == nil { + continue + } + data := strings.TrimSpace(part.InputAudio.Data) + format := strings.TrimSpace(part.InputAudio.Format) + if data == "" || format == "" { + continue + } + items = append(items, core.ResponsesContentItem{ + Type: "input_audio", + InputAudio: &core.InputAudioContent{ + Data: data, + Format: format, + ExtraFields: core.CloneUnknownJSONFields(part.InputAudio.ExtraFields), + }, + }) + } + } + return items +} + +// BuildResponsesOutputItems converts a response message into Responses API output items. +func BuildResponsesOutputItems(msg core.ResponseMessage) []core.ResponsesOutputItem { + output := make([]core.ResponsesOutputItem, 0, len(msg.ToolCalls)+1) + contentItems := buildResponsesMessageContent(msg.Content) + if len(contentItems) > 0 || len(msg.ToolCalls) == 0 { + if len(contentItems) == 0 { + contentItems = []core.ResponsesContentItem{ + { + Type: "output_text", + Text: "", + Annotations: []json.RawMessage{}, + }, + } + } + output = append(output, core.ResponsesOutputItem{ + ID: "msg_" + uuid.New().String(), + Type: "message", + Role: "assistant", + Status: "completed", + Content: contentItems, + }) + } + for _, toolCall := range msg.ToolCalls { + callID := ResponsesFunctionCallCallID(toolCall.ID) + output = append(output, core.ResponsesOutputItem{ + ID: ResponsesFunctionCallItemID(callID), + Type: "function_call", + Status: "completed", + CallID: callID, + Name: toolCall.Function.Name, + Arguments: toolCall.Function.Arguments, + }) + } + return output +} + +// ConvertChatResponseToResponses converts a ChatResponse to a ResponsesResponse. +func ConvertChatResponseToResponses(resp *core.ChatResponse) *core.ResponsesResponse { + output := []core.ResponsesOutputItem{ + { + ID: "msg_" + uuid.New().String(), + Type: "message", + Role: "assistant", + Status: "completed", + Content: []core.ResponsesContentItem{ + { + Type: "output_text", + Text: "", + Annotations: []json.RawMessage{}, + }, + }, + }, + } + if len(resp.Choices) > 0 { + output = BuildResponsesOutputItems(resp.Choices[0].Message) + } + + return &core.ResponsesResponse{ + ID: resp.ID, + Object: "response", + CreatedAt: resp.Created, + Model: resp.Model, + Provider: resp.Provider, + Status: "completed", + Output: output, + Usage: &core.ResponsesUsage{ + InputTokens: resp.Usage.PromptTokens, + OutputTokens: resp.Usage.CompletionTokens, + TotalTokens: resp.Usage.TotalTokens, + PromptTokensDetails: resp.Usage.PromptTokensDetails, + CompletionTokensDetails: resp.Usage.CompletionTokensDetails, + RawUsage: resp.Usage.RawUsage, + }, + } +} From 3752492bfa6d4be7cde815ec969406fec1042d21 Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Wed, 29 Apr 2026 23:18:21 +0200 Subject: [PATCH 04/15] refactor(responsecache): split stream cache by API direction MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit internal/responsecache/stream_cache.go was 1309 lines bundling SSE helpers, chat-stream cache builder, and responses-stream cache builder. Split into peer files in the same package: stream_cache.go (304) entry helpers — cacheKeyRequestBody, writeCachedResponse, SSE parsing (parseSSEJSONEvents, parseCacheEventJSON, nextCacheEventBoundary, parseCacheDataLine), reconstructStreamingResponse dispatcher, stream-options helpers, JSON helpers stream_cache_chat.go (395) chatStreamCacheBuilder + state types, renderCachedChatStream, buildChatToolCalls, renderChatToolCalls, chatUsageMap, chatReasoningContent stream_cache_responses.go (627) responsesStreamCacheBuilder + state types, renderCachedResponsesStream, responsesAddedItem, responsesTerminalEventName, appendResponsesItemDeltaEvents, responsesContentDeltaEvent Pure relocation; tests pass. Co-Authored-By: Claude Opus 4.7 --- internal/responsecache/stream_cache.go | 1005 ----------------- internal/responsecache/stream_cache_chat.go | 395 +++++++ .../responsecache/stream_cache_responses.go | 627 ++++++++++ 3 files changed, 1022 insertions(+), 1005 deletions(-) create mode 100644 internal/responsecache/stream_cache_chat.go create mode 100644 internal/responsecache/stream_cache_responses.go diff --git a/internal/responsecache/stream_cache.go b/internal/responsecache/stream_cache.go index d9de69150..950432cc8 100644 --- a/internal/responsecache/stream_cache.go +++ b/internal/responsecache/stream_cache.go @@ -4,7 +4,6 @@ import ( "bytes" "encoding/json" "net/http" - "sort" "strings" "github.com/labstack/echo/v5" @@ -25,68 +24,6 @@ type streamResponseDefaults struct { Provider string } -type chatToolCallState struct { - Index int - ID string - Type string - Name string - Arguments strings.Builder -} - -type chatChoiceState struct { - Index int - Role string - Content strings.Builder - Reasoning strings.Builder - FinishReason string - Logprobs json.RawMessage - HasLogprobs bool - ToolCalls map[int]*chatToolCallState -} - -type chatStreamCacheBuilder struct { - defaults streamResponseDefaults - seen bool - ID string - Model string - Provider string - Object string - SystemFingerprint string - Created int64 - Usage map[string]any - Choices map[int]*chatChoiceState -} - -type responsesOutputState struct { - Index int - Item map[string]any - - TextParts map[int]*strings.Builder - ReasoningParts map[int]*strings.Builder - Arguments strings.Builder - HasArgs bool -} - -type responsesStreamCacheBuilder struct { - defaults streamResponseDefaults - seen bool - Response map[string]any - ID string - Object string - Model string - Provider string - Status string - CreatedAt int64 - Usage map[string]any - Error map[string]any - Output map[int]*responsesOutputState - ItemIDs map[string]int - AssistantIndex int - HasAssistant bool - ReasoningIndex int - HasReasoning bool -} - func cacheKeyRequestBody(path string, body []byte) []byte { switch path { case "/v1/chat/completions": @@ -164,163 +101,6 @@ func cacheHeaderValue(cacheType string) string { } } -func renderCachedChatStream(requestBody, cached []byte) ([]byte, error) { - var resp core.ChatResponse - if err := json.Unmarshal(cached, &resp); err != nil { - return nil, err - } - - var out bytes.Buffer - includeUsage := streamIncludeUsageRequested("/v1/chat/completions", requestBody) - usage := chatUsageMap(resp.Usage) - if !includeUsage { - usage = nil - } - for _, choice := range resp.Choices { - delta := map[string]any{} - role := strings.TrimSpace(choice.Message.Role) - if role == "" { - role = "assistant" - } - delta["role"] = role - - if content := core.ExtractTextContent(choice.Message.Content); content != "" { - delta["content"] = content - } - if reasoning := chatReasoningContent(choice.Message); reasoning != "" { - delta["reasoning_content"] = reasoning - } - if len(choice.Message.ToolCalls) > 0 { - delta["tool_calls"] = renderChatToolCalls(choice.Message.ToolCalls) - } - - renderedChoice := map[string]any{ - "index": choice.Index, - "delta": delta, - "finish_reason": choice.FinishReason, - } - if len(choice.Logprobs) > 0 { - renderedChoice["logprobs"] = choice.Logprobs - } - - chunk := map[string]any{ - "id": resp.ID, - "object": "chat.completion.chunk", - "model": resp.Model, - "choices": []map[string]any{renderedChoice}, - } - if resp.Created != 0 { - chunk["created"] = resp.Created - } - if resp.Provider != "" { - chunk["provider"] = resp.Provider - } - if resp.SystemFingerprint != "" { - chunk["system_fingerprint"] = resp.SystemFingerprint - } - if err := appendSSEJSONEvent(&out, "", chunk); err != nil { - return nil, err - } - } - - if usage != nil { - chunk := map[string]any{ - "id": resp.ID, - "object": "chat.completion.chunk", - "model": resp.Model, - "choices": []map[string]any{}, - "usage": usage, - } - if resp.Created != 0 { - chunk["created"] = resp.Created - } - if resp.Provider != "" { - chunk["provider"] = resp.Provider - } - if err := appendSSEJSONEvent(&out, "", chunk); err != nil { - return nil, err - } - } - - out.WriteString("data: [DONE]\n\n") - return out.Bytes(), nil -} - -func renderCachedResponsesStream(requestBody, cached []byte) ([]byte, error) { - var resp map[string]any - if err := json.Unmarshal(cached, &resp); err != nil { - return nil, err - } - - var out bytes.Buffer - includeUsage := streamIncludeUsageRequested("/v1/responses", requestBody) - responseWithUsage := cloneJSONMap(resp) - if !includeUsage { - delete(responseWithUsage, "usage") - } - - respID, _ := responseWithUsage["id"].(string) - respObject, _ := responseWithUsage["object"].(string) - respModel, _ := responseWithUsage["model"].(string) - respProvider, _ := responseWithUsage["provider"].(string) - respCreatedAt := responseWithUsage["created_at"] - created := map[string]any{ - "id": respID, - "object": nonEmpty(respObject, "response"), - "status": "in_progress", - "model": respModel, - "provider": respProvider, - "created_at": respCreatedAt, - } - if err := appendSSEJSONEvent(&out, "response.created", map[string]any{ - "type": "response.created", - "response": created, - }); err != nil { - return nil, err - } - - output, _ := responseWithUsage["output"].([]any) - for i, itemAny := range output { - itemMap, ok := itemAny.(map[string]any) - if !ok { - continue - } - itemID, _ := itemMap["id"].(string) - added := responsesAddedItem(itemMap) - if err := appendSSEJSONEvent(&out, "response.output_item.added", map[string]any{ - "type": "response.output_item.added", - "item": added, - "output_index": i, - }); err != nil { - return nil, err - } - if err := appendResponsesItemDeltaEvents(&out, itemMap, itemID, i); err != nil { - return nil, err - } - done := cloneJSONMap(itemMap) - if _, ok := done["status"]; !ok || done["status"] == "" { - done["status"] = "completed" - } - if err := appendSSEJSONEvent(&out, "response.output_item.done", map[string]any{ - "type": "response.output_item.done", - "item": done, - "output_index": i, - }); err != nil { - return nil, err - } - } - - terminalEventName := responsesTerminalEventName(responseWithUsage) - if err := appendSSEJSONEvent(&out, terminalEventName, map[string]any{ - "type": terminalEventName, - "response": responseWithUsage, - }); err != nil { - return nil, err - } - out.WriteString("data: [DONE]\n\n") - return out.Bytes(), nil -} - func reconstructStreamingResponse(path string, raw []byte, defaults streamResponseDefaults) ([]byte, bool) { switch path { case "/v1/chat/completions": @@ -423,658 +203,6 @@ func parseCacheDataLine(line []byte) ([]byte, bool) { return payload, true } -func (b *chatStreamCacheBuilder) OnJSONEvent(event map[string]any) { - if b == nil { - return - } - b.seen = true - - if id, ok := event["id"].(string); ok && id != "" { - b.ID = id - } - if model, ok := event["model"].(string); ok && model != "" { - b.Model = model - } - if provider, ok := event["provider"].(string); ok && provider != "" { - b.Provider = provider - } - if object, ok := event["object"].(string); ok && object != "" { - b.Object = object - } - if fingerprint, ok := event["system_fingerprint"].(string); ok && fingerprint != "" { - b.SystemFingerprint = fingerprint - } - if created, ok := jsonNumberToInt64(event["created"]); ok { - b.Created = created - } - if usage, ok := event["usage"].(map[string]any); ok { - b.Usage = cloneJSONMap(usage) - } - - choices, ok := event["choices"].([]any) - if !ok { - return - } - for _, choiceAny := range choices { - choiceMap, ok := choiceAny.(map[string]any) - if !ok { - continue - } - index, ok := jsonNumberToInt(choiceMap["index"]) - if !ok { - index = len(b.Choices) - } - state := b.choice(index) - if finish, ok := choiceMap["finish_reason"].(string); ok && finish != "" { - state.FinishReason = finish - } - if logprobs, ok := choiceMap["logprobs"]; ok { - raw, err := json.Marshal(logprobs) - if err == nil { - state.Logprobs = raw - state.HasLogprobs = true - } - } - - delta, ok := choiceMap["delta"].(map[string]any) - if !ok { - continue - } - if role, ok := delta["role"].(string); ok && role != "" { - state.Role = role - } - if content, ok := delta["content"].(string); ok && content != "" { - _, _ = state.Content.WriteString(content) - } - if reasoning, ok := delta["reasoning_content"].(string); ok && reasoning != "" { - _, _ = state.Reasoning.WriteString(reasoning) - } - - toolCalls, ok := delta["tool_calls"].([]any) - if !ok { - continue - } - for _, toolAny := range toolCalls { - toolMap, ok := toolAny.(map[string]any) - if !ok { - continue - } - toolIndex, ok := jsonNumberToInt(toolMap["index"]) - if !ok { - toolIndex = len(state.ToolCalls) - } - toolState := state.toolCall(toolIndex) - if id, ok := toolMap["id"].(string); ok && id != "" { - toolState.ID = id - } - if typ, ok := toolMap["type"].(string); ok && typ != "" { - toolState.Type = typ - } - function, ok := toolMap["function"].(map[string]any) - if !ok { - continue - } - if name, ok := function["name"].(string); ok && name != "" { - toolState.Name = name - } - if arguments, ok := function["arguments"].(string); ok && arguments != "" { - _, _ = toolState.Arguments.WriteString(arguments) - } - } - } -} - -func (b *chatStreamCacheBuilder) Build() ([]byte, bool) { - if b == nil || !b.seen { - return nil, false - } - - choiceIndexes := make([]int, 0, len(b.Choices)) - for index := range b.Choices { - choiceIndexes = append(choiceIndexes, index) - } - sort.Ints(choiceIndexes) - - choices := make([]map[string]any, 0, len(choiceIndexes)) - for _, index := range choiceIndexes { - state := b.Choices[index] - message := map[string]any{ - "role": nonEmpty(state.Role, "assistant"), - } - - content := state.Content.String() - toolCalls := buildChatToolCalls(state.ToolCalls) - switch { - case content != "": - message["content"] = content - case len(toolCalls) > 0: - message["content"] = nil - default: - message["content"] = "" - } - if len(toolCalls) > 0 { - message["tool_calls"] = toolCalls - } - if reasoning := state.Reasoning.String(); reasoning != "" { - message["reasoning_content"] = reasoning - } - - choice := map[string]any{ - "index": index, - "message": message, - "finish_reason": state.FinishReason, - } - if state.HasLogprobs { - choice["logprobs"] = state.Logprobs - } - choices = append(choices, choice) - } - - response := map[string]any{ - "id": b.ID, - "object": "chat.completion", - "model": nonEmpty(b.Model, b.defaults.Model), - "choices": choices, - } - if provider := nonEmpty(b.Provider, b.defaults.Provider); provider != "" { - response["provider"] = provider - } - if b.Created != 0 { - response["created"] = b.Created - } - if b.SystemFingerprint != "" { - response["system_fingerprint"] = b.SystemFingerprint - } - if b.Usage != nil { - response["usage"] = b.Usage - } - - data, err := json.Marshal(response) - if err != nil { - return nil, false - } - return data, true -} - -func (b *chatStreamCacheBuilder) choice(index int) *chatChoiceState { - state, ok := b.Choices[index] - if ok { - return state - } - state = &chatChoiceState{ - Index: index, - ToolCalls: make(map[int]*chatToolCallState), - } - b.Choices[index] = state - return state -} - -func (c *chatChoiceState) toolCall(index int) *chatToolCallState { - state, ok := c.ToolCalls[index] - if ok { - return state - } - state = &chatToolCallState{Index: index} - c.ToolCalls[index] = state - return state -} - -func (b *responsesStreamCacheBuilder) OnJSONEvent(event map[string]any) { - if b == nil { - return - } - b.seen = true - - eventType, _ := event["type"].(string) - switch eventType { - case "response.created", "response.completed", "response.failed", "response.incomplete", "response.done": - response, ok := event["response"].(map[string]any) - if !ok { - return - } - b.captureResponseMetadata(response) - if output, ok := response["output"].([]any); ok { - for index, itemAny := range output { - itemMap, ok := itemAny.(map[string]any) - if !ok { - continue - } - b.output(index).SetItem(itemMap) - if itemID, _ := itemMap["id"].(string); itemID != "" { - b.ItemIDs[itemID] = index - } - if itemType, _ := itemMap["type"].(string); itemType == "message" { - if role, _ := itemMap["role"].(string); role == "assistant" { - b.AssistantIndex = index - b.HasAssistant = true - } - } else if itemType == "reasoning" { - b.ReasoningIndex = index - b.HasReasoning = true - } - } - } - case "response.output_item.added", "response.output_item.done": - index, ok := jsonNumberToInt(event["output_index"]) - if !ok { - return - } - item, ok := event["item"].(map[string]any) - if !ok { - return - } - state := b.output(index) - state.SetItem(item) - if itemID, _ := item["id"].(string); itemID != "" { - b.ItemIDs[itemID] = index - } - if itemType, _ := item["type"].(string); itemType == "message" { - if role, _ := item["role"].(string); role == "assistant" { - b.AssistantIndex = index - b.HasAssistant = true - } - } else if itemType == "reasoning" { - b.ReasoningIndex = index - b.HasReasoning = true - } - case "response.output_text.delta": - delta, _ := event["delta"].(string) - if delta == "" { - return - } - contentIndex, _ := jsonNumberToInt(event["content_index"]) - index, ok := b.lookupOutputIndex(event) - if !ok { - index = 0 - if b.HasAssistant { - index = b.AssistantIndex - } - } - b.rememberOutputLocator(event, index) - b.AssistantIndex = index - b.HasAssistant = true - b.output(index).AppendText(contentIndex, delta) - case "response.reasoning_text.delta": - delta, _ := event["delta"].(string) - if delta == "" { - return - } - contentIndex, _ := jsonNumberToInt(event["content_index"]) - outputIndex, hasOutputIndex := jsonNumberToInt(event["output_index"]) - index, ok := b.lookupOutputIndex(event) - if !ok { - index = b.ensureReasoningOutputIndex(outputIndex, hasOutputIndex) - } - b.rememberOutputLocator(event, index) - b.ReasoningIndex = index - b.HasReasoning = true - b.output(index).AppendReasoning(contentIndex, delta) - case "response.function_call_arguments.delta": - index, ok := b.lookupOutputIndex(event) - if !ok { - return - } - delta, _ := event["delta"].(string) - if delta == "" { - return - } - b.output(index).AppendArguments(delta) - case "response.function_call_arguments.done": - index, ok := b.lookupOutputIndex(event) - if !ok { - return - } - arguments, _ := event["arguments"].(string) - b.output(index).SetArguments(arguments) - } -} - -func (b *responsesStreamCacheBuilder) Build() ([]byte, bool) { - if b == nil || !b.seen { - return nil, false - } - - indexes := make([]int, 0, len(b.Output)) - for index := range b.Output { - indexes = append(indexes, index) - } - sort.Ints(indexes) - - output := make([]map[string]any, 0, len(indexes)) - for _, index := range indexes { - item := b.Output[index].BuildItem() - if len(item) == 0 { - continue - } - output = append(output, item) - } - - response := cloneJSONMap(b.Response) - if response == nil { - response = map[string]any{ - "id": b.ID, - "object": nonEmpty(b.Object, "response"), - "created_at": b.CreatedAt, - "model": nonEmpty(b.Model, b.defaults.Model), - "status": nonEmpty(b.Status, "completed"), - } - if provider := nonEmpty(b.Provider, b.defaults.Provider); provider != "" { - response["provider"] = provider - } - if b.Usage != nil { - response["usage"] = b.Usage - } - if b.Error != nil { - response["error"] = b.Error - } - } - response["output"] = output - if _, ok := response["id"]; !ok { - response["id"] = b.ID - } - if _, ok := response["object"]; !ok { - response["object"] = nonEmpty(b.Object, "response") - } - if _, ok := response["created_at"]; !ok && b.CreatedAt != 0 { - response["created_at"] = b.CreatedAt - } - if _, ok := response["model"]; !ok { - response["model"] = nonEmpty(b.Model, b.defaults.Model) - } - if _, ok := response["status"]; !ok { - response["status"] = nonEmpty(b.Status, "completed") - } - if provider := nonEmpty(b.Provider, b.defaults.Provider); provider != "" { - if _, ok := response["provider"]; !ok { - response["provider"] = provider - } - } - if b.Usage != nil { - if _, ok := response["usage"]; !ok { - response["usage"] = b.Usage - } - } - if b.Error != nil { - if _, ok := response["error"]; !ok { - response["error"] = b.Error - } - } - - data, err := json.Marshal(response) - if err != nil { - return nil, false - } - return data, true -} - -func (b *responsesStreamCacheBuilder) captureResponseMetadata(response map[string]any) { - b.Response = cloneJSONMap(response) - if id, ok := response["id"].(string); ok && id != "" { - b.ID = id - } - if object, ok := response["object"].(string); ok && object != "" { - b.Object = object - } - if model, ok := response["model"].(string); ok && model != "" { - b.Model = model - } - if provider, ok := response["provider"].(string); ok && provider != "" { - b.Provider = provider - } - if status, ok := response["status"].(string); ok && status != "" { - b.Status = status - } - if createdAt, ok := jsonNumberToInt64(response["created_at"]); ok { - b.CreatedAt = createdAt - } - if usage, ok := response["usage"].(map[string]any); ok { - b.Usage = cloneJSONMap(usage) - } - if errMap, ok := response["error"].(map[string]any); ok { - b.Error = cloneJSONMap(errMap) - } -} - -func (b *responsesStreamCacheBuilder) output(index int) *responsesOutputState { - state, ok := b.Output[index] - if ok { - return state - } - state = &responsesOutputState{Index: index} - b.Output[index] = state - return state -} - -func (b *responsesStreamCacheBuilder) lookupOutputIndex(event map[string]any) (int, bool) { - if index, ok := jsonNumberToInt(event["output_index"]); ok { - return index, true - } - itemID, _ := event["item_id"].(string) - index, ok := b.ItemIDs[itemID] - return index, ok -} - -func (b *responsesStreamCacheBuilder) rememberOutputLocator(event map[string]any, index int) { - itemID, _ := event["item_id"].(string) - if itemID == "" { - return - } - b.ItemIDs[itemID] = index -} - -func (b *responsesStreamCacheBuilder) ensureReasoningOutputIndex(outputIndex int, hasOutputIndex bool) int { - if hasOutputIndex { - b.ReasoningIndex = outputIndex - b.HasReasoning = true - return outputIndex - } - if b.HasReasoning { - return b.ReasoningIndex - } - index := len(b.Output) - b.ReasoningIndex = index - b.HasReasoning = true - return index -} - -func (s *responsesOutputState) SetItem(item map[string]any) { - s.Item = cloneJSONMap(item) -} - -func (s *responsesOutputState) AppendText(contentIndex int, delta string) { - if delta == "" { - return - } - part := s.textPart(contentIndex) - _, _ = part.WriteString(delta) -} - -func (s *responsesOutputState) AppendReasoning(contentIndex int, delta string) { - if delta == "" { - return - } - part := s.reasoningPart(contentIndex) - _, _ = part.WriteString(delta) -} - -func (s *responsesOutputState) AppendArguments(delta string) { - if delta == "" { - return - } - _, _ = s.Arguments.WriteString(delta) - s.HasArgs = true -} - -func (s *responsesOutputState) SetArguments(arguments string) { - s.Arguments = strings.Builder{} - _, _ = s.Arguments.WriteString(arguments) - s.HasArgs = true -} - -func (s *responsesOutputState) textPart(contentIndex int) *strings.Builder { - if s.TextParts == nil { - s.TextParts = make(map[int]*strings.Builder) - } - part, ok := s.TextParts[contentIndex] - if ok { - return part - } - part = &strings.Builder{} - s.TextParts[contentIndex] = part - return part -} - -func (s *responsesOutputState) reasoningPart(contentIndex int) *strings.Builder { - if s.ReasoningParts == nil { - s.ReasoningParts = make(map[int]*strings.Builder) - } - part, ok := s.ReasoningParts[contentIndex] - if ok { - return part - } - part = &strings.Builder{} - s.ReasoningParts[contentIndex] = part - return part -} - -func (s *responsesOutputState) BuildItem() map[string]any { - item := cloneJSONMap(s.Item) - if item == nil { - item = make(map[string]any) - } - - itemType, _ := item["type"].(string) - if len(s.TextParts) > 0 { - if itemType == "" { - itemType = "message" - item["type"] = itemType - } - if item["role"] == nil { - item["role"] = "assistant" - } - item["content"] = buildResponsesContentParts(item["content"], "output_text", s.TextParts) - } - if len(s.ReasoningParts) > 0 { - if itemType == "" { - itemType = "reasoning" - item["type"] = itemType - } - targetField := "content" - if itemType == "reasoning" { - targetField = "summary" - } else { - if _, ok := item["summary"].([]any); ok { - targetField = "summary" - } - } - item[targetField] = buildResponsesContentParts(item[targetField], "reasoning_text", s.ReasoningParts) - } - if s.HasArgs { - if itemType == "" { - itemType = "function_call" - item["type"] = itemType - } - item["arguments"] = s.Arguments.String() - } - if _, ok := item["status"]; !ok || item["status"] == "" { - item["status"] = "completed" - } - - return item -} - -func buildResponsesContentParts(existing any, partType string, parts map[int]*strings.Builder) []map[string]any { - if len(parts) == 0 { - return nil - } - - existingParts, _ := existing.([]any) - maxIndex := len(existingParts) - 1 - for index := range parts { - if index > maxIndex { - maxIndex = index - } - } - - built := make([]map[string]any, 0, maxIndex+1) - for index := 0; index <= maxIndex; index++ { - existingPart, existingOK := cloneJSONPart(existingParts, index) - partBuilder, hasPart := parts[index] - - switch { - case hasPart: - if existingPart == nil { - existingPart = make(map[string]any) - } - existingPart["type"] = partType - existingPart["text"] = partBuilder.String() - built = append(built, existingPart) - case existingOK: - built = append(built, existingPart) - } - } - - return built -} - -func cloneJSONPart(parts []any, index int) (map[string]any, bool) { - if index < 0 || index >= len(parts) { - return nil, false - } - part, ok := parts[index].(map[string]any) - if !ok { - return nil, false - } - return cloneJSONMap(part), true -} - -func buildChatToolCalls(states map[int]*chatToolCallState) []map[string]any { - if len(states) == 0 { - return nil - } - - indexes := make([]int, 0, len(states)) - for index := range states { - indexes = append(indexes, index) - } - sort.Ints(indexes) - - toolCalls := make([]map[string]any, 0, len(indexes)) - for _, index := range indexes { - state := states[index] - toolCall := map[string]any{ - "id": state.ID, - "type": nonEmpty(state.Type, "function"), - "index": index, - "function": map[string]any{ - "name": state.Name, - "arguments": state.Arguments.String(), - }, - } - toolCalls = append(toolCalls, toolCall) - } - return toolCalls -} - -func renderChatToolCalls(toolCalls []core.ToolCall) []map[string]any { - if len(toolCalls) == 0 { - return nil - } - rendered := make([]map[string]any, 0, len(toolCalls)) - for index, toolCall := range toolCalls { - rendered = append(rendered, map[string]any{ - "index": index, - "id": toolCall.ID, - "type": nonEmpty(toolCall.Type, "function"), - "function": map[string]any{ - "name": toolCall.Function.Name, - "arguments": toolCall.Function.Arguments, - }, - }) - } - return rendered -} - func normalizeStreamOptionsForCache(src *core.StreamOptions) *core.StreamOptions { if src == nil || !src.IncludeUsage { return nil @@ -1102,139 +230,6 @@ func streamIncludeUsageRequested(path string, requestBody []byte) bool { } } -func chatReasoningContent(message core.ResponseMessage) string { - raw := message.ExtraFields.Lookup("reasoning_content") - if len(raw) == 0 { - return "" - } - var reasoning string - if err := json.Unmarshal(raw, &reasoning); err != nil { - return "" - } - return reasoning -} - -func responsesAddedItem(item map[string]any) map[string]any { - added := cloneJSONMap(item) - if added == nil { - return nil - } - added["status"] = "in_progress" - delete(added, "arguments") - if _, ok := added["content"].([]any); ok { - added["content"] = []any{} - } - if _, ok := added["summary"].([]any); ok { - added["summary"] = []any{} - } - return added -} - -func responsesTerminalEventName(response map[string]any) string { - status, _ := response["status"].(string) - switch status { - case "failed": - return "response.failed" - case "incomplete": - return "response.incomplete" - default: - return "response.completed" - } -} - -func appendResponsesItemDeltaEvents(out *bytes.Buffer, item map[string]any, itemID string, outputIndex int) error { - if out == nil || item == nil { - return nil - } - - if arguments, ok := item["arguments"].(string); ok && arguments != "" { - if err := appendSSEJSONEvent(out, "response.function_call_arguments.delta", map[string]any{ - "type": "response.function_call_arguments.delta", - "item_id": itemID, - "output_index": outputIndex, - "delta": arguments, - }); err != nil { - return err - } - if err := appendSSEJSONEvent(out, "response.function_call_arguments.done", map[string]any{ - "type": "response.function_call_arguments.done", - "item_id": itemID, - "output_index": outputIndex, - "arguments": arguments, - }); err != nil { - return err - } - } - - for _, key := range []string{"content", "summary"} { - parts, ok := item[key].([]any) - if !ok { - continue - } - for contentIndex, partAny := range parts { - part, ok := partAny.(map[string]any) - if !ok { - continue - } - eventName, payload, ok := responsesContentDeltaEvent(part, itemID, outputIndex, contentIndex) - if !ok { - continue - } - if err := appendSSEJSONEvent(out, eventName, payload); err != nil { - return err - } - } - } - - return nil -} - -func responsesContentDeltaEvent(part map[string]any, itemID string, outputIndex, contentIndex int) (string, map[string]any, bool) { - partType, _ := part["type"].(string) - text, _ := part["text"].(string) - if partType == "" || text == "" { - return "", nil, false - } - - var eventName string - switch partType { - case "output_text": - eventName = "response.output_text.delta" - case "reasoning_text": - eventName = "response.reasoning_text.delta" - default: - return "", nil, false - } - - payload := map[string]any{ - "type": eventName, - "delta": text, - "output_index": outputIndex, - "content_index": contentIndex, - } - if itemID != "" { - payload["item_id"] = itemID - } - - return eventName, payload, true -} - -func chatUsageMap(usage core.Usage) map[string]any { - if usage.PromptTokens == 0 && - usage.CompletionTokens == 0 && - usage.TotalTokens == 0 && - usage.PromptTokensDetails == nil && - usage.CompletionTokensDetails == nil && - len(usage.RawUsage) == 0 { - return nil - } - result, err := toJSONMap(usage) - if err != nil { - return nil - } - return result -} - func appendSSEJSONEvent(out *bytes.Buffer, eventName string, payload any) error { data, err := json.Marshal(payload) if err != nil { diff --git a/internal/responsecache/stream_cache_chat.go b/internal/responsecache/stream_cache_chat.go new file mode 100644 index 000000000..a654ddfba --- /dev/null +++ b/internal/responsecache/stream_cache_chat.go @@ -0,0 +1,395 @@ +package responsecache + +import ( + "bytes" + "encoding/json" + "sort" + "strings" + + "gomodel/internal/core" +) + +type chatToolCallState struct { + Index int + ID string + Type string + Name string + Arguments strings.Builder +} + +type chatChoiceState struct { + Index int + Role string + Content strings.Builder + Reasoning strings.Builder + FinishReason string + Logprobs json.RawMessage + HasLogprobs bool + ToolCalls map[int]*chatToolCallState +} + +type chatStreamCacheBuilder struct { + defaults streamResponseDefaults + seen bool + ID string + Model string + Provider string + Object string + SystemFingerprint string + Created int64 + Usage map[string]any + Choices map[int]*chatChoiceState +} + +func renderCachedChatStream(requestBody, cached []byte) ([]byte, error) { + var resp core.ChatResponse + if err := json.Unmarshal(cached, &resp); err != nil { + return nil, err + } + + var out bytes.Buffer + includeUsage := streamIncludeUsageRequested("/v1/chat/completions", requestBody) + usage := chatUsageMap(resp.Usage) + if !includeUsage { + usage = nil + } + for _, choice := range resp.Choices { + delta := map[string]any{} + role := strings.TrimSpace(choice.Message.Role) + if role == "" { + role = "assistant" + } + delta["role"] = role + + if content := core.ExtractTextContent(choice.Message.Content); content != "" { + delta["content"] = content + } + if reasoning := chatReasoningContent(choice.Message); reasoning != "" { + delta["reasoning_content"] = reasoning + } + if len(choice.Message.ToolCalls) > 0 { + delta["tool_calls"] = renderChatToolCalls(choice.Message.ToolCalls) + } + + renderedChoice := map[string]any{ + "index": choice.Index, + "delta": delta, + "finish_reason": choice.FinishReason, + } + if len(choice.Logprobs) > 0 { + renderedChoice["logprobs"] = choice.Logprobs + } + + chunk := map[string]any{ + "id": resp.ID, + "object": "chat.completion.chunk", + "model": resp.Model, + "choices": []map[string]any{renderedChoice}, + } + if resp.Created != 0 { + chunk["created"] = resp.Created + } + if resp.Provider != "" { + chunk["provider"] = resp.Provider + } + if resp.SystemFingerprint != "" { + chunk["system_fingerprint"] = resp.SystemFingerprint + } + if err := appendSSEJSONEvent(&out, "", chunk); err != nil { + return nil, err + } + } + + if usage != nil { + chunk := map[string]any{ + "id": resp.ID, + "object": "chat.completion.chunk", + "model": resp.Model, + "choices": []map[string]any{}, + "usage": usage, + } + if resp.Created != 0 { + chunk["created"] = resp.Created + } + if resp.Provider != "" { + chunk["provider"] = resp.Provider + } + if err := appendSSEJSONEvent(&out, "", chunk); err != nil { + return nil, err + } + } + + out.WriteString("data: [DONE]\n\n") + return out.Bytes(), nil +} + +func (b *chatStreamCacheBuilder) OnJSONEvent(event map[string]any) { + if b == nil { + return + } + b.seen = true + + if id, ok := event["id"].(string); ok && id != "" { + b.ID = id + } + if model, ok := event["model"].(string); ok && model != "" { + b.Model = model + } + if provider, ok := event["provider"].(string); ok && provider != "" { + b.Provider = provider + } + if object, ok := event["object"].(string); ok && object != "" { + b.Object = object + } + if fingerprint, ok := event["system_fingerprint"].(string); ok && fingerprint != "" { + b.SystemFingerprint = fingerprint + } + if created, ok := jsonNumberToInt64(event["created"]); ok { + b.Created = created + } + if usage, ok := event["usage"].(map[string]any); ok { + b.Usage = cloneJSONMap(usage) + } + + choices, ok := event["choices"].([]any) + if !ok { + return + } + for _, choiceAny := range choices { + choiceMap, ok := choiceAny.(map[string]any) + if !ok { + continue + } + index, ok := jsonNumberToInt(choiceMap["index"]) + if !ok { + index = len(b.Choices) + } + state := b.choice(index) + if finish, ok := choiceMap["finish_reason"].(string); ok && finish != "" { + state.FinishReason = finish + } + if logprobs, ok := choiceMap["logprobs"]; ok { + raw, err := json.Marshal(logprobs) + if err == nil { + state.Logprobs = raw + state.HasLogprobs = true + } + } + + delta, ok := choiceMap["delta"].(map[string]any) + if !ok { + continue + } + if role, ok := delta["role"].(string); ok && role != "" { + state.Role = role + } + if content, ok := delta["content"].(string); ok && content != "" { + _, _ = state.Content.WriteString(content) + } + if reasoning, ok := delta["reasoning_content"].(string); ok && reasoning != "" { + _, _ = state.Reasoning.WriteString(reasoning) + } + + toolCalls, ok := delta["tool_calls"].([]any) + if !ok { + continue + } + for _, toolAny := range toolCalls { + toolMap, ok := toolAny.(map[string]any) + if !ok { + continue + } + toolIndex, ok := jsonNumberToInt(toolMap["index"]) + if !ok { + toolIndex = len(state.ToolCalls) + } + toolState := state.toolCall(toolIndex) + if id, ok := toolMap["id"].(string); ok && id != "" { + toolState.ID = id + } + if typ, ok := toolMap["type"].(string); ok && typ != "" { + toolState.Type = typ + } + function, ok := toolMap["function"].(map[string]any) + if !ok { + continue + } + if name, ok := function["name"].(string); ok && name != "" { + toolState.Name = name + } + if arguments, ok := function["arguments"].(string); ok && arguments != "" { + _, _ = toolState.Arguments.WriteString(arguments) + } + } + } +} + +func (b *chatStreamCacheBuilder) Build() ([]byte, bool) { + if b == nil || !b.seen { + return nil, false + } + + choiceIndexes := make([]int, 0, len(b.Choices)) + for index := range b.Choices { + choiceIndexes = append(choiceIndexes, index) + } + sort.Ints(choiceIndexes) + + choices := make([]map[string]any, 0, len(choiceIndexes)) + for _, index := range choiceIndexes { + state := b.Choices[index] + message := map[string]any{ + "role": nonEmpty(state.Role, "assistant"), + } + + content := state.Content.String() + toolCalls := buildChatToolCalls(state.ToolCalls) + switch { + case content != "": + message["content"] = content + case len(toolCalls) > 0: + message["content"] = nil + default: + message["content"] = "" + } + if len(toolCalls) > 0 { + message["tool_calls"] = toolCalls + } + if reasoning := state.Reasoning.String(); reasoning != "" { + message["reasoning_content"] = reasoning + } + + choice := map[string]any{ + "index": index, + "message": message, + "finish_reason": state.FinishReason, + } + if state.HasLogprobs { + choice["logprobs"] = state.Logprobs + } + choices = append(choices, choice) + } + + response := map[string]any{ + "id": b.ID, + "object": "chat.completion", + "model": nonEmpty(b.Model, b.defaults.Model), + "choices": choices, + } + if provider := nonEmpty(b.Provider, b.defaults.Provider); provider != "" { + response["provider"] = provider + } + if b.Created != 0 { + response["created"] = b.Created + } + if b.SystemFingerprint != "" { + response["system_fingerprint"] = b.SystemFingerprint + } + if b.Usage != nil { + response["usage"] = b.Usage + } + + data, err := json.Marshal(response) + if err != nil { + return nil, false + } + return data, true +} + +func (b *chatStreamCacheBuilder) choice(index int) *chatChoiceState { + state, ok := b.Choices[index] + if ok { + return state + } + state = &chatChoiceState{ + Index: index, + ToolCalls: make(map[int]*chatToolCallState), + } + b.Choices[index] = state + return state +} + +func (c *chatChoiceState) toolCall(index int) *chatToolCallState { + state, ok := c.ToolCalls[index] + if ok { + return state + } + state = &chatToolCallState{Index: index} + c.ToolCalls[index] = state + return state +} + +func buildChatToolCalls(states map[int]*chatToolCallState) []map[string]any { + if len(states) == 0 { + return nil + } + + indexes := make([]int, 0, len(states)) + for index := range states { + indexes = append(indexes, index) + } + sort.Ints(indexes) + + toolCalls := make([]map[string]any, 0, len(indexes)) + for _, index := range indexes { + state := states[index] + toolCall := map[string]any{ + "id": state.ID, + "type": nonEmpty(state.Type, "function"), + "index": index, + "function": map[string]any{ + "name": state.Name, + "arguments": state.Arguments.String(), + }, + } + toolCalls = append(toolCalls, toolCall) + } + return toolCalls +} + +func renderChatToolCalls(toolCalls []core.ToolCall) []map[string]any { + if len(toolCalls) == 0 { + return nil + } + rendered := make([]map[string]any, 0, len(toolCalls)) + for index, toolCall := range toolCalls { + rendered = append(rendered, map[string]any{ + "index": index, + "id": toolCall.ID, + "type": nonEmpty(toolCall.Type, "function"), + "function": map[string]any{ + "name": toolCall.Function.Name, + "arguments": toolCall.Function.Arguments, + }, + }) + } + return rendered +} + +func chatUsageMap(usage core.Usage) map[string]any { + if usage.PromptTokens == 0 && + usage.CompletionTokens == 0 && + usage.TotalTokens == 0 && + usage.PromptTokensDetails == nil && + usage.CompletionTokensDetails == nil && + len(usage.RawUsage) == 0 { + return nil + } + result, err := toJSONMap(usage) + if err != nil { + return nil + } + return result +} + +func chatReasoningContent(message core.ResponseMessage) string { + raw := message.ExtraFields.Lookup("reasoning_content") + if len(raw) == 0 { + return "" + } + var reasoning string + if err := json.Unmarshal(raw, &reasoning); err != nil { + return "" + } + return reasoning +} diff --git a/internal/responsecache/stream_cache_responses.go b/internal/responsecache/stream_cache_responses.go new file mode 100644 index 000000000..aaf5da36f --- /dev/null +++ b/internal/responsecache/stream_cache_responses.go @@ -0,0 +1,627 @@ +package responsecache + +import ( + "bytes" + "encoding/json" + "sort" + "strings" +) + +type responsesOutputState struct { + Index int + Item map[string]any + + TextParts map[int]*strings.Builder + ReasoningParts map[int]*strings.Builder + Arguments strings.Builder + HasArgs bool +} + +type responsesStreamCacheBuilder struct { + defaults streamResponseDefaults + seen bool + Response map[string]any + ID string + Object string + Model string + Provider string + Status string + CreatedAt int64 + Usage map[string]any + Error map[string]any + Output map[int]*responsesOutputState + ItemIDs map[string]int + AssistantIndex int + HasAssistant bool + ReasoningIndex int + HasReasoning bool +} + +func renderCachedResponsesStream(requestBody, cached []byte) ([]byte, error) { + var resp map[string]any + if err := json.Unmarshal(cached, &resp); err != nil { + return nil, err + } + + var out bytes.Buffer + includeUsage := streamIncludeUsageRequested("/v1/responses", requestBody) + responseWithUsage := cloneJSONMap(resp) + if !includeUsage { + delete(responseWithUsage, "usage") + } + + respID, _ := responseWithUsage["id"].(string) + respObject, _ := responseWithUsage["object"].(string) + respModel, _ := responseWithUsage["model"].(string) + respProvider, _ := responseWithUsage["provider"].(string) + respCreatedAt := responseWithUsage["created_at"] + created := map[string]any{ + "id": respID, + "object": nonEmpty(respObject, "response"), + "status": "in_progress", + "model": respModel, + "provider": respProvider, + "created_at": respCreatedAt, + } + if err := appendSSEJSONEvent(&out, "response.created", map[string]any{ + "type": "response.created", + "response": created, + }); err != nil { + return nil, err + } + + output, _ := responseWithUsage["output"].([]any) + for i, itemAny := range output { + itemMap, ok := itemAny.(map[string]any) + if !ok { + continue + } + itemID, _ := itemMap["id"].(string) + added := responsesAddedItem(itemMap) + if err := appendSSEJSONEvent(&out, "response.output_item.added", map[string]any{ + "type": "response.output_item.added", + "item": added, + "output_index": i, + }); err != nil { + return nil, err + } + if err := appendResponsesItemDeltaEvents(&out, itemMap, itemID, i); err != nil { + return nil, err + } + done := cloneJSONMap(itemMap) + if _, ok := done["status"]; !ok || done["status"] == "" { + done["status"] = "completed" + } + if err := appendSSEJSONEvent(&out, "response.output_item.done", map[string]any{ + "type": "response.output_item.done", + "item": done, + "output_index": i, + }); err != nil { + return nil, err + } + } + + terminalEventName := responsesTerminalEventName(responseWithUsage) + if err := appendSSEJSONEvent(&out, terminalEventName, map[string]any{ + "type": terminalEventName, + "response": responseWithUsage, + }); err != nil { + return nil, err + } + out.WriteString("data: [DONE]\n\n") + return out.Bytes(), nil +} + +func (b *responsesStreamCacheBuilder) OnJSONEvent(event map[string]any) { + if b == nil { + return + } + b.seen = true + + eventType, _ := event["type"].(string) + switch eventType { + case "response.created", "response.completed", "response.failed", "response.incomplete", "response.done": + response, ok := event["response"].(map[string]any) + if !ok { + return + } + b.captureResponseMetadata(response) + if output, ok := response["output"].([]any); ok { + for index, itemAny := range output { + itemMap, ok := itemAny.(map[string]any) + if !ok { + continue + } + b.output(index).SetItem(itemMap) + if itemID, _ := itemMap["id"].(string); itemID != "" { + b.ItemIDs[itemID] = index + } + if itemType, _ := itemMap["type"].(string); itemType == "message" { + if role, _ := itemMap["role"].(string); role == "assistant" { + b.AssistantIndex = index + b.HasAssistant = true + } + } else if itemType == "reasoning" { + b.ReasoningIndex = index + b.HasReasoning = true + } + } + } + case "response.output_item.added", "response.output_item.done": + index, ok := jsonNumberToInt(event["output_index"]) + if !ok { + return + } + item, ok := event["item"].(map[string]any) + if !ok { + return + } + state := b.output(index) + state.SetItem(item) + if itemID, _ := item["id"].(string); itemID != "" { + b.ItemIDs[itemID] = index + } + if itemType, _ := item["type"].(string); itemType == "message" { + if role, _ := item["role"].(string); role == "assistant" { + b.AssistantIndex = index + b.HasAssistant = true + } + } else if itemType == "reasoning" { + b.ReasoningIndex = index + b.HasReasoning = true + } + case "response.output_text.delta": + delta, _ := event["delta"].(string) + if delta == "" { + return + } + contentIndex, _ := jsonNumberToInt(event["content_index"]) + index, ok := b.lookupOutputIndex(event) + if !ok { + index = 0 + if b.HasAssistant { + index = b.AssistantIndex + } + } + b.rememberOutputLocator(event, index) + b.AssistantIndex = index + b.HasAssistant = true + b.output(index).AppendText(contentIndex, delta) + case "response.reasoning_text.delta": + delta, _ := event["delta"].(string) + if delta == "" { + return + } + contentIndex, _ := jsonNumberToInt(event["content_index"]) + outputIndex, hasOutputIndex := jsonNumberToInt(event["output_index"]) + index, ok := b.lookupOutputIndex(event) + if !ok { + index = b.ensureReasoningOutputIndex(outputIndex, hasOutputIndex) + } + b.rememberOutputLocator(event, index) + b.ReasoningIndex = index + b.HasReasoning = true + b.output(index).AppendReasoning(contentIndex, delta) + case "response.function_call_arguments.delta": + index, ok := b.lookupOutputIndex(event) + if !ok { + return + } + delta, _ := event["delta"].(string) + if delta == "" { + return + } + b.output(index).AppendArguments(delta) + case "response.function_call_arguments.done": + index, ok := b.lookupOutputIndex(event) + if !ok { + return + } + arguments, _ := event["arguments"].(string) + b.output(index).SetArguments(arguments) + } +} + +func (b *responsesStreamCacheBuilder) Build() ([]byte, bool) { + if b == nil || !b.seen { + return nil, false + } + + indexes := make([]int, 0, len(b.Output)) + for index := range b.Output { + indexes = append(indexes, index) + } + sort.Ints(indexes) + + output := make([]map[string]any, 0, len(indexes)) + for _, index := range indexes { + item := b.Output[index].BuildItem() + if len(item) == 0 { + continue + } + output = append(output, item) + } + + response := cloneJSONMap(b.Response) + if response == nil { + response = map[string]any{ + "id": b.ID, + "object": nonEmpty(b.Object, "response"), + "created_at": b.CreatedAt, + "model": nonEmpty(b.Model, b.defaults.Model), + "status": nonEmpty(b.Status, "completed"), + } + if provider := nonEmpty(b.Provider, b.defaults.Provider); provider != "" { + response["provider"] = provider + } + if b.Usage != nil { + response["usage"] = b.Usage + } + if b.Error != nil { + response["error"] = b.Error + } + } + response["output"] = output + if _, ok := response["id"]; !ok { + response["id"] = b.ID + } + if _, ok := response["object"]; !ok { + response["object"] = nonEmpty(b.Object, "response") + } + if _, ok := response["created_at"]; !ok && b.CreatedAt != 0 { + response["created_at"] = b.CreatedAt + } + if _, ok := response["model"]; !ok { + response["model"] = nonEmpty(b.Model, b.defaults.Model) + } + if _, ok := response["status"]; !ok { + response["status"] = nonEmpty(b.Status, "completed") + } + if provider := nonEmpty(b.Provider, b.defaults.Provider); provider != "" { + if _, ok := response["provider"]; !ok { + response["provider"] = provider + } + } + if b.Usage != nil { + if _, ok := response["usage"]; !ok { + response["usage"] = b.Usage + } + } + if b.Error != nil { + if _, ok := response["error"]; !ok { + response["error"] = b.Error + } + } + + data, err := json.Marshal(response) + if err != nil { + return nil, false + } + return data, true +} + +func (b *responsesStreamCacheBuilder) captureResponseMetadata(response map[string]any) { + b.Response = cloneJSONMap(response) + if id, ok := response["id"].(string); ok && id != "" { + b.ID = id + } + if object, ok := response["object"].(string); ok && object != "" { + b.Object = object + } + if model, ok := response["model"].(string); ok && model != "" { + b.Model = model + } + if provider, ok := response["provider"].(string); ok && provider != "" { + b.Provider = provider + } + if status, ok := response["status"].(string); ok && status != "" { + b.Status = status + } + if createdAt, ok := jsonNumberToInt64(response["created_at"]); ok { + b.CreatedAt = createdAt + } + if usage, ok := response["usage"].(map[string]any); ok { + b.Usage = cloneJSONMap(usage) + } + if errMap, ok := response["error"].(map[string]any); ok { + b.Error = cloneJSONMap(errMap) + } +} + +func (b *responsesStreamCacheBuilder) output(index int) *responsesOutputState { + state, ok := b.Output[index] + if ok { + return state + } + state = &responsesOutputState{Index: index} + b.Output[index] = state + return state +} + +func (b *responsesStreamCacheBuilder) lookupOutputIndex(event map[string]any) (int, bool) { + if index, ok := jsonNumberToInt(event["output_index"]); ok { + return index, true + } + itemID, _ := event["item_id"].(string) + index, ok := b.ItemIDs[itemID] + return index, ok +} + +func (b *responsesStreamCacheBuilder) rememberOutputLocator(event map[string]any, index int) { + itemID, _ := event["item_id"].(string) + if itemID == "" { + return + } + b.ItemIDs[itemID] = index +} + +func (b *responsesStreamCacheBuilder) ensureReasoningOutputIndex(outputIndex int, hasOutputIndex bool) int { + if hasOutputIndex { + b.ReasoningIndex = outputIndex + b.HasReasoning = true + return outputIndex + } + if b.HasReasoning { + return b.ReasoningIndex + } + index := len(b.Output) + b.ReasoningIndex = index + b.HasReasoning = true + return index +} + +func (s *responsesOutputState) SetItem(item map[string]any) { + s.Item = cloneJSONMap(item) +} + +func (s *responsesOutputState) AppendText(contentIndex int, delta string) { + if delta == "" { + return + } + part := s.textPart(contentIndex) + _, _ = part.WriteString(delta) +} + +func (s *responsesOutputState) AppendReasoning(contentIndex int, delta string) { + if delta == "" { + return + } + part := s.reasoningPart(contentIndex) + _, _ = part.WriteString(delta) +} + +func (s *responsesOutputState) AppendArguments(delta string) { + if delta == "" { + return + } + _, _ = s.Arguments.WriteString(delta) + s.HasArgs = true +} + +func (s *responsesOutputState) SetArguments(arguments string) { + s.Arguments = strings.Builder{} + _, _ = s.Arguments.WriteString(arguments) + s.HasArgs = true +} + +func (s *responsesOutputState) textPart(contentIndex int) *strings.Builder { + if s.TextParts == nil { + s.TextParts = make(map[int]*strings.Builder) + } + part, ok := s.TextParts[contentIndex] + if ok { + return part + } + part = &strings.Builder{} + s.TextParts[contentIndex] = part + return part +} + +func (s *responsesOutputState) reasoningPart(contentIndex int) *strings.Builder { + if s.ReasoningParts == nil { + s.ReasoningParts = make(map[int]*strings.Builder) + } + part, ok := s.ReasoningParts[contentIndex] + if ok { + return part + } + part = &strings.Builder{} + s.ReasoningParts[contentIndex] = part + return part +} + +func (s *responsesOutputState) BuildItem() map[string]any { + item := cloneJSONMap(s.Item) + if item == nil { + item = make(map[string]any) + } + + itemType, _ := item["type"].(string) + if len(s.TextParts) > 0 { + if itemType == "" { + itemType = "message" + item["type"] = itemType + } + if item["role"] == nil { + item["role"] = "assistant" + } + item["content"] = buildResponsesContentParts(item["content"], "output_text", s.TextParts) + } + if len(s.ReasoningParts) > 0 { + if itemType == "" { + itemType = "reasoning" + item["type"] = itemType + } + targetField := "content" + if itemType == "reasoning" { + targetField = "summary" + } else { + if _, ok := item["summary"].([]any); ok { + targetField = "summary" + } + } + item[targetField] = buildResponsesContentParts(item[targetField], "reasoning_text", s.ReasoningParts) + } + if s.HasArgs { + if itemType == "" { + itemType = "function_call" + item["type"] = itemType + } + item["arguments"] = s.Arguments.String() + } + if _, ok := item["status"]; !ok || item["status"] == "" { + item["status"] = "completed" + } + + return item +} + +func buildResponsesContentParts(existing any, partType string, parts map[int]*strings.Builder) []map[string]any { + if len(parts) == 0 { + return nil + } + + existingParts, _ := existing.([]any) + maxIndex := len(existingParts) - 1 + for index := range parts { + if index > maxIndex { + maxIndex = index + } + } + + built := make([]map[string]any, 0, maxIndex+1) + for index := 0; index <= maxIndex; index++ { + existingPart, existingOK := cloneJSONPart(existingParts, index) + partBuilder, hasPart := parts[index] + + switch { + case hasPart: + if existingPart == nil { + existingPart = make(map[string]any) + } + existingPart["type"] = partType + existingPart["text"] = partBuilder.String() + built = append(built, existingPart) + case existingOK: + built = append(built, existingPart) + } + } + + return built +} + +func cloneJSONPart(parts []any, index int) (map[string]any, bool) { + if index < 0 || index >= len(parts) { + return nil, false + } + part, ok := parts[index].(map[string]any) + if !ok { + return nil, false + } + return cloneJSONMap(part), true +} + +func responsesAddedItem(item map[string]any) map[string]any { + added := cloneJSONMap(item) + if added == nil { + return nil + } + added["status"] = "in_progress" + delete(added, "arguments") + if _, ok := added["content"].([]any); ok { + added["content"] = []any{} + } + if _, ok := added["summary"].([]any); ok { + added["summary"] = []any{} + } + return added +} + +func responsesTerminalEventName(response map[string]any) string { + status, _ := response["status"].(string) + switch status { + case "failed": + return "response.failed" + case "incomplete": + return "response.incomplete" + default: + return "response.completed" + } +} + +func appendResponsesItemDeltaEvents(out *bytes.Buffer, item map[string]any, itemID string, outputIndex int) error { + if out == nil || item == nil { + return nil + } + + if arguments, ok := item["arguments"].(string); ok && arguments != "" { + if err := appendSSEJSONEvent(out, "response.function_call_arguments.delta", map[string]any{ + "type": "response.function_call_arguments.delta", + "item_id": itemID, + "output_index": outputIndex, + "delta": arguments, + }); err != nil { + return err + } + if err := appendSSEJSONEvent(out, "response.function_call_arguments.done", map[string]any{ + "type": "response.function_call_arguments.done", + "item_id": itemID, + "output_index": outputIndex, + "arguments": arguments, + }); err != nil { + return err + } + } + + for _, key := range []string{"content", "summary"} { + parts, ok := item[key].([]any) + if !ok { + continue + } + for contentIndex, partAny := range parts { + part, ok := partAny.(map[string]any) + if !ok { + continue + } + eventName, payload, ok := responsesContentDeltaEvent(part, itemID, outputIndex, contentIndex) + if !ok { + continue + } + if err := appendSSEJSONEvent(out, eventName, payload); err != nil { + return err + } + } + } + + return nil +} + +func responsesContentDeltaEvent(part map[string]any, itemID string, outputIndex, contentIndex int) (string, map[string]any, bool) { + partType, _ := part["type"].(string) + text, _ := part["text"].(string) + if partType == "" || text == "" { + return "", nil, false + } + + var eventName string + switch partType { + case "output_text": + eventName = "response.output_text.delta" + case "reasoning_text": + eventName = "response.reasoning_text.delta" + default: + return "", nil, false + } + + payload := map[string]any{ + "type": eventName, + "delta": text, + "output_index": outputIndex, + "content_index": contentIndex, + } + if itemID != "" { + payload["item_id"] = itemID + } + + return eventName, payload, true +} From 7f28eeab46e6850ec518689d45e6920a350be80f Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Wed, 29 Apr 2026 23:24:10 +0200 Subject: [PATCH 05/15] refactor(providers): split ModelRegistry into per-concern peer files internal/providers/registry.go was 1902 lines mixing seven concerns (registration, init/refresh, cache load+save, metadata enrichment, lookup, categories, runtime snapshots). Split into peer files in the same package: registry.go (808) struct, registration, lookup, ListModels, categories, runtime snapshots, provider-name/type indirection, splitModelSelector registry_init.go (494) Initialize/Refresh/InitializeAsync, fetchProviderInventory, refresh semaphore, StartBackgroundRefresh, RefreshModelList, provider runtime updates registry_cache.go (227) LoadFromCache, SaveToCache registry_metadata.go (417) SetModelList, EnrichModels, ResolveMetadata, GetModelMetadata, ResolvePricing, snapshot helpers, applyConfigMetadataOverrides, enrichProviderModelMaps, registryAccessor Pure relocation; tests pass. Co-Authored-By: Claude Opus 4.7 --- internal/providers/registry.go | 1096 ----------------------- internal/providers/registry_cache.go | 227 +++++ internal/providers/registry_init.go | 494 ++++++++++ internal/providers/registry_metadata.go | 417 +++++++++ 4 files changed, 1138 insertions(+), 1096 deletions(-) create mode 100644 internal/providers/registry_cache.go create mode 100644 internal/providers/registry_init.go create mode 100644 internal/providers/registry_metadata.go diff --git a/internal/providers/registry.go b/internal/providers/registry.go index 2f5eb44a4..d9cb28e1c 100644 --- a/internal/providers/registry.go +++ b/internal/providers/registry.go @@ -2,14 +2,8 @@ package providers import ( - "context" "encoding/json" - "errors" "fmt" - "log/slog" - "maps" - "net/http" - "reflect" "slices" "sort" "strings" @@ -197,558 +191,8 @@ func (r *ModelRegistry) RegisterProviderWithNameAndType(provider core.Provider, r.providerRuntime[providerName] = state } -// Initialize fetches models from all registered providers and populates the registry. -// This should be called on application startup. -func (r *ModelRegistry) Initialize(ctx context.Context) error { - release, err := r.acquireRefresh(ctx) - if err != nil { - return err - } - defer release() - return r.initialize(ctx) -} - -func (r *ModelRegistry) initialize(ctx context.Context) error { - // Get a snapshot of providers with a read lock - r.mu.RLock() - providers := make([]core.Provider, len(r.providers)) - copy(providers, r.providers) - r.mu.RUnlock() - - // Build new model maps without holding the lock. - // This allows concurrent reads to continue using the existing map - // while we fetch models from providers (which may involve network calls). - newModels := make(map[string]*ModelInfo) - newModelsByProvider := make(map[string]map[string]*ModelInfo) - var totalModels int - var failedProviders int - runtimeUpdates := make(map[string]providerRuntimeState) - - r.mu.RLock() - providerTypes := make(map[core.Provider]string, len(r.providerTypes)) - providerNames := make(map[core.Provider]string, len(r.providerNames)) - maps.Copy(providerTypes, r.providerTypes) - maps.Copy(providerNames, r.providerNames) - r.mu.RUnlock() - configuredProviderModels, configuredProviderModelsMode := r.snapshotConfiguredProviderModels() - - for _, provider := range providers { - providerName := providerNames[provider] - if providerName == "" { - providerName = providerTypes[provider] - } - if providerName == "" { - providerName = fmt.Sprintf("%p", provider) - } - - configuredModels := configuredProviderModels[providerName] - resp, configuredReason, fetchAt, err := fetchProviderInventory( - ctx, - provider, - providerName, - providerTypes[provider], - configuredProviderModelsMode, - configuredModels, - ) - var configuredUpstreamError string - if configuredReason != configuredProviderModelsNotApplied { - attrs := []any{ - "provider", providerName, - "reason", string(configuredReason), - "configured_models", len(configuredModels), - } - if err != nil { - configuredUpstreamError = err.Error() - attrs = append(attrs, "error", err) - slog.Warn("upstream ListModels failed, using configured provider models", attrs...) - } else if configuredReason == configuredProviderModelsAllowlist { - slog.Debug("using configured provider models", attrs...) - } else { - slog.Warn("using configured provider models", attrs...) - } - err = nil - } - if err != nil { - slog.Warn("failed to fetch models from provider", - "provider", providerName, - "error", err, - ) - failedProviders++ - runtimeUpdates[providerName] = providerRuntimeState{ - registered: true, - lastModelFetchAt: fetchAt, - lastModelFetchError: err.Error(), - } - continue - } - - if resp == nil { - err := errors.New("provider returned nil model list") - slog.Warn("failed to fetch models from provider", - "provider", providerName, - "error", err, - ) - failedProviders++ - runtimeUpdates[providerName] = providerRuntimeState{ - registered: true, - lastModelFetchAt: fetchAt, - lastModelFetchError: err.Error(), - } - continue - } - - if len(resp.Data) == 0 { - err := errors.New("provider returned empty model list") - slog.Warn("provider returned empty model list", - "provider", providerName, - ) - runtimeUpdates[providerName] = providerRuntimeState{ - registered: true, - lastModelFetchAt: fetchAt, - lastModelFetchError: err.Error(), - } - if _, ok := newModelsByProvider[providerName]; !ok { - newModelsByProvider[providerName] = make(map[string]*ModelInfo) - } - continue - } - - runtimeUpdate := providerRuntimeState{ - registered: true, - lastModelFetchAt: fetchAt, - lastModelFetchError: configuredUpstreamError, - } - if configuredReason == configuredProviderModelsNotApplied { - runtimeUpdate.lastModelFetchSuccessAt = fetchAt - } - runtimeUpdates[providerName] = runtimeUpdate - - if _, ok := newModelsByProvider[providerName]; !ok { - newModelsByProvider[providerName] = make(map[string]*ModelInfo, len(resp.Data)) - } - - for _, model := range resp.Data { - info := &ModelInfo{ - Model: model, - Provider: provider, - ProviderName: providerName, - ProviderType: providerTypes[provider], - } - newModelsByProvider[providerName][model.ID] = info - - if _, exists := newModels[model.ID]; exists { - // Model already registered by another provider, skip - // First provider wins for unqualified lookups. - slog.Debug("model already registered, skipping", - "model", model.ID, - "provider", providerName, - "owner", model.OwnedBy, - ) - continue - } - - newModels[model.ID] = info - totalModels++ - } - } - - if totalModels == 0 { - r.applyProviderRuntimeUpdates(runtimeUpdates) - if failedProviders == len(providers) { - return fmt.Errorf("failed to fetch models from any provider") - } - return fmt.Errorf("no models available: providers returned empty model lists") - } - - // Enrich models with metadata from the model list (if loaded) - r.mu.RLock() - list := r.modelList - r.mu.RUnlock() - configOverrides := r.snapshotConfigOverrides() - metadataStats := metadataEnrichmentStats{} - if list != nil { - metadataStats = enrichProviderModelMaps(list, providerTypes, newModelsByProvider, nil) - } - metadataStats.Enriched += applyConfigMetadataOverrides(configOverrides, newModelsByProvider, nil) - - // Atomically swap the models map and invalidate sorted caches - r.mu.Lock() - r.models = newModels - r.modelsByProvider = newModelsByProvider - r.applyProviderRuntimeUpdatesLocked(runtimeUpdates) - r.invalidateSortedCaches() - r.mu.Unlock() - - // Mark as initialized - r.initMu.Lock() - r.initialized = true - r.initMu.Unlock() - - attrs := []any{ - "total_models", totalModels, - "providers", len(providers), - "failed_providers", failedProviders, - } - attrs = append(attrs, metadataStats.slogAttrs()...) - slog.Info("model registry initialized", attrs...) - - return nil -} - -func fetchProviderInventory( - ctx context.Context, - provider core.Provider, - providerName string, - providerType string, - mode config.ConfiguredProviderModelsMode, - configuredModels []string, -) (*core.ModelsResponse, configuredProviderModelsApplyReason, time.Time, error) { - fetchAt := time.Now().UTC() - if mode == config.ConfiguredProviderModelsModeAllowlist && len(configuredModels) > 0 { - resp, reason := applyConfiguredProviderModels( - providerName, - providerType, - mode, - configuredModels, - nil, - nil, - fetchAt.Unix(), - ) - return resp, reason, fetchAt, nil - } - - resp, err := provider.ListModels(ctx) - fetchAt = time.Now().UTC() - resp, reason := applyConfiguredProviderModels( - providerName, - providerType, - mode, - configuredModels, - resp, - err, - fetchAt.Unix(), - ) - return resp, reason, fetchAt, err -} - -func (r *ModelRegistry) applyProviderRuntimeUpdates(updates map[string]providerRuntimeState) { - if len(updates) == 0 { - return - } - - r.mu.Lock() - defer r.mu.Unlock() - - r.applyProviderRuntimeUpdatesLocked(updates) -} - -func (r *ModelRegistry) applyProviderRuntimeUpdatesLocked(updates map[string]providerRuntimeState) { - for providerName, update := range updates { - current := r.providerRuntime[providerName] - current.registered = update.registered || current.registered - if !update.lastModelFetchAt.IsZero() { - current.lastModelFetchAt = update.lastModelFetchAt - } - if !update.lastModelFetchSuccessAt.IsZero() { - current.lastModelFetchSuccessAt = update.lastModelFetchSuccessAt - if strings.TrimSpace(update.lastModelFetchError) == "" { - current.lastModelFetchError = "" - } - } - if strings.TrimSpace(update.lastModelFetchError) != "" { - current.lastModelFetchError = update.lastModelFetchError - } - r.providerRuntime[providerName] = current - } -} - -// Refresh updates the model registry by fetching fresh model lists from providers. -// This can be called periodically to keep the registry up to date. -func (r *ModelRegistry) Refresh(ctx context.Context) error { - return r.Initialize(ctx) -} - -func (r *ModelRegistry) acquireRefresh(ctx context.Context) (func(), error) { - if ctx == nil { - ctx = context.Background() - } - if err := ctx.Err(); err != nil { - return nil, registryRefreshAcquireError(err) - } - ch := r.refreshSemaphore() - select { - case ch <- struct{}{}: - return func() { <-ch }, nil - case <-ctx.Done(): - return nil, registryRefreshAcquireError(ctx.Err()) - } -} - -func (r *ModelRegistry) refreshSemaphore() chan struct{} { - r.refreshOnce.Do(func() { - if r.refreshCh == nil { - r.refreshCh = make(chan struct{}, 1) - } - }) - return r.refreshCh -} - -func registryRefreshAcquireError(err error) *core.GatewayError { - if errors.Is(err, context.DeadlineExceeded) { - return core.NewProviderError("model_registry", http.StatusGatewayTimeout, "model registry refresh timed out before start", err) - } - return core.NewProviderError("model_registry", http.StatusRequestTimeout, "model registry refresh canceled before start", err) -} - -// LoadFromCache loads the model list from the cache backend. -// Returns the number of models loaded and any error encountered. -func (r *ModelRegistry) LoadFromCache(ctx context.Context) (int, error) { - r.mu.RLock() - cacheBackend := r.cache - r.mu.RUnlock() - - if cacheBackend == nil { - return 0, nil - } - - modelCache, err := cacheBackend.Get(ctx) - if err != nil { - return 0, fmt.Errorf("failed to read cache: %w", err) - } - - if modelCache == nil { - return 0, nil // No cache yet, not an error - } - - // Build lookup maps from configured providers. - r.mu.RLock() - nameToProvider := make(map[string]core.Provider, len(r.providerNames)) - nameToProviderType := make(map[string]string, len(r.providerNames)) - providerOrderNames := make([]string, 0, len(r.providers)) - for _, provider := range r.providers { - providerName := r.providerNames[provider] - if providerName == "" { - continue - } - providerOrderNames = append(providerOrderNames, providerName) - } - for provider, pName := range r.providerNames { - nameToProvider[pName] = provider - nameToProviderType[pName] = r.providerTypes[provider] - } - r.mu.RUnlock() - - // Populate model maps from grouped cache structure. Unqualified lookups keep "first provider wins". - newModels := make(map[string]*ModelInfo) - newModelsByProvider := make(map[string]map[string]*ModelInfo) - cachedProviderTypes := make(map[string]string, len(modelCache.Providers)) - for providerName, cachedProv := range modelCache.Providers { - provider, ok := nameToProvider[providerName] - if !ok { - // Provider not configured, skip all its models - continue - } - cachedProviderTypes[providerName] = strings.TrimSpace(cachedProv.ProviderType) - providerType := strings.TrimSpace(nameToProviderType[providerName]) - if providerType == "" { - providerType = strings.TrimSpace(cachedProv.ProviderType) - } - providerModels := make(map[string]*ModelInfo, len(cachedProv.Models)) - for _, cached := range cachedProv.Models { - info := &ModelInfo{ - Model: core.Model{ - ID: cached.ID, - Object: "model", - OwnedBy: cachedProv.OwnedBy, - Created: cached.Created, - }, - Provider: provider, - ProviderName: providerName, - ProviderType: providerType, - } - providerModels[cached.ID] = info - if _, exists := newModels[cached.ID]; !exists { - newModels[cached.ID] = info - } - } - newModelsByProvider[providerName] = providerModels - } - - configuredProviderModels, configuredProviderModelsMode := r.snapshotConfiguredProviderModels() - if len(configuredProviderModels) > 0 { - for providerName, configuredModels := range configuredProviderModels { - provider, ok := nameToProvider[providerName] - if !ok { - continue - } - providerType := strings.TrimSpace(nameToProviderType[providerName]) - if providerType == "" { - providerType = strings.TrimSpace(cachedProviderTypes[providerName]) - } - providerModels := newModelsByProvider[providerName] - upstream := modelsResponseFromProviderMap(providerModels) - resp, reason := applyConfiguredProviderModels(providerName, providerType, configuredProviderModelsMode, configuredModels, upstream, nil, modelCache.UpdatedAt.Unix()) - if reason == configuredProviderModelsNotApplied { - continue - } - newModelsByProvider[providerName] = modelInfoMapFromResponse(resp, provider, providerName, providerType) - } - } - newModels = rebuildGlobalModelMap(newModelsByProvider, providerOrderNames) - - // Load model list data from cache if available - var list *modeldata.ModelList - if len(modelCache.ModelListData) > 0 { - parsed, parseErr := modeldata.Parse(modelCache.ModelListData) - if parseErr != nil { - slog.Warn("failed to parse cached model list data", "error", parseErr) - } else { - list = parsed - } - } - - // Enrich cached models with model list metadata - metadataStats := metadataEnrichmentStats{} - if list != nil { - metadataStats = enrichProviderModelMaps(list, r.snapshotProviderTypes(), newModelsByProvider, nil) - } - configOverrides := r.snapshotConfigOverrides() - metadataStats.Enriched += applyConfigMetadataOverrides(configOverrides, newModelsByProvider, nil) - - r.mu.Lock() - r.models = newModels - r.modelsByProvider = newModelsByProvider - r.invalidateSortedCaches() - if list != nil { - r.modelList = list - r.modelListRaw = modelCache.ModelListData - } - r.mu.Unlock() - - attrs := []any{ - "models", len(newModels), - "cache_updated_at", modelCache.UpdatedAt, - } - attrs = append(attrs, metadataStats.slogAttrs()...) - slog.Info("loaded models from cache", attrs...) - - return len(newModels), nil -} - -// SaveToCache saves the current model list to the cache backend. -func (r *ModelRegistry) SaveToCache(ctx context.Context) error { - r.mu.RLock() - cacheBackend := r.cache - modelsByProvider := make(map[string]map[string]*ModelInfo, len(r.modelsByProvider)) - for providerName, models := range r.modelsByProvider { - modelsByProvider[providerName] = make(map[string]*ModelInfo, len(models)) - maps.Copy(modelsByProvider[providerName], models) - } - providerTypes := make(map[core.Provider]string, len(r.providerTypes)) - maps.Copy(providerTypes, r.providerTypes) - modelListRaw := r.modelListRaw - r.mu.RUnlock() - if cacheBackend == nil { - return nil - } - - mc := &modelcache.ModelCache{ - UpdatedAt: time.Now().UTC(), - Providers: make(map[string]modelcache.CachedProvider, len(modelsByProvider)), - ModelListData: modelListRaw, - } - - var totalModels int - for providerName, models := range modelsByProvider { - // Determine provider type and owned_by from any model in this provider group. - var pType, ownedBy string - for _, info := range models { - if ownedBy == "" { - ownedBy = info.Model.OwnedBy - } - if pType == "" { - pType = strings.TrimSpace(info.ProviderType) - if pType == "" { - pType = strings.TrimSpace(providerTypes[info.Provider]) - } - } - if pType != "" && ownedBy != "" { - break - } - } - if pType == "" { - // No known provider type for this provider, skip entirely. - continue - } - modelIDs := make([]string, 0, len(models)) - for modelID := range models { - modelIDs = append(modelIDs, modelID) - } - sort.Strings(modelIDs) - - cachedModels := make([]modelcache.CachedModel, 0, len(modelIDs)) - for _, modelID := range modelIDs { - info := models[modelID] - cachedModels = append(cachedModels, modelcache.CachedModel{ - ID: modelID, - Created: info.Model.Created, - }) - } - mc.Providers[providerName] = modelcache.CachedProvider{ - ProviderType: pType, - OwnedBy: ownedBy, - Models: cachedModels, - } - totalModels += len(cachedModels) - } - - if err := cacheBackend.Set(ctx, mc); err != nil { - return fmt.Errorf("failed to save cache: %w", err) - } - - slog.Debug("saved models to cache", "models", totalModels) - return nil -} - -// InitializeAsync starts model fetching in a background goroutine. -// It first loads any cached models for immediate availability, then refreshes from network. -// Returns immediately after loading cache. The background goroutine will update models -// and save to cache when network fetch completes. -func (r *ModelRegistry) InitializeAsync(ctx context.Context) { - // First, try to load from cache for instant startup - cached, err := r.LoadFromCache(ctx) - if err != nil { - slog.Warn("failed to load models from cache", "error", err) - } else if cached > 0 { - slog.Info("serving traffic with cached models while refreshing", "cached_models", cached) - } - - // Start background initialization - go func() { - initCtx, cancel := context.WithTimeout(context.Background(), 60*time.Second) - defer cancel() - - if err := r.Initialize(initCtx); err != nil { - slog.Warn("background model initialization failed", "error", err) - return - } - - // Save to cache for next startup - if err := r.SaveToCache(initCtx); err != nil { - slog.Warn("failed to save models to cache", "error", err) - } - }() -} - -// IsInitialized returns true if at least one successful network fetch has completed. -// This can be used to check if the registry has fresh data or is only serving from cache. -func (r *ModelRegistry) IsInitialized() bool { - r.initMu.Lock() - defer r.initMu.Unlock() - return r.initialized -} // GetProvider returns the provider for the given model, or nil if not found func (r *ModelRegistry) GetProvider(model string) core.Provider { @@ -1360,543 +804,3 @@ func (r *ModelRegistry) ProviderRuntimeSnapshots() []ProviderRuntimeSnapshot { return result } - -// SetModelList stores the parsed model list and its raw bytes for cache persistence. -func (r *ModelRegistry) SetModelList(list *modeldata.ModelList, raw json.RawMessage) { - r.mu.Lock() - defer r.mu.Unlock() - r.modelList = list - r.modelListRaw = raw -} - -// EnrichModels re-applies model list metadata to all currently registered models. -// Call this after SetModelList to update existing models with the new metadata. -// Holds the write lock for the entire operation and replaces published ModelInfo -// entries instead of mutating them in place so concurrent readers can safely keep -// using older snapshots after unlocking. -func (r *ModelRegistry) EnrichModels() { - _ = r.enrichModels() -} - -func (r *ModelRegistry) enrichModels() metadataEnrichmentStats { - r.mu.Lock() - defer r.mu.Unlock() - return r.enrichModelsLocked() -} - -func (r *ModelRegistry) enrichModelsLocked() metadataEnrichmentStats { - if len(r.models) == 0 { - return metadataEnrichmentStats{} - } - if r.modelList == nil && len(r.configMetadataOverrides) == 0 { - return metadataEnrichmentStats{} - } - - providerTypes := make(map[core.Provider]string, len(r.providerTypes)) - maps.Copy(providerTypes, r.providerTypes) - - replacements := make(map[*ModelInfo]*ModelInfo, len(r.models)) - stats := metadataEnrichmentStats{} - if r.modelList != nil { - stats = enrichProviderModelMaps(r.modelList, providerTypes, r.modelsByProvider, replacements) - } - stats.Enriched += applyConfigMetadataOverrides(r.configMetadataOverrides, r.modelsByProvider, replacements) - for modelID, info := range r.models { - if replacement, ok := replacements[info]; ok { - r.models[modelID] = replacement - } - } - r.invalidateSortedCaches() - return stats -} - -func (r *ModelRegistry) setModelListAndEnrich(list *modeldata.ModelList, raw json.RawMessage) metadataEnrichmentStats { - r.mu.Lock() - defer r.mu.Unlock() - r.modelList = list - r.modelListRaw = raw - return r.enrichModelsLocked() -} - -// ResolveMetadata resolves metadata for a model directly via the stored model list, -// bypassing the registry key lookup. This handles cases where the usage DB stores -// a response model ID (e.g., "gpt-4o-2024-08-06") that differs from the registry -// key (e.g., "gpt-4o") by using the reverse index in the model list. -func (r *ModelRegistry) ResolveMetadata(providerType, modelID string) *core.ModelMetadata { - r.mu.RLock() - defer r.mu.RUnlock() - if r.modelList == nil { - return nil - } - return modeldata.Resolve(r.modelList, providerType, modelID) -} - -// GetModelMetadata returns the metadata for a model, or nil if not found or not enriched. -func (r *ModelRegistry) GetModelMetadata(modelID string) *core.ModelMetadata { - r.mu.RLock() - defer r.mu.RUnlock() - if info, ok := r.models[modelID]; ok { - return info.Model.Metadata - } - return nil -} - -// ResolvePricing returns the pricing metadata for a model, trying the registry first -// and falling back to a reverse-index lookup via the model list. -// Returns nil if no pricing is available. -func (r *ModelRegistry) ResolvePricing(model, providerType string) *core.ModelPricing { - providerSelector := strings.TrimSpace(providerType) - if meta := r.getProviderModelMetadata(providerSelector, model); meta != nil && meta.Pricing != nil { - return meta.Pricing - } - - meta := r.GetModelMetadata(model) - if meta != nil && meta.Pricing != nil { - return meta.Pricing - } - if providerSelector != "" { - meta = r.ResolveMetadata(r.metadataProviderType(providerSelector), r.metadataModelID(model)) - if meta != nil && meta.Pricing != nil { - return meta.Pricing - } - } - return nil -} - -func (r *ModelRegistry) getProviderModelMetadata(providerSelector, model string) *core.ModelMetadata { - providerSelector = strings.TrimSpace(providerSelector) - model = strings.TrimSpace(model) - if model == "" { - return nil - } - - modelProviderName, modelID := splitModelSelector(model) - r.mu.RLock() - defer r.mu.RUnlock() - - if modelProviderName != "" { - if meta := metadataFromProviderModel(r.modelsByProvider[modelProviderName], modelID); meta != nil { - return meta - } - if r.hasConfiguredProviderNameLocked(modelProviderName) { - return nil - } - } - - if providerSelector == "" { - return nil - } - if meta := metadataFromProviderModel(r.modelsByProvider[providerSelector], model); meta != nil { - return meta - } - return nil -} - -func metadataFromProviderModel(providerModels map[string]*ModelInfo, model string) *core.ModelMetadata { - if len(providerModels) == 0 { - return nil - } - info := providerModels[strings.TrimSpace(model)] - if info == nil { - return nil - } - return info.Model.Metadata -} - -func (r *ModelRegistry) metadataProviderType(providerSelector string) string { - providerSelector = strings.TrimSpace(providerSelector) - if providerSelector == "" { - return "" - } - if providerType := r.GetProviderTypeForName(providerSelector); providerType != "" { - return providerType - } - return providerSelector -} - -func (r *ModelRegistry) metadataModelID(model string) string { - model = strings.TrimSpace(model) - providerName, modelID := splitModelSelector(model) - if providerName == "" { - return model - } - r.mu.RLock() - defer r.mu.RUnlock() - if r.hasConfiguredProviderNameLocked(providerName) { - return modelID - } - return model -} - -// snapshotProviderTypes returns a copy of the providerTypes map for use outside the lock. -func (r *ModelRegistry) snapshotProviderTypes() map[core.Provider]string { - r.mu.RLock() - defer r.mu.RUnlock() - m := make(map[core.Provider]string, len(r.providerTypes)) - maps.Copy(m, r.providerTypes) - return m -} - -// snapshotConfigOverrides returns a copy of the configMetadataOverrides outer -// and inner maps for use outside the lock. The inner *core.ModelMetadata -// pointers are shared, which is safe because SetProviderMetadataOverrides -// deep-clones on insertion and the registry never hands those values back out. -func (r *ModelRegistry) snapshotConfigOverrides() map[string]map[string]*core.ModelMetadata { - r.mu.RLock() - defer r.mu.RUnlock() - if len(r.configMetadataOverrides) == 0 { - return nil - } - out := make(map[string]map[string]*core.ModelMetadata, len(r.configMetadataOverrides)) - for provider, inner := range r.configMetadataOverrides { - innerCopy := make(map[string]*core.ModelMetadata, len(inner)) - for modelID, meta := range inner { - innerCopy[modelID] = meta - } - out[provider] = innerCopy - } - return out -} - -func (r *ModelRegistry) snapshotConfiguredProviderModels() (map[string][]string, config.ConfiguredProviderModelsMode) { - r.mu.RLock() - defer r.mu.RUnlock() - mode := config.ResolveConfiguredProviderModelsMode(r.configuredProviderModelsMode) - if len(r.configuredProviderModels) == 0 { - return nil, mode - } - out := make(map[string][]string, len(r.configuredProviderModels)) - for provider, models := range r.configuredProviderModels { - out[provider] = slices.Clone(models) - } - return out, mode -} - -// collectionEmpty reports whether a reflect.Value representing a slice, array, -// or map has no elements (covering both nil and non-nil-but-zero-length), and -// falls back to reflect.Value.IsZero for other kinds. This lets override- -// emptiness checks treat `modes: []` the same as an omitted field, which -// IsZero alone would not. -func collectionEmpty(v reflect.Value) bool { - switch v.Kind() { - case reflect.Slice, reflect.Map, reflect.Array: - return v.Len() == 0 - } - return v.IsZero() -} - -// structFieldsEmpty returns true if every field of the given struct value -// passes collectionEmpty. -func structFieldsEmpty(v reflect.Value) bool { - for i := 0; i < v.NumField(); i++ { - if !collectionEmpty(v.Field(i)) { - return false - } - } - return true -} - -// metadataOverrideEmpty reports whether an override has no effective content. -// An empty override (either nil or zero-valued on every field) would turn a -// nil current metadata into a non-nil empty struct after MergeMetadata, so -// callers should short-circuit on it. Uses reflect-based field inspection so -// new fields on core.ModelMetadata are picked up automatically; Pricing is -// handled separately so a non-nil pointer to an empty pricing block still -// counts as empty. -func metadataOverrideEmpty(m *core.ModelMetadata) bool { - if m == nil { - return true - } - if !pricingOverrideEmpty(m.Pricing) { - return false - } - tmp := *m - tmp.Pricing = nil - return structFieldsEmpty(reflect.ValueOf(tmp)) -} - -// pricingOverrideEmpty reports whether a pricing override has no effective -// content — nil or every field at its zero value (with collections treated as -// empty when length==0). -func pricingOverrideEmpty(p *core.ModelPricing) bool { - if p == nil { - return true - } - return structFieldsEmpty(reflect.ValueOf(*p)) -} - -// applyConfigMetadataOverrides layers operator-declared metadata onto already- -// enriched models. Call it after enrichProviderModelMaps with the same -// replacements map (pass nil replacements for fresh, unpublished maps). -// Returns the number of models whose metadata was updated. -func applyConfigMetadataOverrides( - overrides map[string]map[string]*core.ModelMetadata, - modelsByProvider map[string]map[string]*ModelInfo, - replacements map[*ModelInfo]*ModelInfo, -) int { - if len(overrides) == 0 { - return 0 - } - // reverse lets us find the pre-enrichment pointer when an entry has - // already been replaced by enrichProviderModelMaps, so our replacement - // chain stays consistent from the caller's perspective. Always allocated - // when replacements is non-nil so the else-branch write below cannot hit - // a nil map when enrichment made no replacements. - var reverse map[*ModelInfo]*ModelInfo - if replacements != nil { - reverse = make(map[*ModelInfo]*ModelInfo, len(replacements)) - for orig, repl := range replacements { - reverse[repl] = orig - } - } - applied := 0 - for providerName, modelOverrides := range overrides { - providerModels, ok := modelsByProvider[providerName] - if !ok { - continue - } - for modelID, override := range modelOverrides { - if metadataOverrideEmpty(override) { - // A nil or effectively-empty override has nothing to - // contribute. Skipping avoids turning a nil current metadata - // into a non-nil empty struct, which the DeepEqual check - // below would not catch. - continue - } - current, ok := providerModels[modelID] - if !ok { - continue - } - merged := modeldata.MergeMetadata(current.Model.Metadata, override) - // Skip no-op merges so concurrent readers holding the current - // pointer keep a stable view when the override adds no new info. - if reflect.DeepEqual(current.Model.Metadata, merged) { - continue - } - if replacements == nil { - current.Model.Metadata = merged - applied++ - continue - } - cloned := *current - cloned.Model.Metadata = merged - next := &cloned - providerModels[modelID] = next - if orig, hasOrig := reverse[current]; hasOrig { - replacements[orig] = next - reverse[next] = orig - } else { - replacements[current] = next - reverse[next] = current - } - applied++ - } - } - return applied -} - -func enrichProviderModelMaps( - list *modeldata.ModelList, - providerTypes map[core.Provider]string, - modelsByProvider map[string]map[string]*ModelInfo, - replacements map[*ModelInfo]*ModelInfo, -) metadataEnrichmentStats { - if list == nil { - return metadataEnrichmentStats{} - } - stats := metadataEnrichmentStats{} - for _, providerModels := range modelsByProvider { - if len(providerModels) == 0 { - continue - } - stats.Providers++ - accessor := ®istryAccessor{ - models: providerModels, - providerTypes: providerTypes, - replacements: replacements, - } - enrichStats := modeldata.Enrich(accessor, list) - stats.Enriched += enrichStats.Enriched - stats.Total += enrichStats.Total - } - return stats -} - -// registryAccessor implements modeldata.ModelInfoAccessor. -// The models map may be either an unpublished snapshot (Initialize, LoadFromCache) -// or the live registry map (EnrichModels, which uses replacements to preserve -// immutability of already-published ModelInfo values). -type registryAccessor struct { - models map[string]*ModelInfo - providerTypes map[core.Provider]string - replacements map[*ModelInfo]*ModelInfo -} - -func (a *registryAccessor) ModelIDs() []string { - ids := make([]string, 0, len(a.models)) - for id := range a.models { - ids = append(ids, id) - } - return ids -} - -func (a *registryAccessor) GetProviderType(modelID string) string { - info, ok := a.models[modelID] - if !ok { - return "" - } - if providerType := strings.TrimSpace(info.ProviderType); providerType != "" { - return providerType - } - return strings.TrimSpace(a.providerTypes[info.Provider]) -} - -func (a *registryAccessor) SetMetadata(modelID string, meta *core.ModelMetadata) { - if info, ok := a.models[modelID]; ok { - if a.replacements != nil { - cloned := *info - cloned.Model.Metadata = meta - replacement := &cloned - a.models[modelID] = replacement - a.replacements[info] = replacement - return - } - info.Model.Metadata = meta - } -} - -// StartBackgroundRefresh starts a goroutine that periodically refreshes the model registry. -// If modelListURL is non-empty, the model list is also re-fetched on each tick. -// The returned stop function is blocking: it cancels the refresh loop and waits -// for the goroutine to exit before returning, so callers should expect it to -// block during shutdown until any in-flight refresh work unwinds. -func (r *ModelRegistry) StartBackgroundRefresh(interval time.Duration, modelListURL string) func() { - ctx, cancel := context.WithCancel(context.Background()) - done := make(chan struct{}) - var stopOnce sync.Once - - go func() { - defer close(done) - ticker := time.NewTicker(interval) - defer ticker.Stop() - - for { - select { - case <-ctx.Done(): - return - case <-ticker.C: - refreshCtx, refreshCancel := context.WithTimeout(ctx, 30*time.Second) - err := r.Initialize(refreshCtx) - refreshCancel() - if err != nil { - if !isBenignBackgroundRefreshError(ctx, err) { - slog.Warn("background model refresh failed", "error", err) - } - } else { - func() { - cacheCtx, cacheCancel := context.WithTimeout(ctx, 10*time.Second) - defer cacheCancel() - if err := r.SaveToCache(cacheCtx); err != nil { - if !isBenignBackgroundRefreshError(ctx, err) { - slog.Warn("failed to save models to cache after refresh", "error", err) - } - } - }() - } - - // Also refresh model list if configured - if modelListURL != "" { - r.refreshModelList(ctx, modelListURL) - } - } - } - }() - - return func() { - stopOnce.Do(func() { - cancel() - <-done - }) - } -} - -// RefreshModelList fetches the external model metadata list and re-enriches all -// currently registered models. It does not persist the model cache; callers that -// want durable startup data should call SaveToCache after this succeeds. -func (r *ModelRegistry) RefreshModelList(ctx context.Context, url string) (int, error) { - if strings.TrimSpace(url) == "" { - return 0, nil - } - - release, err := r.acquireRefresh(ctx) - if err != nil { - return 0, err - } - defer release() - - models, _, err := r.refreshModelListLocked(ctx, url) - return models, err -} - -func (r *ModelRegistry) refreshModelListLocked(ctx context.Context, url string) (int, metadataEnrichmentStats, error) { - list, raw, err := modeldata.Fetch(ctx, url) - if err != nil { - return 0, metadataEnrichmentStats{}, err - } - if list == nil { - return 0, metadataEnrichmentStats{}, nil - } - - metadataStats := r.setModelListAndEnrich(list, raw) - return len(list.Models), metadataStats, nil -} - -// refreshModelList fetches the model list and re-enriches all models. -func (r *ModelRegistry) refreshModelList(ctx context.Context, url string) { - fetchCtx, cancel := context.WithTimeout(ctx, 45*time.Second) - defer cancel() - - release, err := r.acquireRefresh(fetchCtx) - if err != nil { - if !isBenignBackgroundRefreshError(ctx, err) { - slog.Warn("failed to acquire model list refresh", "url", url, "error", err) - } - return - } - var ( - models int - metadataStats metadataEnrichmentStats - ) - func() { - defer release() - models, metadataStats, err = r.refreshModelListLocked(fetchCtx, url) - }() - if err != nil { - if !isBenignBackgroundRefreshError(ctx, err) { - slog.Warn("failed to refresh model list", "url", url, "error", err) - } - return - } - if models == 0 { - return - } - - if err := r.SaveToCache(fetchCtx); err != nil { - if !isBenignBackgroundRefreshError(ctx, err) { - slog.Warn("failed to save cache after model list refresh", "error", err) - } - } - attrs := []any{"models", models} - attrs = append(attrs, metadataStats.slogAttrs()...) - slog.Debug("model list refreshed", attrs...) -} - -func isBenignBackgroundRefreshError(parent context.Context, err error) bool { - if err == nil { - return true - } - if parent == nil || parent.Err() == nil { - return false - } - return errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) -} diff --git a/internal/providers/registry_cache.go b/internal/providers/registry_cache.go new file mode 100644 index 000000000..b4c0537fe --- /dev/null +++ b/internal/providers/registry_cache.go @@ -0,0 +1,227 @@ +package providers + +import ( + "context" + "fmt" + "log/slog" + "maps" + "sort" + "strings" + "time" + + "gomodel/internal/cache/modelcache" + "gomodel/internal/core" + "gomodel/internal/modeldata" +) + +// LoadFromCache loads the model list from the cache backend. +// Returns the number of models loaded and any error encountered. +func (r *ModelRegistry) LoadFromCache(ctx context.Context) (int, error) { + r.mu.RLock() + cacheBackend := r.cache + r.mu.RUnlock() + + if cacheBackend == nil { + return 0, nil + } + + modelCache, err := cacheBackend.Get(ctx) + if err != nil { + return 0, fmt.Errorf("failed to read cache: %w", err) + } + + if modelCache == nil { + return 0, nil // No cache yet, not an error + } + + // Build lookup maps from configured providers. + r.mu.RLock() + nameToProvider := make(map[string]core.Provider, len(r.providerNames)) + nameToProviderType := make(map[string]string, len(r.providerNames)) + providerOrderNames := make([]string, 0, len(r.providers)) + for _, provider := range r.providers { + providerName := r.providerNames[provider] + if providerName == "" { + continue + } + providerOrderNames = append(providerOrderNames, providerName) + } + for provider, pName := range r.providerNames { + nameToProvider[pName] = provider + nameToProviderType[pName] = r.providerTypes[provider] + } + r.mu.RUnlock() + + // Populate model maps from grouped cache structure. Unqualified lookups keep "first provider wins". + newModels := make(map[string]*ModelInfo) + newModelsByProvider := make(map[string]map[string]*ModelInfo) + cachedProviderTypes := make(map[string]string, len(modelCache.Providers)) + for providerName, cachedProv := range modelCache.Providers { + provider, ok := nameToProvider[providerName] + if !ok { + // Provider not configured, skip all its models + continue + } + cachedProviderTypes[providerName] = strings.TrimSpace(cachedProv.ProviderType) + providerType := strings.TrimSpace(nameToProviderType[providerName]) + if providerType == "" { + providerType = strings.TrimSpace(cachedProv.ProviderType) + } + providerModels := make(map[string]*ModelInfo, len(cachedProv.Models)) + for _, cached := range cachedProv.Models { + info := &ModelInfo{ + Model: core.Model{ + ID: cached.ID, + Object: "model", + OwnedBy: cachedProv.OwnedBy, + Created: cached.Created, + }, + Provider: provider, + ProviderName: providerName, + ProviderType: providerType, + } + providerModels[cached.ID] = info + if _, exists := newModels[cached.ID]; !exists { + newModels[cached.ID] = info + } + } + newModelsByProvider[providerName] = providerModels + } + + configuredProviderModels, configuredProviderModelsMode := r.snapshotConfiguredProviderModels() + if len(configuredProviderModels) > 0 { + for providerName, configuredModels := range configuredProviderModels { + provider, ok := nameToProvider[providerName] + if !ok { + continue + } + providerType := strings.TrimSpace(nameToProviderType[providerName]) + if providerType == "" { + providerType = strings.TrimSpace(cachedProviderTypes[providerName]) + } + providerModels := newModelsByProvider[providerName] + upstream := modelsResponseFromProviderMap(providerModels) + resp, reason := applyConfiguredProviderModels(providerName, providerType, configuredProviderModelsMode, configuredModels, upstream, nil, modelCache.UpdatedAt.Unix()) + if reason == configuredProviderModelsNotApplied { + continue + } + newModelsByProvider[providerName] = modelInfoMapFromResponse(resp, provider, providerName, providerType) + } + } + newModels = rebuildGlobalModelMap(newModelsByProvider, providerOrderNames) + + // Load model list data from cache if available + var list *modeldata.ModelList + if len(modelCache.ModelListData) > 0 { + parsed, parseErr := modeldata.Parse(modelCache.ModelListData) + if parseErr != nil { + slog.Warn("failed to parse cached model list data", "error", parseErr) + } else { + list = parsed + } + } + + // Enrich cached models with model list metadata + metadataStats := metadataEnrichmentStats{} + if list != nil { + metadataStats = enrichProviderModelMaps(list, r.snapshotProviderTypes(), newModelsByProvider, nil) + } + configOverrides := r.snapshotConfigOverrides() + metadataStats.Enriched += applyConfigMetadataOverrides(configOverrides, newModelsByProvider, nil) + + r.mu.Lock() + r.models = newModels + r.modelsByProvider = newModelsByProvider + r.invalidateSortedCaches() + if list != nil { + r.modelList = list + r.modelListRaw = modelCache.ModelListData + } + r.mu.Unlock() + + attrs := []any{ + "models", len(newModels), + "cache_updated_at", modelCache.UpdatedAt, + } + attrs = append(attrs, metadataStats.slogAttrs()...) + slog.Info("loaded models from cache", attrs...) + + return len(newModels), nil +} + +// SaveToCache saves the current model list to the cache backend. +func (r *ModelRegistry) SaveToCache(ctx context.Context) error { + r.mu.RLock() + cacheBackend := r.cache + modelsByProvider := make(map[string]map[string]*ModelInfo, len(r.modelsByProvider)) + for providerName, models := range r.modelsByProvider { + modelsByProvider[providerName] = make(map[string]*ModelInfo, len(models)) + maps.Copy(modelsByProvider[providerName], models) + } + providerTypes := make(map[core.Provider]string, len(r.providerTypes)) + maps.Copy(providerTypes, r.providerTypes) + modelListRaw := r.modelListRaw + r.mu.RUnlock() + + if cacheBackend == nil { + return nil + } + + mc := &modelcache.ModelCache{ + UpdatedAt: time.Now().UTC(), + Providers: make(map[string]modelcache.CachedProvider, len(modelsByProvider)), + ModelListData: modelListRaw, + } + + var totalModels int + for providerName, models := range modelsByProvider { + // Determine provider type and owned_by from any model in this provider group. + var pType, ownedBy string + for _, info := range models { + if ownedBy == "" { + ownedBy = info.Model.OwnedBy + } + if pType == "" { + pType = strings.TrimSpace(info.ProviderType) + if pType == "" { + pType = strings.TrimSpace(providerTypes[info.Provider]) + } + } + if pType != "" && ownedBy != "" { + break + } + } + if pType == "" { + // No known provider type for this provider, skip entirely. + continue + } + + modelIDs := make([]string, 0, len(models)) + for modelID := range models { + modelIDs = append(modelIDs, modelID) + } + sort.Strings(modelIDs) + + cachedModels := make([]modelcache.CachedModel, 0, len(modelIDs)) + for _, modelID := range modelIDs { + info := models[modelID] + cachedModels = append(cachedModels, modelcache.CachedModel{ + ID: modelID, + Created: info.Model.Created, + }) + } + mc.Providers[providerName] = modelcache.CachedProvider{ + ProviderType: pType, + OwnedBy: ownedBy, + Models: cachedModels, + } + totalModels += len(cachedModels) + } + + if err := cacheBackend.Set(ctx, mc); err != nil { + return fmt.Errorf("failed to save cache: %w", err) + } + + slog.Debug("saved models to cache", "models", totalModels) + return nil +} diff --git a/internal/providers/registry_init.go b/internal/providers/registry_init.go new file mode 100644 index 000000000..55a28788e --- /dev/null +++ b/internal/providers/registry_init.go @@ -0,0 +1,494 @@ +package providers + +import ( + "context" + "errors" + "fmt" + "log/slog" + "maps" + "net/http" + "strings" + "sync" + "time" + + "gomodel/config" + "gomodel/internal/core" + "gomodel/internal/modeldata" +) + +// Initialize fetches models from all registered providers and populates the registry. +// This should be called on application startup. +func (r *ModelRegistry) Initialize(ctx context.Context) error { + release, err := r.acquireRefresh(ctx) + if err != nil { + return err + } + defer release() + return r.initialize(ctx) +} + +func (r *ModelRegistry) initialize(ctx context.Context) error { + // Get a snapshot of providers with a read lock + r.mu.RLock() + providers := make([]core.Provider, len(r.providers)) + copy(providers, r.providers) + r.mu.RUnlock() + + // Build new model maps without holding the lock. + // This allows concurrent reads to continue using the existing map + // while we fetch models from providers (which may involve network calls). + newModels := make(map[string]*ModelInfo) + newModelsByProvider := make(map[string]map[string]*ModelInfo) + var totalModels int + var failedProviders int + runtimeUpdates := make(map[string]providerRuntimeState) + + r.mu.RLock() + providerTypes := make(map[core.Provider]string, len(r.providerTypes)) + providerNames := make(map[core.Provider]string, len(r.providerNames)) + maps.Copy(providerTypes, r.providerTypes) + maps.Copy(providerNames, r.providerNames) + r.mu.RUnlock() + configuredProviderModels, configuredProviderModelsMode := r.snapshotConfiguredProviderModels() + + for _, provider := range providers { + providerName := providerNames[provider] + if providerName == "" { + providerName = providerTypes[provider] + } + if providerName == "" { + providerName = fmt.Sprintf("%p", provider) + } + + configuredModels := configuredProviderModels[providerName] + resp, configuredReason, fetchAt, err := fetchProviderInventory( + ctx, + provider, + providerName, + providerTypes[provider], + configuredProviderModelsMode, + configuredModels, + ) + var configuredUpstreamError string + if configuredReason != configuredProviderModelsNotApplied { + attrs := []any{ + "provider", providerName, + "reason", string(configuredReason), + "configured_models", len(configuredModels), + } + if err != nil { + configuredUpstreamError = err.Error() + attrs = append(attrs, "error", err) + slog.Warn("upstream ListModels failed, using configured provider models", attrs...) + } else if configuredReason == configuredProviderModelsAllowlist { + slog.Debug("using configured provider models", attrs...) + } else { + slog.Warn("using configured provider models", attrs...) + } + err = nil + } + if err != nil { + slog.Warn("failed to fetch models from provider", + "provider", providerName, + "error", err, + ) + failedProviders++ + runtimeUpdates[providerName] = providerRuntimeState{ + registered: true, + lastModelFetchAt: fetchAt, + lastModelFetchError: err.Error(), + } + continue + } + + if resp == nil { + err := errors.New("provider returned nil model list") + slog.Warn("failed to fetch models from provider", + "provider", providerName, + "error", err, + ) + failedProviders++ + runtimeUpdates[providerName] = providerRuntimeState{ + registered: true, + lastModelFetchAt: fetchAt, + lastModelFetchError: err.Error(), + } + continue + } + + if len(resp.Data) == 0 { + err := errors.New("provider returned empty model list") + slog.Warn("provider returned empty model list", + "provider", providerName, + ) + runtimeUpdates[providerName] = providerRuntimeState{ + registered: true, + lastModelFetchAt: fetchAt, + lastModelFetchError: err.Error(), + } + if _, ok := newModelsByProvider[providerName]; !ok { + newModelsByProvider[providerName] = make(map[string]*ModelInfo) + } + continue + } + + runtimeUpdate := providerRuntimeState{ + registered: true, + lastModelFetchAt: fetchAt, + lastModelFetchError: configuredUpstreamError, + } + if configuredReason == configuredProviderModelsNotApplied { + runtimeUpdate.lastModelFetchSuccessAt = fetchAt + } + runtimeUpdates[providerName] = runtimeUpdate + + if _, ok := newModelsByProvider[providerName]; !ok { + newModelsByProvider[providerName] = make(map[string]*ModelInfo, len(resp.Data)) + } + + for _, model := range resp.Data { + info := &ModelInfo{ + Model: model, + Provider: provider, + ProviderName: providerName, + ProviderType: providerTypes[provider], + } + newModelsByProvider[providerName][model.ID] = info + + if _, exists := newModels[model.ID]; exists { + // Model already registered by another provider, skip + // First provider wins for unqualified lookups. + slog.Debug("model already registered, skipping", + "model", model.ID, + "provider", providerName, + "owner", model.OwnedBy, + ) + continue + } + + newModels[model.ID] = info + totalModels++ + } + } + + if totalModels == 0 { + r.applyProviderRuntimeUpdates(runtimeUpdates) + if failedProviders == len(providers) { + return fmt.Errorf("failed to fetch models from any provider") + } + return fmt.Errorf("no models available: providers returned empty model lists") + } + + // Enrich models with metadata from the model list (if loaded) + r.mu.RLock() + list := r.modelList + r.mu.RUnlock() + configOverrides := r.snapshotConfigOverrides() + metadataStats := metadataEnrichmentStats{} + if list != nil { + metadataStats = enrichProviderModelMaps(list, providerTypes, newModelsByProvider, nil) + } + metadataStats.Enriched += applyConfigMetadataOverrides(configOverrides, newModelsByProvider, nil) + + // Atomically swap the models map and invalidate sorted caches + r.mu.Lock() + r.models = newModels + r.modelsByProvider = newModelsByProvider + r.applyProviderRuntimeUpdatesLocked(runtimeUpdates) + r.invalidateSortedCaches() + r.mu.Unlock() + + // Mark as initialized + r.initMu.Lock() + r.initialized = true + r.initMu.Unlock() + + attrs := []any{ + "total_models", totalModels, + "providers", len(providers), + "failed_providers", failedProviders, + } + attrs = append(attrs, metadataStats.slogAttrs()...) + slog.Info("model registry initialized", attrs...) + + return nil +} + +func fetchProviderInventory( + ctx context.Context, + provider core.Provider, + providerName string, + providerType string, + mode config.ConfiguredProviderModelsMode, + configuredModels []string, +) (*core.ModelsResponse, configuredProviderModelsApplyReason, time.Time, error) { + fetchAt := time.Now().UTC() + if mode == config.ConfiguredProviderModelsModeAllowlist && len(configuredModels) > 0 { + resp, reason := applyConfiguredProviderModels( + providerName, + providerType, + mode, + configuredModels, + nil, + nil, + fetchAt.Unix(), + ) + return resp, reason, fetchAt, nil + } + + resp, err := provider.ListModels(ctx) + fetchAt = time.Now().UTC() + resp, reason := applyConfiguredProviderModels( + providerName, + providerType, + mode, + configuredModels, + resp, + err, + fetchAt.Unix(), + ) + return resp, reason, fetchAt, err +} + +func (r *ModelRegistry) applyProviderRuntimeUpdates(updates map[string]providerRuntimeState) { + if len(updates) == 0 { + return + } + + r.mu.Lock() + defer r.mu.Unlock() + + r.applyProviderRuntimeUpdatesLocked(updates) +} + +func (r *ModelRegistry) applyProviderRuntimeUpdatesLocked(updates map[string]providerRuntimeState) { + for providerName, update := range updates { + current := r.providerRuntime[providerName] + current.registered = update.registered || current.registered + if !update.lastModelFetchAt.IsZero() { + current.lastModelFetchAt = update.lastModelFetchAt + } + if !update.lastModelFetchSuccessAt.IsZero() { + current.lastModelFetchSuccessAt = update.lastModelFetchSuccessAt + if strings.TrimSpace(update.lastModelFetchError) == "" { + current.lastModelFetchError = "" + } + } + if strings.TrimSpace(update.lastModelFetchError) != "" { + current.lastModelFetchError = update.lastModelFetchError + } + r.providerRuntime[providerName] = current + } +} + +// Refresh updates the model registry by fetching fresh model lists from providers. +// This can be called periodically to keep the registry up to date. +func (r *ModelRegistry) Refresh(ctx context.Context) error { + return r.Initialize(ctx) +} + +func (r *ModelRegistry) acquireRefresh(ctx context.Context) (func(), error) { + if ctx == nil { + ctx = context.Background() + } + if err := ctx.Err(); err != nil { + return nil, registryRefreshAcquireError(err) + } + ch := r.refreshSemaphore() + select { + case ch <- struct{}{}: + return func() { <-ch }, nil + case <-ctx.Done(): + return nil, registryRefreshAcquireError(ctx.Err()) + } +} + +func (r *ModelRegistry) refreshSemaphore() chan struct{} { + r.refreshOnce.Do(func() { + if r.refreshCh == nil { + r.refreshCh = make(chan struct{}, 1) + } + }) + return r.refreshCh +} + +func registryRefreshAcquireError(err error) *core.GatewayError { + if errors.Is(err, context.DeadlineExceeded) { + return core.NewProviderError("model_registry", http.StatusGatewayTimeout, "model registry refresh timed out before start", err) + } + return core.NewProviderError("model_registry", http.StatusRequestTimeout, "model registry refresh canceled before start", err) +} + +// InitializeAsync starts model fetching in a background goroutine. +// It first loads any cached models for immediate availability, then refreshes from network. +// Returns immediately after loading cache. The background goroutine will update models +// and save to cache when network fetch completes. +func (r *ModelRegistry) InitializeAsync(ctx context.Context) { + // First, try to load from cache for instant startup + cached, err := r.LoadFromCache(ctx) + if err != nil { + slog.Warn("failed to load models from cache", "error", err) + } else if cached > 0 { + slog.Info("serving traffic with cached models while refreshing", "cached_models", cached) + } + + // Start background initialization + go func() { + initCtx, cancel := context.WithTimeout(context.Background(), 60*time.Second) + defer cancel() + + if err := r.Initialize(initCtx); err != nil { + slog.Warn("background model initialization failed", "error", err) + return + } + + // Save to cache for next startup + if err := r.SaveToCache(initCtx); err != nil { + slog.Warn("failed to save models to cache", "error", err) + } + }() +} + +// IsInitialized returns true if at least one successful network fetch has completed. +// This can be used to check if the registry has fresh data or is only serving from cache. +func (r *ModelRegistry) IsInitialized() bool { + r.initMu.Lock() + defer r.initMu.Unlock() + return r.initialized +} + +// StartBackgroundRefresh starts a goroutine that periodically refreshes the model registry. +// If modelListURL is non-empty, the model list is also re-fetched on each tick. +// The returned stop function is blocking: it cancels the refresh loop and waits +// for the goroutine to exit before returning, so callers should expect it to +// block during shutdown until any in-flight refresh work unwinds. +func (r *ModelRegistry) StartBackgroundRefresh(interval time.Duration, modelListURL string) func() { + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + var stopOnce sync.Once + + go func() { + defer close(done) + ticker := time.NewTicker(interval) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + refreshCtx, refreshCancel := context.WithTimeout(ctx, 30*time.Second) + err := r.Initialize(refreshCtx) + refreshCancel() + if err != nil { + if !isBenignBackgroundRefreshError(ctx, err) { + slog.Warn("background model refresh failed", "error", err) + } + } else { + func() { + cacheCtx, cacheCancel := context.WithTimeout(ctx, 10*time.Second) + defer cacheCancel() + if err := r.SaveToCache(cacheCtx); err != nil { + if !isBenignBackgroundRefreshError(ctx, err) { + slog.Warn("failed to save models to cache after refresh", "error", err) + } + } + }() + } + + // Also refresh model list if configured + if modelListURL != "" { + r.refreshModelList(ctx, modelListURL) + } + } + } + }() + + return func() { + stopOnce.Do(func() { + cancel() + <-done + }) + } +} + +// RefreshModelList fetches the external model metadata list and re-enriches all +// currently registered models. It does not persist the model cache; callers that +// want durable startup data should call SaveToCache after this succeeds. +func (r *ModelRegistry) RefreshModelList(ctx context.Context, url string) (int, error) { + if strings.TrimSpace(url) == "" { + return 0, nil + } + + release, err := r.acquireRefresh(ctx) + if err != nil { + return 0, err + } + defer release() + + models, _, err := r.refreshModelListLocked(ctx, url) + return models, err +} + +func (r *ModelRegistry) refreshModelListLocked(ctx context.Context, url string) (int, metadataEnrichmentStats, error) { + list, raw, err := modeldata.Fetch(ctx, url) + if err != nil { + return 0, metadataEnrichmentStats{}, err + } + if list == nil { + return 0, metadataEnrichmentStats{}, nil + } + + metadataStats := r.setModelListAndEnrich(list, raw) + return len(list.Models), metadataStats, nil +} + +// refreshModelList fetches the model list and re-enriches all models. +func (r *ModelRegistry) refreshModelList(ctx context.Context, url string) { + fetchCtx, cancel := context.WithTimeout(ctx, 45*time.Second) + defer cancel() + + release, err := r.acquireRefresh(fetchCtx) + if err != nil { + if !isBenignBackgroundRefreshError(ctx, err) { + slog.Warn("failed to acquire model list refresh", "url", url, "error", err) + } + return + } + var ( + models int + metadataStats metadataEnrichmentStats + ) + func() { + defer release() + models, metadataStats, err = r.refreshModelListLocked(fetchCtx, url) + }() + if err != nil { + if !isBenignBackgroundRefreshError(ctx, err) { + slog.Warn("failed to refresh model list", "url", url, "error", err) + } + return + } + if models == 0 { + return + } + + if err := r.SaveToCache(fetchCtx); err != nil { + if !isBenignBackgroundRefreshError(ctx, err) { + slog.Warn("failed to save cache after model list refresh", "error", err) + } + } + attrs := []any{"models", models} + attrs = append(attrs, metadataStats.slogAttrs()...) + slog.Debug("model list refreshed", attrs...) +} + +func isBenignBackgroundRefreshError(parent context.Context, err error) bool { + if err == nil { + return true + } + if parent == nil || parent.Err() == nil { + return false + } + return errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) +} diff --git a/internal/providers/registry_metadata.go b/internal/providers/registry_metadata.go new file mode 100644 index 000000000..5a53fc146 --- /dev/null +++ b/internal/providers/registry_metadata.go @@ -0,0 +1,417 @@ +package providers + +import ( + "encoding/json" + "maps" + "reflect" + "slices" + "strings" + + "gomodel/config" + "gomodel/internal/core" + "gomodel/internal/modeldata" +) + +// SetModelList stores the parsed model list and its raw bytes for cache persistence. +func (r *ModelRegistry) SetModelList(list *modeldata.ModelList, raw json.RawMessage) { + r.mu.Lock() + defer r.mu.Unlock() + r.modelList = list + r.modelListRaw = raw +} + +// EnrichModels re-applies model list metadata to all currently registered models. +// Call this after SetModelList to update existing models with the new metadata. +// Holds the write lock for the entire operation and replaces published ModelInfo +// entries instead of mutating them in place so concurrent readers can safely keep +// using older snapshots after unlocking. +func (r *ModelRegistry) EnrichModels() { + _ = r.enrichModels() +} + +func (r *ModelRegistry) enrichModels() metadataEnrichmentStats { + r.mu.Lock() + defer r.mu.Unlock() + return r.enrichModelsLocked() +} + +func (r *ModelRegistry) enrichModelsLocked() metadataEnrichmentStats { + if len(r.models) == 0 { + return metadataEnrichmentStats{} + } + if r.modelList == nil && len(r.configMetadataOverrides) == 0 { + return metadataEnrichmentStats{} + } + + providerTypes := make(map[core.Provider]string, len(r.providerTypes)) + maps.Copy(providerTypes, r.providerTypes) + + replacements := make(map[*ModelInfo]*ModelInfo, len(r.models)) + stats := metadataEnrichmentStats{} + if r.modelList != nil { + stats = enrichProviderModelMaps(r.modelList, providerTypes, r.modelsByProvider, replacements) + } + stats.Enriched += applyConfigMetadataOverrides(r.configMetadataOverrides, r.modelsByProvider, replacements) + for modelID, info := range r.models { + if replacement, ok := replacements[info]; ok { + r.models[modelID] = replacement + } + } + r.invalidateSortedCaches() + return stats +} + +func (r *ModelRegistry) setModelListAndEnrich(list *modeldata.ModelList, raw json.RawMessage) metadataEnrichmentStats { + r.mu.Lock() + defer r.mu.Unlock() + r.modelList = list + r.modelListRaw = raw + return r.enrichModelsLocked() +} + +// ResolveMetadata resolves metadata for a model directly via the stored model list, +// bypassing the registry key lookup. This handles cases where the usage DB stores +// a response model ID (e.g., "gpt-4o-2024-08-06") that differs from the registry +// key (e.g., "gpt-4o") by using the reverse index in the model list. +func (r *ModelRegistry) ResolveMetadata(providerType, modelID string) *core.ModelMetadata { + r.mu.RLock() + defer r.mu.RUnlock() + if r.modelList == nil { + return nil + } + return modeldata.Resolve(r.modelList, providerType, modelID) +} + +// GetModelMetadata returns the metadata for a model, or nil if not found or not enriched. +func (r *ModelRegistry) GetModelMetadata(modelID string) *core.ModelMetadata { + r.mu.RLock() + defer r.mu.RUnlock() + if info, ok := r.models[modelID]; ok { + return info.Model.Metadata + } + return nil +} + +// ResolvePricing returns the pricing metadata for a model, trying the registry first +// and falling back to a reverse-index lookup via the model list. +// Returns nil if no pricing is available. +func (r *ModelRegistry) ResolvePricing(model, providerType string) *core.ModelPricing { + providerSelector := strings.TrimSpace(providerType) + if meta := r.getProviderModelMetadata(providerSelector, model); meta != nil && meta.Pricing != nil { + return meta.Pricing + } + + meta := r.GetModelMetadata(model) + if meta != nil && meta.Pricing != nil { + return meta.Pricing + } + if providerSelector != "" { + meta = r.ResolveMetadata(r.metadataProviderType(providerSelector), r.metadataModelID(model)) + if meta != nil && meta.Pricing != nil { + return meta.Pricing + } + } + return nil +} + +func (r *ModelRegistry) getProviderModelMetadata(providerSelector, model string) *core.ModelMetadata { + providerSelector = strings.TrimSpace(providerSelector) + model = strings.TrimSpace(model) + if model == "" { + return nil + } + + modelProviderName, modelID := splitModelSelector(model) + r.mu.RLock() + defer r.mu.RUnlock() + + if modelProviderName != "" { + if meta := metadataFromProviderModel(r.modelsByProvider[modelProviderName], modelID); meta != nil { + return meta + } + if r.hasConfiguredProviderNameLocked(modelProviderName) { + return nil + } + } + + if providerSelector == "" { + return nil + } + if meta := metadataFromProviderModel(r.modelsByProvider[providerSelector], model); meta != nil { + return meta + } + return nil +} + +func metadataFromProviderModel(providerModels map[string]*ModelInfo, model string) *core.ModelMetadata { + if len(providerModels) == 0 { + return nil + } + info := providerModels[strings.TrimSpace(model)] + if info == nil { + return nil + } + return info.Model.Metadata +} + +func (r *ModelRegistry) metadataProviderType(providerSelector string) string { + providerSelector = strings.TrimSpace(providerSelector) + if providerSelector == "" { + return "" + } + if providerType := r.GetProviderTypeForName(providerSelector); providerType != "" { + return providerType + } + return providerSelector +} + +func (r *ModelRegistry) metadataModelID(model string) string { + model = strings.TrimSpace(model) + providerName, modelID := splitModelSelector(model) + if providerName == "" { + return model + } + r.mu.RLock() + defer r.mu.RUnlock() + if r.hasConfiguredProviderNameLocked(providerName) { + return modelID + } + return model +} + +// snapshotProviderTypes returns a copy of the providerTypes map for use outside the lock. +func (r *ModelRegistry) snapshotProviderTypes() map[core.Provider]string { + r.mu.RLock() + defer r.mu.RUnlock() + m := make(map[core.Provider]string, len(r.providerTypes)) + maps.Copy(m, r.providerTypes) + return m +} + +// snapshotConfigOverrides returns a copy of the configMetadataOverrides outer +// and inner maps for use outside the lock. The inner *core.ModelMetadata +// pointers are shared, which is safe because SetProviderMetadataOverrides +// deep-clones on insertion and the registry never hands those values back out. +func (r *ModelRegistry) snapshotConfigOverrides() map[string]map[string]*core.ModelMetadata { + r.mu.RLock() + defer r.mu.RUnlock() + if len(r.configMetadataOverrides) == 0 { + return nil + } + out := make(map[string]map[string]*core.ModelMetadata, len(r.configMetadataOverrides)) + for provider, inner := range r.configMetadataOverrides { + innerCopy := make(map[string]*core.ModelMetadata, len(inner)) + for modelID, meta := range inner { + innerCopy[modelID] = meta + } + out[provider] = innerCopy + } + return out +} + +func (r *ModelRegistry) snapshotConfiguredProviderModels() (map[string][]string, config.ConfiguredProviderModelsMode) { + r.mu.RLock() + defer r.mu.RUnlock() + mode := config.ResolveConfiguredProviderModelsMode(r.configuredProviderModelsMode) + if len(r.configuredProviderModels) == 0 { + return nil, mode + } + out := make(map[string][]string, len(r.configuredProviderModels)) + for provider, models := range r.configuredProviderModels { + out[provider] = slices.Clone(models) + } + return out, mode +} + +// collectionEmpty reports whether a reflect.Value representing a slice, array, +// or map has no elements (covering both nil and non-nil-but-zero-length), and +// falls back to reflect.Value.IsZero for other kinds. This lets override- +// emptiness checks treat `modes: []` the same as an omitted field, which +// IsZero alone would not. +func collectionEmpty(v reflect.Value) bool { + switch v.Kind() { + case reflect.Slice, reflect.Map, reflect.Array: + return v.Len() == 0 + } + return v.IsZero() +} + +// structFieldsEmpty returns true if every field of the given struct value +// passes collectionEmpty. +func structFieldsEmpty(v reflect.Value) bool { + for i := 0; i < v.NumField(); i++ { + if !collectionEmpty(v.Field(i)) { + return false + } + } + return true +} + +// metadataOverrideEmpty reports whether an override has no effective content. +// An empty override (either nil or zero-valued on every field) would turn a +// nil current metadata into a non-nil empty struct after MergeMetadata, so +// callers should short-circuit on it. Uses reflect-based field inspection so +// new fields on core.ModelMetadata are picked up automatically; Pricing is +// handled separately so a non-nil pointer to an empty pricing block still +// counts as empty. +func metadataOverrideEmpty(m *core.ModelMetadata) bool { + if m == nil { + return true + } + if !pricingOverrideEmpty(m.Pricing) { + return false + } + tmp := *m + tmp.Pricing = nil + return structFieldsEmpty(reflect.ValueOf(tmp)) +} + +// pricingOverrideEmpty reports whether a pricing override has no effective +// content — nil or every field at its zero value (with collections treated as +// empty when length==0). +func pricingOverrideEmpty(p *core.ModelPricing) bool { + if p == nil { + return true + } + return structFieldsEmpty(reflect.ValueOf(*p)) +} + +// applyConfigMetadataOverrides layers operator-declared metadata onto already- +// enriched models. Call it after enrichProviderModelMaps with the same +// replacements map (pass nil replacements for fresh, unpublished maps). +// Returns the number of models whose metadata was updated. +func applyConfigMetadataOverrides( + overrides map[string]map[string]*core.ModelMetadata, + modelsByProvider map[string]map[string]*ModelInfo, + replacements map[*ModelInfo]*ModelInfo, +) int { + if len(overrides) == 0 { + return 0 + } + // reverse lets us find the pre-enrichment pointer when an entry has + // already been replaced by enrichProviderModelMaps, so our replacement + // chain stays consistent from the caller's perspective. Always allocated + // when replacements is non-nil so the else-branch write below cannot hit + // a nil map when enrichment made no replacements. + var reverse map[*ModelInfo]*ModelInfo + if replacements != nil { + reverse = make(map[*ModelInfo]*ModelInfo, len(replacements)) + for orig, repl := range replacements { + reverse[repl] = orig + } + } + applied := 0 + for providerName, modelOverrides := range overrides { + providerModels, ok := modelsByProvider[providerName] + if !ok { + continue + } + for modelID, override := range modelOverrides { + if metadataOverrideEmpty(override) { + // A nil or effectively-empty override has nothing to + // contribute. Skipping avoids turning a nil current metadata + // into a non-nil empty struct, which the DeepEqual check + // below would not catch. + continue + } + current, ok := providerModels[modelID] + if !ok { + continue + } + merged := modeldata.MergeMetadata(current.Model.Metadata, override) + // Skip no-op merges so concurrent readers holding the current + // pointer keep a stable view when the override adds no new info. + if reflect.DeepEqual(current.Model.Metadata, merged) { + continue + } + if replacements == nil { + current.Model.Metadata = merged + applied++ + continue + } + cloned := *current + cloned.Model.Metadata = merged + next := &cloned + providerModels[modelID] = next + if orig, hasOrig := reverse[current]; hasOrig { + replacements[orig] = next + reverse[next] = orig + } else { + replacements[current] = next + reverse[next] = current + } + applied++ + } + } + return applied +} + +func enrichProviderModelMaps( + list *modeldata.ModelList, + providerTypes map[core.Provider]string, + modelsByProvider map[string]map[string]*ModelInfo, + replacements map[*ModelInfo]*ModelInfo, +) metadataEnrichmentStats { + if list == nil { + return metadataEnrichmentStats{} + } + stats := metadataEnrichmentStats{} + for _, providerModels := range modelsByProvider { + if len(providerModels) == 0 { + continue + } + stats.Providers++ + accessor := ®istryAccessor{ + models: providerModels, + providerTypes: providerTypes, + replacements: replacements, + } + enrichStats := modeldata.Enrich(accessor, list) + stats.Enriched += enrichStats.Enriched + stats.Total += enrichStats.Total + } + return stats +} + +// registryAccessor implements modeldata.ModelInfoAccessor. +// The models map may be either an unpublished snapshot (Initialize, LoadFromCache) +// or the live registry map (EnrichModels, which uses replacements to preserve +// immutability of already-published ModelInfo values). +type registryAccessor struct { + models map[string]*ModelInfo + providerTypes map[core.Provider]string + replacements map[*ModelInfo]*ModelInfo +} + +func (a *registryAccessor) ModelIDs() []string { + ids := make([]string, 0, len(a.models)) + for id := range a.models { + ids = append(ids, id) + } + return ids +} + +func (a *registryAccessor) GetProviderType(modelID string) string { + info, ok := a.models[modelID] + if !ok { + return "" + } + if providerType := strings.TrimSpace(info.ProviderType); providerType != "" { + return providerType + } + return strings.TrimSpace(a.providerTypes[info.Provider]) +} + +func (a *registryAccessor) SetMetadata(modelID string, meta *core.ModelMetadata) { + if info, ok := a.models[modelID]; ok { + if a.replacements != nil { + cloned := *info + cloned.Model.Metadata = meta + replacement := &cloned + a.models[modelID] = replacement + a.replacements[info] = replacement + return + } + info.Model.Metadata = meta + } +} From a2bdfc0113863c07bc6f42969de8bbdcb74c36b7 Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Wed, 29 Apr 2026 23:28:03 +0200 Subject: [PATCH 06/15] refactor(guardrails): split GuardedProvider helpers into peer files internal/guardrails/provider.go was 1569 lines. The GuardedProvider decorator is small; the bulk was three other concerns (raw JSON batch-body rewriting, chat envelope preservation, responses input preservation). Split into peer files in the same package: provider.go (512) GuardedProvider struct + all decorator methods, message-clone helpers batch_rewrite.go (188) rewriteGuardedChatBatchBody, patchGuardedChatBatchBody, patchChatMessagesJSON, patchRawChatMessage, rewriteGuardedResponsesBatchBody, patchJSONObjectFields + jsonFieldPatch, unmarshalJSONArray, isZeroJSONFieldValue chat_message_apply.go (250) chatToMessages, applyMessagesToChat- PreservingEnvelope, tailMatchedSystem- Offsets, applyGuardedMessageToOriginal, newChatMessageFromGuardrail, applyGuardedContentToOriginal, rewriteStructuredContentWithTextRewrite, normalizeGuardrailMessageText responses_message_apply.go (641) responsesToMessages, responsesInput- ToMessages, coerceResponsesInputElements, responsesInputElementToGuardrailMessage, applyMessagesToResponses + envelope/element patchers, structured-content rewriters, cloneStringAnyMap Pure relocation; tests pass. Co-Authored-By: Claude Opus 4.7 --- internal/guardrails/batch_rewrite.go | 188 +++ internal/guardrails/chat_message_apply.go | 250 ++++ internal/guardrails/provider.go | 1057 ----------------- .../guardrails/responses_message_apply.go | 641 ++++++++++ 4 files changed, 1079 insertions(+), 1057 deletions(-) create mode 100644 internal/guardrails/batch_rewrite.go create mode 100644 internal/guardrails/chat_message_apply.go create mode 100644 internal/guardrails/responses_message_apply.go diff --git a/internal/guardrails/batch_rewrite.go b/internal/guardrails/batch_rewrite.go new file mode 100644 index 000000000..a18d792af --- /dev/null +++ b/internal/guardrails/batch_rewrite.go @@ -0,0 +1,188 @@ +package guardrails + +import ( + "encoding/json" + + "gomodel/internal/core" +) + +func rewriteGuardedChatBatchBody(originalBody json.RawMessage, original *core.ChatRequest, modified *core.ChatRequest) (json.RawMessage, error) { + body, err := patchGuardedChatBatchBody(originalBody, original, modified) + if err == nil { + return body, nil + } + return json.Marshal(modified) +} + +func patchGuardedChatBatchBody(originalBody json.RawMessage, original *core.ChatRequest, modified *core.ChatRequest) (json.RawMessage, error) { + if modified == nil { + return nil, core.NewInvalidRequestError("missing guarded chat request", nil) + } + + var raw map[string]json.RawMessage + if err := json.Unmarshal(originalBody, &raw); err != nil { + return nil, err + } + + patchedMessages, err := patchChatMessagesJSON(raw["messages"], original.Messages, modified.Messages) + if err != nil { + return nil, err + } + raw["messages"] = patchedMessages + return json.Marshal(raw) +} + +func patchChatMessagesJSON(originalRaw json.RawMessage, original, modified []core.Message) (json.RawMessage, error) { + originalRawItems, err := unmarshalJSONArray(originalRaw) + if err != nil { + return nil, err + } + if len(originalRawItems) != len(original) { + return nil, core.NewInvalidRequestError("guardrails chat message payload does not match parsed request", nil) + } + + systemOriginals := make([]json.RawMessage, 0, len(original)) + nonSystemOriginals := make([]json.RawMessage, 0, len(original)) + nonSystemMessages := make([]core.Message, 0, len(original)) + for i, msg := range original { + if msg.Role == "system" { + systemOriginals = append(systemOriginals, originalRawItems[i]) + continue + } + nonSystemOriginals = append(nonSystemOriginals, originalRawItems[i]) + nonSystemMessages = append(nonSystemMessages, msg) + } + + patched := make([]json.RawMessage, 0, len(modified)) + modifiedSystemCount := 0 + for _, msg := range modified { + if msg.Role == "system" { + modifiedSystemCount++ + } + } + systemMatchStart, originalSystemStart := tailMatchedSystemOffsets(len(systemOriginals), modifiedSystemCount) + nextSystem := 0 + nextNonSystem := 0 + for _, msg := range modified { + if msg.Role == "system" { + if nextSystem >= systemMatchStart { + item, err := patchRawChatMessage(systemOriginals[originalSystemStart+(nextSystem-systemMatchStart)], msg) + if err != nil { + return nil, err + } + patched = append(patched, item) + } else { + item, err := json.Marshal(msg) + if err != nil { + return nil, err + } + patched = append(patched, item) + } + nextSystem++ + continue + } + + if nextNonSystem >= len(nonSystemOriginals) { + return nil, core.NewInvalidRequestError("guardrails cannot insert non-system chat messages", nil) + } + if nonSystemMessages[nextNonSystem].Role != msg.Role { + return nil, core.NewInvalidRequestError("guardrails cannot reorder non-system chat messages", nil) + } + item, err := patchRawChatMessage(nonSystemOriginals[nextNonSystem], msg) + if err != nil { + return nil, err + } + patched = append(patched, item) + nextNonSystem++ + } + if nextNonSystem != len(nonSystemOriginals) { + return nil, core.NewInvalidRequestError("guardrails cannot add or remove non-system chat messages", nil) + } + + return json.Marshal(patched) +} + +func patchRawChatMessage(original json.RawMessage, modified core.Message) (json.RawMessage, error) { + var raw map[string]json.RawMessage + if err := json.Unmarshal(original, &raw); err != nil { + return nil, err + } + + updatedBody, err := json.Marshal(modified) + if err != nil { + return nil, err + } + + var updated map[string]json.RawMessage + if err := json.Unmarshal(updatedBody, &updated); err != nil { + return nil, err + } + + for _, field := range []string{"role", "content", "tool_calls", "tool_call_id"} { + delete(raw, field) + if value, ok := updated[field]; ok { + raw[field] = value + } + } + + return json.Marshal(raw) +} + +func rewriteGuardedResponsesBatchBody(originalBody json.RawMessage, modified *core.ResponsesRequest) (json.RawMessage, error) { + if modified == nil { + return nil, core.NewInvalidRequestError("missing guarded responses request", nil) + } + + body, err := patchJSONObjectFields(originalBody, map[string]jsonFieldPatch{ + "instructions": {value: modified.Instructions, omitWhenEmpty: modified.Instructions == ""}, + "input": {value: modified.Input}, + }) + if err == nil { + return body, nil + } + return json.Marshal(modified) +} + +type jsonFieldPatch struct { + value any + omitWhenEmpty bool +} + +func patchJSONObjectFields(originalBody json.RawMessage, patches map[string]jsonFieldPatch) (json.RawMessage, error) { + var raw map[string]json.RawMessage + if err := json.Unmarshal(originalBody, &raw); err != nil { + return nil, err + } + + for field, patch := range patches { + if patch.omitWhenEmpty && isZeroJSONFieldValue(patch.value) { + delete(raw, field) + continue + } + + encoded, err := json.Marshal(patch.value) + if err != nil { + return nil, err + } + raw[field] = encoded + } + + return json.Marshal(raw) +} + +func unmarshalJSONArray(raw json.RawMessage) ([]json.RawMessage, error) { + var items []json.RawMessage + if err := json.Unmarshal(raw, &items); err != nil { + return nil, err + } + return items, nil +} + +func isZeroJSONFieldValue(value any) bool { + switch v := value.(type) { + case string: + return v == "" + default: + return value == nil + } +} diff --git a/internal/guardrails/chat_message_apply.go b/internal/guardrails/chat_message_apply.go new file mode 100644 index 000000000..ac78320b3 --- /dev/null +++ b/internal/guardrails/chat_message_apply.go @@ -0,0 +1,250 @@ +package guardrails + +import ( + "reflect" + "strings" + + "gomodel/internal/core" +) + +// chatToMessages extracts the normalized message list from a ChatRequest. +func chatToMessages(req *core.ChatRequest) ([]Message, error) { + msgs := make([]Message, len(req.Messages)) + for i, m := range req.Messages { + text, err := normalizeGuardrailMessageText(m.Content) + if err != nil { + return nil, core.NewInvalidRequestError("invalid chat message content", err) + } + msgs[i] = Message{ + Role: m.Role, + Content: text, + ToolCalls: cloneToolCalls(m.ToolCalls), + ToolCallID: m.ToolCallID, + ContentNull: m.ContentNull || m.Content == nil, + } + } + return msgs, nil +} + +// applyMessagesToChatPreservingEnvelope applies guardrail message updates while +// preserving the original chat message envelopes and structured content shapes. +func applyMessagesToChatPreservingEnvelope(req *core.ChatRequest, msgs []Message) (*core.ChatRequest, error) { + systemOriginal := make([]core.Message, 0, len(req.Messages)) + nonSystemOriginal := make([]core.Message, 0, len(req.Messages)) + for _, original := range req.Messages { + if original.Role == "system" { + systemOriginal = append(systemOriginal, original) + continue + } + nonSystemOriginal = append(nonSystemOriginal, original) + } + + coreMessages := make([]core.Message, 0, len(msgs)) + modifiedSystemCount := 0 + for _, modified := range msgs { + if modified.Role == "system" { + modifiedSystemCount++ + } + } + systemMatchStart, originalSystemStart := tailMatchedSystemOffsets(len(systemOriginal), modifiedSystemCount) + nextSystem := 0 + nextNonSystem := 0 + for _, modified := range msgs { + if modified.Role == "system" { + if nextSystem >= systemMatchStart { + preserved, err := applyGuardedMessageToOriginal(systemOriginal[originalSystemStart+(nextSystem-systemMatchStart)], modified) + if err != nil { + return nil, err + } + coreMessages = append(coreMessages, preserved) + } else { + coreMessages = append(coreMessages, newChatMessageFromGuardrail(modified)) + } + nextSystem++ + continue + } + + if nextNonSystem >= len(nonSystemOriginal) { + return nil, core.NewInvalidRequestError("guardrails cannot insert non-system chat messages", nil) + } + original := nonSystemOriginal[nextNonSystem] + if modified.Role != original.Role { + return nil, core.NewInvalidRequestError("guardrails cannot reorder non-system chat messages", nil) + } + preserved, err := applyGuardedMessageToOriginal(original, modified) + if err != nil { + return nil, err + } + coreMessages = append(coreMessages, preserved) + nextNonSystem++ + } + + if nextNonSystem != len(nonSystemOriginal) { + return nil, core.NewInvalidRequestError("guardrails cannot add or remove non-system chat messages", nil) + } + + result := *req + result.Messages = coreMessages + return &result, nil +} + +func tailMatchedSystemOffsets(originalSystemCount, modifiedSystemCount int) (matchStart, originalStart int) { + matched := min(modifiedSystemCount, originalSystemCount) + return modifiedSystemCount - matched, originalSystemCount - matched +} + +func applyGuardedMessageToOriginal(original core.Message, modified Message) (core.Message, error) { + preserved := cloneChatMessageEnvelope(original) + preserved.Role = modified.Role + preserved.ToolCalls = cloneToolCalls(modified.ToolCalls) + preserved.ToolCallID = modified.ToolCallID + + content, contentNull, err := applyGuardedContentToOriginal(original.Content, modified.Content, modified.ContentNull) + if err != nil { + return core.Message{}, err + } + preserved.Content = content + preserved.ContentNull = contentNull + return preserved, nil +} + +func newChatMessageFromGuardrail(m Message) core.Message { + contentNull := m.ContentNull + if m.Content != "" { + contentNull = false + } + + content := any(m.Content) + if contentNull { + content = nil + } + + return core.Message{ + Role: m.Role, + Content: content, + ToolCalls: cloneToolCalls(m.ToolCalls), + ToolCallID: m.ToolCallID, + ContentNull: contentNull, + } +} + +func applyGuardedContentToOriginal(originalContent any, rewrittenText string, contentNull bool) (any, bool, error) { + if core.HasStructuredContent(originalContent) { + mergedContent, err := rewriteStructuredContentWithTextRewrite(originalContent, rewrittenText) + if err != nil { + return nil, false, err + } + return mergedContent, false, nil + } + + if rewrittenText != "" { + contentNull = false + } + if contentNull { + return nil, true, nil + } + return rewrittenText, false, nil +} + +func rewriteStructuredContentWithTextRewrite(originalContent any, rewrittenText string) (any, error) { + parts, ok := core.NormalizeContentParts(originalContent) + if !ok { + return nil, core.NewInvalidRequestError("guardrails cannot merge rewritten text into structured message", nil) + } + + // Guard against pathological numbers of content parts that could cause size + // computations for allocations to overflow on some platforms. + const maxContentParts = 1_000_000 + if len(parts) >= maxContentParts { + return nil, core.NewInvalidRequestError("guardrails cannot merge structured message with too many content parts", nil) + } + + originalTexts := make([]string, 0, len(parts)) + textPartIndexes := make([]int, 0, len(parts)) + for i, part := range parts { + if part.Type == "text" { + textPartIndexes = append(textPartIndexes, i) + originalTexts = append(originalTexts, part.Text) + } + } + + if len(textPartIndexes) == 0 { + merged := cloneContentParts(parts) + if rewrittenText != "" { + merged = append([]core.ContentPart{{Type: "text", Text: rewrittenText}}, merged...) + } + if len(merged) == 0 { + return nil, core.NewInvalidRequestError("guardrails produced empty structured message after rewrite", nil) + } + return merged, nil + } + + if len(textPartIndexes) == 1 { + merged := cloneContentParts(parts) + textIndex := textPartIndexes[0] + if rewrittenText == "" { + merged = append(merged[:textIndex], merged[textIndex+1:]...) + } else { + merged[textIndex].Text = rewrittenText + } + if len(merged) == 0 { + return nil, core.NewInvalidRequestError("guardrails produced empty structured message after rewrite", nil) + } + return merged, nil + } + + if rewrittenText == strings.Join(originalTexts, " ") { + return cloneContentParts(parts), nil + } + + merged := make([]core.ContentPart, 0, len(parts)) + insertedRewrittenText := false + for _, part := range parts { + if part.Type == "text" { + if !insertedRewrittenText && rewrittenText != "" { + rewrittenPart := cloneContentPart(part) + rewrittenPart.Text = rewrittenText + merged = append(merged, rewrittenPart) + insertedRewrittenText = true + } + continue + } + merged = append(merged, cloneContentPart(part)) + } + + if len(merged) == 0 { + return nil, core.NewInvalidRequestError("guardrails produced empty structured message after rewrite", nil) + } + return merged, nil +} + +func normalizeGuardrailMessageText(content any) (string, error) { + normalizedContent := content + switch typed := content.(type) { + case []map[string]any: + parts := make([]any, len(typed)) + for i, part := range typed { + parts[i] = part + } + normalizedContent = parts + default: + value := reflect.ValueOf(content) + if value.IsValid() && (value.Kind() == reflect.Slice || value.Kind() == reflect.Array) { + switch content.(type) { + case []any, []core.ContentPart: + default: + parts := make([]any, value.Len()) + for i := 0; i < value.Len(); i++ { + parts[i] = value.Index(i).Interface() + } + normalizedContent = parts + } + } + } + + normalized, err := core.NormalizeMessageContent(normalizedContent) + if err != nil { + return "", err + } + return core.ExtractTextContent(normalized), nil +} diff --git a/internal/guardrails/provider.go b/internal/guardrails/provider.go index b2d7bf5a3..6cf804e8b 100644 --- a/internal/guardrails/provider.go +++ b/internal/guardrails/provider.go @@ -2,10 +2,7 @@ package guardrails import ( "context" - "encoding/json" "io" - "reflect" - "strings" "gomodel/internal/batchrewrite" "gomodel/internal/core" @@ -196,186 +193,6 @@ func (g *GuardedProvider) passthroughRouter() (core.RoutablePassthrough, error) return pp, nil } -func rewriteGuardedChatBatchBody(originalBody json.RawMessage, original *core.ChatRequest, modified *core.ChatRequest) (json.RawMessage, error) { - body, err := patchGuardedChatBatchBody(originalBody, original, modified) - if err == nil { - return body, nil - } - return json.Marshal(modified) -} - -func patchGuardedChatBatchBody(originalBody json.RawMessage, original *core.ChatRequest, modified *core.ChatRequest) (json.RawMessage, error) { - if modified == nil { - return nil, core.NewInvalidRequestError("missing guarded chat request", nil) - } - - var raw map[string]json.RawMessage - if err := json.Unmarshal(originalBody, &raw); err != nil { - return nil, err - } - - patchedMessages, err := patchChatMessagesJSON(raw["messages"], original.Messages, modified.Messages) - if err != nil { - return nil, err - } - raw["messages"] = patchedMessages - return json.Marshal(raw) -} - -func patchChatMessagesJSON(originalRaw json.RawMessage, original, modified []core.Message) (json.RawMessage, error) { - originalRawItems, err := unmarshalJSONArray(originalRaw) - if err != nil { - return nil, err - } - if len(originalRawItems) != len(original) { - return nil, core.NewInvalidRequestError("guardrails chat message payload does not match parsed request", nil) - } - - systemOriginals := make([]json.RawMessage, 0, len(original)) - nonSystemOriginals := make([]json.RawMessage, 0, len(original)) - nonSystemMessages := make([]core.Message, 0, len(original)) - for i, msg := range original { - if msg.Role == "system" { - systemOriginals = append(systemOriginals, originalRawItems[i]) - continue - } - nonSystemOriginals = append(nonSystemOriginals, originalRawItems[i]) - nonSystemMessages = append(nonSystemMessages, msg) - } - - patched := make([]json.RawMessage, 0, len(modified)) - modifiedSystemCount := 0 - for _, msg := range modified { - if msg.Role == "system" { - modifiedSystemCount++ - } - } - systemMatchStart, originalSystemStart := tailMatchedSystemOffsets(len(systemOriginals), modifiedSystemCount) - nextSystem := 0 - nextNonSystem := 0 - for _, msg := range modified { - if msg.Role == "system" { - if nextSystem >= systemMatchStart { - item, err := patchRawChatMessage(systemOriginals[originalSystemStart+(nextSystem-systemMatchStart)], msg) - if err != nil { - return nil, err - } - patched = append(patched, item) - } else { - item, err := json.Marshal(msg) - if err != nil { - return nil, err - } - patched = append(patched, item) - } - nextSystem++ - continue - } - - if nextNonSystem >= len(nonSystemOriginals) { - return nil, core.NewInvalidRequestError("guardrails cannot insert non-system chat messages", nil) - } - if nonSystemMessages[nextNonSystem].Role != msg.Role { - return nil, core.NewInvalidRequestError("guardrails cannot reorder non-system chat messages", nil) - } - item, err := patchRawChatMessage(nonSystemOriginals[nextNonSystem], msg) - if err != nil { - return nil, err - } - patched = append(patched, item) - nextNonSystem++ - } - if nextNonSystem != len(nonSystemOriginals) { - return nil, core.NewInvalidRequestError("guardrails cannot add or remove non-system chat messages", nil) - } - - return json.Marshal(patched) -} - -func patchRawChatMessage(original json.RawMessage, modified core.Message) (json.RawMessage, error) { - var raw map[string]json.RawMessage - if err := json.Unmarshal(original, &raw); err != nil { - return nil, err - } - - updatedBody, err := json.Marshal(modified) - if err != nil { - return nil, err - } - - var updated map[string]json.RawMessage - if err := json.Unmarshal(updatedBody, &updated); err != nil { - return nil, err - } - - for _, field := range []string{"role", "content", "tool_calls", "tool_call_id"} { - delete(raw, field) - if value, ok := updated[field]; ok { - raw[field] = value - } - } - - return json.Marshal(raw) -} - -func rewriteGuardedResponsesBatchBody(originalBody json.RawMessage, modified *core.ResponsesRequest) (json.RawMessage, error) { - if modified == nil { - return nil, core.NewInvalidRequestError("missing guarded responses request", nil) - } - - body, err := patchJSONObjectFields(originalBody, map[string]jsonFieldPatch{ - "instructions": {value: modified.Instructions, omitWhenEmpty: modified.Instructions == ""}, - "input": {value: modified.Input}, - }) - if err == nil { - return body, nil - } - return json.Marshal(modified) -} - -type jsonFieldPatch struct { - value any - omitWhenEmpty bool -} - -func patchJSONObjectFields(originalBody json.RawMessage, patches map[string]jsonFieldPatch) (json.RawMessage, error) { - var raw map[string]json.RawMessage - if err := json.Unmarshal(originalBody, &raw); err != nil { - return nil, err - } - - for field, patch := range patches { - if patch.omitWhenEmpty && isZeroJSONFieldValue(patch.value) { - delete(raw, field) - continue - } - - encoded, err := json.Marshal(patch.value) - if err != nil { - return nil, err - } - raw[field] = encoded - } - - return json.Marshal(raw) -} - -func unmarshalJSONArray(raw json.RawMessage) ([]json.RawMessage, error) { - var items []json.RawMessage - if err := json.Unmarshal(raw, &items); err != nil { - return nil, err - } - return items, nil -} - -func isZeroJSONFieldValue(value any) bool { - switch v := value.(type) { - case string: - return v == "" - default: - return value == nil - } -} // CreateBatch delegates native batch creation and optionally applies guardrails to inline items. func (g *GuardedProvider) CreateBatch(ctx context.Context, providerType string, req *core.BatchRequest) (*core.BatchResponse, error) { @@ -609,881 +426,7 @@ func (g *GuardedProvider) PrepareBatchRequest(ctx context.Context, providerType return processGuardedBatchRequest(ctx, providerType, req, g.pipeline, g.batchFileTransport()) } -// --- Adapters: concrete requests ↔ normalized []Message --- - -// chatToMessages extracts the normalized message list from a ChatRequest. -func chatToMessages(req *core.ChatRequest) ([]Message, error) { - msgs := make([]Message, len(req.Messages)) - for i, m := range req.Messages { - text, err := normalizeGuardrailMessageText(m.Content) - if err != nil { - return nil, core.NewInvalidRequestError("invalid chat message content", err) - } - msgs[i] = Message{ - Role: m.Role, - Content: text, - ToolCalls: cloneToolCalls(m.ToolCalls), - ToolCallID: m.ToolCallID, - ContentNull: m.ContentNull || m.Content == nil, - } - } - return msgs, nil -} - -// applyMessagesToChatPreservingEnvelope applies guardrail message updates while -// preserving the original chat message envelopes and structured content shapes. -func applyMessagesToChatPreservingEnvelope(req *core.ChatRequest, msgs []Message) (*core.ChatRequest, error) { - systemOriginal := make([]core.Message, 0, len(req.Messages)) - nonSystemOriginal := make([]core.Message, 0, len(req.Messages)) - for _, original := range req.Messages { - if original.Role == "system" { - systemOriginal = append(systemOriginal, original) - continue - } - nonSystemOriginal = append(nonSystemOriginal, original) - } - - coreMessages := make([]core.Message, 0, len(msgs)) - modifiedSystemCount := 0 - for _, modified := range msgs { - if modified.Role == "system" { - modifiedSystemCount++ - } - } - systemMatchStart, originalSystemStart := tailMatchedSystemOffsets(len(systemOriginal), modifiedSystemCount) - nextSystem := 0 - nextNonSystem := 0 - for _, modified := range msgs { - if modified.Role == "system" { - if nextSystem >= systemMatchStart { - preserved, err := applyGuardedMessageToOriginal(systemOriginal[originalSystemStart+(nextSystem-systemMatchStart)], modified) - if err != nil { - return nil, err - } - coreMessages = append(coreMessages, preserved) - } else { - coreMessages = append(coreMessages, newChatMessageFromGuardrail(modified)) - } - nextSystem++ - continue - } - - if nextNonSystem >= len(nonSystemOriginal) { - return nil, core.NewInvalidRequestError("guardrails cannot insert non-system chat messages", nil) - } - original := nonSystemOriginal[nextNonSystem] - if modified.Role != original.Role { - return nil, core.NewInvalidRequestError("guardrails cannot reorder non-system chat messages", nil) - } - preserved, err := applyGuardedMessageToOriginal(original, modified) - if err != nil { - return nil, err - } - coreMessages = append(coreMessages, preserved) - nextNonSystem++ - } - - if nextNonSystem != len(nonSystemOriginal) { - return nil, core.NewInvalidRequestError("guardrails cannot add or remove non-system chat messages", nil) - } - - result := *req - result.Messages = coreMessages - return &result, nil -} - -func tailMatchedSystemOffsets(originalSystemCount, modifiedSystemCount int) (matchStart, originalStart int) { - matched := min(modifiedSystemCount, originalSystemCount) - return modifiedSystemCount - matched, originalSystemCount - matched -} - -func applyGuardedMessageToOriginal(original core.Message, modified Message) (core.Message, error) { - preserved := cloneChatMessageEnvelope(original) - preserved.Role = modified.Role - preserved.ToolCalls = cloneToolCalls(modified.ToolCalls) - preserved.ToolCallID = modified.ToolCallID - - content, contentNull, err := applyGuardedContentToOriginal(original.Content, modified.Content, modified.ContentNull) - if err != nil { - return core.Message{}, err - } - preserved.Content = content - preserved.ContentNull = contentNull - return preserved, nil -} - -func newChatMessageFromGuardrail(m Message) core.Message { - contentNull := m.ContentNull - if m.Content != "" { - contentNull = false - } - - content := any(m.Content) - if contentNull { - content = nil - } - - return core.Message{ - Role: m.Role, - Content: content, - ToolCalls: cloneToolCalls(m.ToolCalls), - ToolCallID: m.ToolCallID, - ContentNull: contentNull, - } -} - -func applyGuardedContentToOriginal(originalContent any, rewrittenText string, contentNull bool) (any, bool, error) { - if core.HasStructuredContent(originalContent) { - mergedContent, err := rewriteStructuredContentWithTextRewrite(originalContent, rewrittenText) - if err != nil { - return nil, false, err - } - return mergedContent, false, nil - } - - if rewrittenText != "" { - contentNull = false - } - if contentNull { - return nil, true, nil - } - return rewrittenText, false, nil -} - -func rewriteStructuredContentWithTextRewrite(originalContent any, rewrittenText string) (any, error) { - parts, ok := core.NormalizeContentParts(originalContent) - if !ok { - return nil, core.NewInvalidRequestError("guardrails cannot merge rewritten text into structured message", nil) - } - - // Guard against pathological numbers of content parts that could cause size - // computations for allocations to overflow on some platforms. - const maxContentParts = 1_000_000 - if len(parts) >= maxContentParts { - return nil, core.NewInvalidRequestError("guardrails cannot merge structured message with too many content parts", nil) - } - - originalTexts := make([]string, 0, len(parts)) - textPartIndexes := make([]int, 0, len(parts)) - for i, part := range parts { - if part.Type == "text" { - textPartIndexes = append(textPartIndexes, i) - originalTexts = append(originalTexts, part.Text) - } - } - - if len(textPartIndexes) == 0 { - merged := cloneContentParts(parts) - if rewrittenText != "" { - merged = append([]core.ContentPart{{Type: "text", Text: rewrittenText}}, merged...) - } - if len(merged) == 0 { - return nil, core.NewInvalidRequestError("guardrails produced empty structured message after rewrite", nil) - } - return merged, nil - } - - if len(textPartIndexes) == 1 { - merged := cloneContentParts(parts) - textIndex := textPartIndexes[0] - if rewrittenText == "" { - merged = append(merged[:textIndex], merged[textIndex+1:]...) - } else { - merged[textIndex].Text = rewrittenText - } - if len(merged) == 0 { - return nil, core.NewInvalidRequestError("guardrails produced empty structured message after rewrite", nil) - } - return merged, nil - } - - if rewrittenText == strings.Join(originalTexts, " ") { - return cloneContentParts(parts), nil - } - - merged := make([]core.ContentPart, 0, len(parts)) - insertedRewrittenText := false - for _, part := range parts { - if part.Type == "text" { - if !insertedRewrittenText && rewrittenText != "" { - rewrittenPart := cloneContentPart(part) - rewrittenPart.Text = rewrittenText - merged = append(merged, rewrittenPart) - insertedRewrittenText = true - } - continue - } - merged = append(merged, cloneContentPart(part)) - } - - if len(merged) == 0 { - return nil, core.NewInvalidRequestError("guardrails produced empty structured message after rewrite", nil) - } - return merged, nil -} - -func normalizeGuardrailMessageText(content any) (string, error) { - normalizedContent := content - switch typed := content.(type) { - case []map[string]any: - parts := make([]any, len(typed)) - for i, part := range typed { - parts[i] = part - } - normalizedContent = parts - default: - value := reflect.ValueOf(content) - if value.IsValid() && (value.Kind() == reflect.Slice || value.Kind() == reflect.Array) { - switch content.(type) { - case []any, []core.ContentPart: - default: - parts := make([]any, value.Len()) - for i := 0; i < value.Len(); i++ { - parts[i] = value.Index(i).Interface() - } - normalizedContent = parts - } - } - } - - normalized, err := core.NormalizeMessageContent(normalizedContent) - if err != nil { - return "", err - } - return core.ExtractTextContent(normalized), nil -} - -// responsesToMessages extracts the normalized message list from a ResponsesRequest. -// The Instructions field maps to a system message and the input items are kept in -// one-to-one order so content-rewriting guardrails can be applied back safely. -func responsesToMessages(req *core.ResponsesRequest) ([]Message, error) { - var msgs []Message - if req.Instructions != "" { - msgs = append(msgs, Message{Role: "system", Content: req.Instructions}) - } - - inputMsgs, err := responsesInputToMessages(req.Input) - if err != nil { - return nil, err - } - msgs = append(msgs, inputMsgs...) - return msgs, nil -} - -func responsesInputToMessages(input any) ([]Message, error) { - switch typed := input.(type) { - case nil: - return nil, nil - case string: - return []Message{{Role: "user", Content: typed}}, nil - } - - elements, err := coerceResponsesInputElements(input) - if err != nil { - return nil, err - } - msgs := make([]Message, len(elements)) - for i, element := range elements { - msg, err := responsesInputElementToGuardrailMessage(element, i) - if err != nil { - return nil, err - } - msgs[i] = msg - } - return msgs, nil -} - -func coerceResponsesInputElements(input any) ([]core.ResponsesInputElement, error) { - switch typed := input.(type) { - case []core.ResponsesInputElement: - elements := make([]core.ResponsesInputElement, len(typed)) - copy(elements, typed) - return elements, nil - case []map[string]any: - elements := make([]core.ResponsesInputElement, len(typed)) - for i, item := range typed { - raw, err := json.Marshal(item) - if err != nil { - return nil, core.NewInvalidRequestError("invalid responses input item", err) - } - if err := json.Unmarshal(raw, &elements[i]); err != nil { - return nil, core.NewInvalidRequestError("invalid responses input item", err) - } - } - return elements, nil - case []any: - elements := make([]core.ResponsesInputElement, len(typed)) - for i, item := range typed { - raw, err := json.Marshal(item) - if err != nil { - return nil, core.NewInvalidRequestError("invalid responses input item", err) - } - if err := json.Unmarshal(raw, &elements[i]); err != nil { - return nil, core.NewInvalidRequestError("invalid responses input item", err) - } - } - return elements, nil - default: - return nil, core.NewInvalidRequestError("invalid responses input: unsupported type", nil) - } -} - -func responsesInputElementToGuardrailMessage(item core.ResponsesInputElement, index int) (Message, error) { - switch item.Type { - case "function_call": - if strings.TrimSpace(item.Name) == "" { - return Message{}, core.NewInvalidRequestError("invalid responses input item: function_call name is required", nil) - } - return Message{ - Role: "assistant", - Content: "", - ContentNull: true, - ToolCalls: []core.ToolCall{ - { - ID: item.CallID, - Type: "function", - Function: core.FunctionCall{ - Name: item.Name, - Arguments: item.Arguments, - }, - }, - }, - }, nil - case "function_call_output": - content, err := stringifyResponsesValue(item.Output) - if err != nil { - return Message{}, core.NewInvalidRequestError("invalid responses input item: function_call_output.output must be JSON-serializable", err) - } - return Message{ - Role: "tool", - ToolCallID: item.CallID, - Content: content, - }, nil - default: - role := strings.TrimSpace(item.Role) - if role == "" { - return Message{}, core.NewInvalidRequestError("invalid responses input item: role is required", nil) - } - text, err := normalizeGuardrailMessageText(item.Content) - if err != nil { - return Message{}, core.NewInvalidRequestError("invalid responses input item: unsupported content", err) - } - return Message{ - Role: role, - Content: text, - ContentNull: item.Content == nil, - }, nil - } -} - -func stringifyResponsesValue(value any) (string, error) { - switch typed := value.(type) { - case nil: - return "", nil - case string: - return typed, nil - default: - raw, err := json.Marshal(typed) - if err != nil { - return "", err - } - return string(raw), nil - } -} - -// applyMessagesToResponses returns a shallow copy of req with system and input -// messages applied back to the original Responses envelope. -func applyMessagesToResponses(req *core.ResponsesRequest, msgs []Message) (*core.ResponsesRequest, error) { - result := *req - originalInputMsgs, err := responsesInputToMessages(req.Input) - if err != nil { - return nil, err - } - - inputMsgs := msgs - switch { - case len(msgs) == len(originalInputMsgs): - result.Instructions = "" - case len(msgs) == len(originalInputMsgs)+1 && len(msgs) > 0 && msgs[0].Role == "system": - result.Instructions = msgs[0].Content - inputMsgs = msgs[1:] - default: - return nil, core.NewInvalidRequestError("guardrails cannot add or remove responses input items", nil) - } - - input, err := applyMessagesToResponsesInput(req.Input, inputMsgs) - if err != nil { - return nil, err - } - result.Input = input - return &result, nil -} - -func applyMessagesToResponsesInput(original any, msgs []Message) (any, error) { - switch original.(type) { - case nil: - if len(msgs) != 0 { - return nil, core.NewInvalidRequestError("guardrails cannot add or remove responses input items", nil) - } - return nil, nil - case string: - if len(msgs) != 1 { - return nil, core.NewInvalidRequestError("guardrails cannot add or remove responses input items", nil) - } - if msgs[0].Role != "user" { - return nil, core.NewInvalidRequestError("guardrails cannot change the role of a string responses input", nil) - } - if msgs[0].ContentNull { - return "", nil - } - return msgs[0].Content, nil - } - - elements, err := coerceResponsesInputElements(original) - if err != nil { - return nil, err - } - patched, err := applyMessagesToResponsesElements(elements, msgs) - if err != nil { - return nil, err - } - return patchResponsesInputEnvelope(original, patched) -} - -func applyMessagesToResponsesElements(elements []core.ResponsesInputElement, msgs []Message) ([]core.ResponsesInputElement, error) { - if len(msgs) != len(elements) { - return nil, core.NewInvalidRequestError("guardrails cannot add or remove responses input items", nil) - } - - result := make([]core.ResponsesInputElement, len(elements)) - for i, original := range elements { - patched, err := applyGuardedResponsesElementToOriginal(original, msgs[i], i) - if err != nil { - return nil, err - } - result[i] = patched - } - return result, nil -} - -func applyGuardedResponsesElementToOriginal(original core.ResponsesInputElement, modified Message, _ int) (core.ResponsesInputElement, error) { - preserved := original - - switch original.Type { - case "function_call": - if modified.Role != "assistant" { - return core.ResponsesInputElement{}, core.NewInvalidRequestError("guardrails cannot reorder or retag responses input items", nil) - } - return preserved, nil - case "function_call_output": - if modified.Role != "tool" { - return core.ResponsesInputElement{}, core.NewInvalidRequestError("guardrails cannot reorder or retag responses input items", nil) - } - preserved.Output = modified.Content - return preserved, nil - default: - role := strings.TrimSpace(original.Role) - if role == "" { - return core.ResponsesInputElement{}, core.NewInvalidRequestError("invalid responses input item: role is required", nil) - } - if modified.Role != role { - return core.ResponsesInputElement{}, core.NewInvalidRequestError("guardrails cannot reorder or retag responses input items", nil) - } - content, err := applyGuardedResponsesContentToOriginal(original.Content, modified.Content, modified.ContentNull) - if err != nil { - return core.ResponsesInputElement{}, core.NewInvalidRequestError("guardrails cannot merge rewritten text into responses input item", err) - } - preserved.Content = content - return preserved, nil - } -} - -func applyGuardedResponsesContentToOriginal(originalContent any, rewrittenText string, contentNull bool) (any, error) { - if isResponsesStructuredContent(originalContent) { - return rewriteStructuredResponsesContentWithTextRewrite(originalContent, rewrittenText) - } - if rewrittenText != "" { - contentNull = false - } - if contentNull { - return nil, nil - } - return rewrittenText, nil -} - -func patchResponsesInputEnvelope(original any, patched []core.ResponsesInputElement) (any, error) { - switch typed := original.(type) { - case []core.ResponsesInputElement: - result := make([]core.ResponsesInputElement, len(patched)) - copy(result, patched) - return result, nil - case []map[string]any: - return patchResponsesInputMapSlice(typed, patched) - case []any: - return patchResponsesInputInterfaceSlice(typed, patched) - default: - return patched, nil - } -} - -func patchResponsesInputMapSlice(original []map[string]any, patched []core.ResponsesInputElement) ([]map[string]any, error) { - if len(original) != len(patched) { - return nil, core.NewInvalidRequestError("guardrails cannot add or remove responses input items", nil) - } - result := make([]map[string]any, len(original)) - for i := range original { - item, err := patchResponsesInputMap(original[i], patched[i]) - if err != nil { - return nil, err - } - result[i] = item - } - return result, nil -} - -func patchResponsesInputInterfaceSlice(original []any, patched []core.ResponsesInputElement) ([]any, error) { - if len(original) != len(patched) { - return nil, core.NewInvalidRequestError("guardrails cannot add or remove responses input items", nil) - } - result := make([]any, len(original)) - for i := range original { - item, err := patchResponsesInputInterfaceElement(original[i], patched[i]) - if err != nil { - return nil, err - } - result[i] = item - } - return result, nil -} - -func patchResponsesInputInterfaceElement(original any, patched core.ResponsesInputElement) (any, error) { - if originalMap, ok := original.(map[string]any); ok { - return patchResponsesInputMap(originalMap, patched) - } - return responsesInputElementAsAny(patched) -} - -func patchResponsesInputMap(original map[string]any, patched core.ResponsesInputElement) (map[string]any, error) { - cloned := cloneStringAnyMap(original) - updated, err := responsesInputElementAsMap(patched) - if err != nil { - return nil, err - } - - for _, key := range []string{"type", "role", "status", "content", "call_id", "id", "name", "arguments", "output"} { - delete(cloned, key) - } - for key, value := range updated { - cloned[key] = value - } - if patched.Type == "function_call_output" { - cloned["output"] = restoreResponsesInputOutputValue(original["output"], patched.Output) - } - return cloned, nil -} - -func restoreResponsesInputOutputValue(original any, rewritten string) any { - if _, ok := original.(string); ok { - return rewritten - } - if strings.TrimSpace(rewritten) == "" { - if original == nil { - return nil - } - return original - } - - var decoded any - if err := json.Unmarshal([]byte(rewritten), &decoded); err == nil { - return decoded - } - if original == nil { - return nil - } - return original -} - -func responsesInputElementAsMap(element core.ResponsesInputElement) (map[string]any, error) { - value, err := responsesInputElementAsAny(element) - if err != nil { - return nil, err - } - itemMap, ok := value.(map[string]any) - if !ok { - return nil, core.NewInvalidRequestError("invalid responses input item", nil) - } - return itemMap, nil -} - -func responsesInputElementAsAny(element core.ResponsesInputElement) (any, error) { - raw, err := json.Marshal(element) - if err != nil { - return nil, core.NewInvalidRequestError("invalid responses input item", err) - } - var value any - if err := json.Unmarshal(raw, &value); err != nil { - return nil, core.NewInvalidRequestError("invalid responses input item", err) - } - return value, nil -} - -func isResponsesStructuredContent(content any) bool { - if content == nil { - return false - } - switch typed := content.(type) { - case []any: - return true - case []core.ContentPart: - return true - case []map[string]any: - return true - default: - _ = typed - } - contentType := reflect.TypeOf(content) - return contentType.Kind() == reflect.Slice || contentType.Kind() == reflect.Array -} - -func rewriteStructuredResponsesContentWithTextRewrite(originalContent any, rewrittenText string) (any, error) { - switch typed := originalContent.(type) { - case []core.ContentPart: - return rewriteStructuredResponsesTypedContentParts(typed, rewrittenText) - case []any: - return rewriteStructuredResponsesInterfaceContentParts(typed, rewrittenText) - case []map[string]any: - return rewriteStructuredResponsesMapContentParts(typed, rewrittenText) - default: - value := reflect.ValueOf(originalContent) - if !value.IsValid() || (value.Kind() != reflect.Slice && value.Kind() != reflect.Array) { - return nil, core.NewInvalidRequestError("unsupported structured responses content", nil) - } - parts := make([]any, value.Len()) - for i := 0; i < value.Len(); i++ { - parts[i] = value.Index(i).Interface() - } - return rewriteStructuredResponsesInterfaceContentParts(parts, rewrittenText) - } -} - -func rewriteStructuredResponsesTypedContentParts(parts []core.ContentPart, rewrittenText string) (any, error) { - textIndexes := make([]int, 0, len(parts)) - originalTexts := make([]string, 0, len(parts)) - for i, part := range parts { - if !isResponsesTextPartType(part.Type) || part.Text == "" { - continue - } - textIndexes = append(textIndexes, i) - originalTexts = append(originalTexts, part.Text) - } - - if len(textIndexes) == 0 { - if rewrittenText == "" { - if len(parts) == 0 { - return []core.ContentPart{}, nil - } - return cloneContentParts(parts), nil - } - prepended := []core.ContentPart{{Type: "input_text", Text: rewrittenText}} - prepended = append(prepended, cloneContentParts(parts)...) - return prepended, nil - } - - if len(textIndexes) == 1 { - merged := cloneContentParts(parts) - textIndex := textIndexes[0] - if rewrittenText == "" { - merged = append(merged[:textIndex], merged[textIndex+1:]...) - } else { - merged[textIndex].Text = rewrittenText - } - if len(merged) == 0 { - return nil, core.NewInvalidRequestError("guardrails produced empty structured responses content after rewrite", nil) - } - return merged, nil - } - - if rewrittenText == strings.Join(originalTexts, " ") { - return cloneContentParts(parts), nil - } - - merged := make([]core.ContentPart, 0, len(parts)) - insertedRewrittenText := false - for _, part := range parts { - if isResponsesTextPartType(part.Type) { - if !insertedRewrittenText && rewrittenText != "" { - rewrittenPart := cloneContentPart(part) - rewrittenPart.Text = rewrittenText - merged = append(merged, rewrittenPart) - insertedRewrittenText = true - } - continue - } - merged = append(merged, cloneContentPart(part)) - } - - if len(merged) == 0 { - return nil, core.NewInvalidRequestError("guardrails produced empty structured responses content after rewrite", nil) - } - return merged, nil -} - -func rewriteStructuredResponsesInterfaceContentParts(parts []any, rewrittenText string) (any, error) { - textIndexes := make([]int, 0, len(parts)) - originalTexts := make([]string, 0, len(parts)) - for i, part := range parts { - partMap, ok := part.(map[string]any) - if !ok { - continue - } - partType, _ := partMap["type"].(string) - if !isResponsesTextPartType(partType) { - continue - } - text, _ := partMap["text"].(string) - if text == "" { - continue - } - textIndexes = append(textIndexes, i) - originalTexts = append(originalTexts, text) - } - - if len(textIndexes) == 0 { - if rewrittenText == "" { - if len(parts) == 0 { - return []any{}, nil - } - return cloneResponsesInterfaceParts(parts), nil - } - prepended := []any{map[string]any{"type": "input_text", "text": rewrittenText}} - prepended = append(prepended, cloneResponsesInterfaceParts(parts)...) - return prepended, nil - } - - if len(textIndexes) == 1 { - merged := make([]any, 0, len(parts)) - textIndex := textIndexes[0] - for i, part := range parts { - if i == textIndex { - if rewrittenText == "" { - continue - } - partMap, ok := part.(map[string]any) - if !ok { - return nil, core.NewInvalidRequestError("guardrails cannot rewrite non-object responses content part", nil) - } - cloned := cloneStringAnyMap(partMap) - cloned["text"] = rewrittenText - merged = append(merged, cloned) - continue - } - merged = append(merged, cloneResponsesInterfacePart(part)) - } - if len(merged) == 0 { - return nil, core.NewInvalidRequestError("guardrails produced empty structured responses content after rewrite", nil) - } - return merged, nil - } - - if rewrittenText == strings.Join(originalTexts, " ") { - return cloneResponsesInterfaceParts(parts), nil - } - - merged := make([]any, 0, len(parts)) - insertedRewrittenText := false - for _, part := range parts { - partMap, ok := part.(map[string]any) - if ok { - partType, _ := partMap["type"].(string) - if isResponsesTextPartType(partType) { - if !insertedRewrittenText && rewrittenText != "" { - cloned := cloneStringAnyMap(partMap) - cloned["text"] = rewrittenText - merged = append(merged, cloned) - insertedRewrittenText = true - } - continue - } - } - merged = append(merged, cloneResponsesInterfacePart(part)) - } - - if len(merged) == 0 { - return nil, core.NewInvalidRequestError("guardrails produced empty structured responses content after rewrite", nil) - } - return merged, nil -} - -func rewriteStructuredResponsesMapContentParts(parts []map[string]any, rewrittenText string) (any, error) { - if len(parts) == 0 && rewrittenText == "" { - return []map[string]any{}, nil - } - - interfaceParts := make([]any, len(parts)) - for i, part := range parts { - interfaceParts[i] = part - } - rewritten, err := rewriteStructuredResponsesInterfaceContentParts(interfaceParts, rewrittenText) - if err != nil { - return nil, err - } - rewrittenParts, ok := rewritten.([]any) - if !ok { - return nil, core.NewInvalidRequestError("unsupported structured responses content", nil) - } - result := make([]map[string]any, len(rewrittenParts)) - for i, part := range rewrittenParts { - partMap, ok := part.(map[string]any) - if !ok { - return nil, core.NewInvalidRequestError("guardrails cannot rewrite non-object responses content part", nil) - } - result[i] = cloneStringAnyMap(partMap) - } - return result, nil -} - -func isResponsesTextPartType(partType string) bool { - switch partType { - case "text", "input_text", "output_text": - return true - default: - return false - } -} - -func cloneResponsesInterfaceParts(parts []any) []any { - if len(parts) == 0 { - return nil - } - cloned := make([]any, len(parts)) - for i, part := range parts { - cloned[i] = cloneResponsesInterfacePart(part) - } - return cloned -} - -func cloneResponsesInterfacePart(part any) any { - partMap, ok := part.(map[string]any) - if !ok { - return part - } - return cloneStringAnyMap(partMap) -} - -// cloneStringAnyMap performs a shallow copy of the map. Nested maps/slices are -// intentionally shared; callers are expected to either preserve them as-is or -// replace whole top-level values instead of mutating nested structures in place. -func cloneStringAnyMap(src map[string]any) map[string]any { - if src == nil { - return nil - } - cloned := make(map[string]any, len(src)) - for key, value := range src { - cloned[key] = value - } - return cloned -} func cloneToolCalls(toolCalls []core.ToolCall) []core.ToolCall { if len(toolCalls) == 0 { diff --git a/internal/guardrails/responses_message_apply.go b/internal/guardrails/responses_message_apply.go new file mode 100644 index 000000000..5922f0f9e --- /dev/null +++ b/internal/guardrails/responses_message_apply.go @@ -0,0 +1,641 @@ +package guardrails + +import ( + "encoding/json" + "reflect" + "strings" + + "gomodel/internal/core" +) + +// responsesToMessages extracts the normalized message list from a ResponsesRequest. +// The Instructions field maps to a system message and the input items are kept in +// one-to-one order so content-rewriting guardrails can be applied back safely. +func responsesToMessages(req *core.ResponsesRequest) ([]Message, error) { + var msgs []Message + if req.Instructions != "" { + msgs = append(msgs, Message{Role: "system", Content: req.Instructions}) + } + + inputMsgs, err := responsesInputToMessages(req.Input) + if err != nil { + return nil, err + } + msgs = append(msgs, inputMsgs...) + return msgs, nil +} + +func responsesInputToMessages(input any) ([]Message, error) { + switch typed := input.(type) { + case nil: + return nil, nil + case string: + return []Message{{Role: "user", Content: typed}}, nil + } + + elements, err := coerceResponsesInputElements(input) + if err != nil { + return nil, err + } + msgs := make([]Message, len(elements)) + for i, element := range elements { + msg, err := responsesInputElementToGuardrailMessage(element, i) + if err != nil { + return nil, err + } + msgs[i] = msg + } + return msgs, nil +} + +func coerceResponsesInputElements(input any) ([]core.ResponsesInputElement, error) { + switch typed := input.(type) { + case []core.ResponsesInputElement: + elements := make([]core.ResponsesInputElement, len(typed)) + copy(elements, typed) + return elements, nil + case []map[string]any: + elements := make([]core.ResponsesInputElement, len(typed)) + for i, item := range typed { + raw, err := json.Marshal(item) + if err != nil { + return nil, core.NewInvalidRequestError("invalid responses input item", err) + } + if err := json.Unmarshal(raw, &elements[i]); err != nil { + return nil, core.NewInvalidRequestError("invalid responses input item", err) + } + } + return elements, nil + case []any: + elements := make([]core.ResponsesInputElement, len(typed)) + for i, item := range typed { + raw, err := json.Marshal(item) + if err != nil { + return nil, core.NewInvalidRequestError("invalid responses input item", err) + } + if err := json.Unmarshal(raw, &elements[i]); err != nil { + return nil, core.NewInvalidRequestError("invalid responses input item", err) + } + } + return elements, nil + default: + return nil, core.NewInvalidRequestError("invalid responses input: unsupported type", nil) + } +} + +func responsesInputElementToGuardrailMessage(item core.ResponsesInputElement, index int) (Message, error) { + switch item.Type { + case "function_call": + if strings.TrimSpace(item.Name) == "" { + return Message{}, core.NewInvalidRequestError("invalid responses input item: function_call name is required", nil) + } + return Message{ + Role: "assistant", + Content: "", + ContentNull: true, + ToolCalls: []core.ToolCall{ + { + ID: item.CallID, + Type: "function", + Function: core.FunctionCall{ + Name: item.Name, + Arguments: item.Arguments, + }, + }, + }, + }, nil + case "function_call_output": + content, err := stringifyResponsesValue(item.Output) + if err != nil { + return Message{}, core.NewInvalidRequestError("invalid responses input item: function_call_output.output must be JSON-serializable", err) + } + return Message{ + Role: "tool", + ToolCallID: item.CallID, + Content: content, + }, nil + default: + role := strings.TrimSpace(item.Role) + if role == "" { + return Message{}, core.NewInvalidRequestError("invalid responses input item: role is required", nil) + } + text, err := normalizeGuardrailMessageText(item.Content) + if err != nil { + return Message{}, core.NewInvalidRequestError("invalid responses input item: unsupported content", err) + } + return Message{ + Role: role, + Content: text, + ContentNull: item.Content == nil, + }, nil + } +} + +func stringifyResponsesValue(value any) (string, error) { + switch typed := value.(type) { + case nil: + return "", nil + case string: + return typed, nil + default: + raw, err := json.Marshal(typed) + if err != nil { + return "", err + } + return string(raw), nil + } +} + +// applyMessagesToResponses returns a shallow copy of req with system and input +// messages applied back to the original Responses envelope. +func applyMessagesToResponses(req *core.ResponsesRequest, msgs []Message) (*core.ResponsesRequest, error) { + result := *req + originalInputMsgs, err := responsesInputToMessages(req.Input) + if err != nil { + return nil, err + } + + inputMsgs := msgs + switch { + case len(msgs) == len(originalInputMsgs): + result.Instructions = "" + case len(msgs) == len(originalInputMsgs)+1 && len(msgs) > 0 && msgs[0].Role == "system": + result.Instructions = msgs[0].Content + inputMsgs = msgs[1:] + default: + return nil, core.NewInvalidRequestError("guardrails cannot add or remove responses input items", nil) + } + + input, err := applyMessagesToResponsesInput(req.Input, inputMsgs) + if err != nil { + return nil, err + } + result.Input = input + return &result, nil +} + +func applyMessagesToResponsesInput(original any, msgs []Message) (any, error) { + switch original.(type) { + case nil: + if len(msgs) != 0 { + return nil, core.NewInvalidRequestError("guardrails cannot add or remove responses input items", nil) + } + return nil, nil + case string: + if len(msgs) != 1 { + return nil, core.NewInvalidRequestError("guardrails cannot add or remove responses input items", nil) + } + if msgs[0].Role != "user" { + return nil, core.NewInvalidRequestError("guardrails cannot change the role of a string responses input", nil) + } + if msgs[0].ContentNull { + return "", nil + } + return msgs[0].Content, nil + } + + elements, err := coerceResponsesInputElements(original) + if err != nil { + return nil, err + } + patched, err := applyMessagesToResponsesElements(elements, msgs) + if err != nil { + return nil, err + } + return patchResponsesInputEnvelope(original, patched) +} + +func applyMessagesToResponsesElements(elements []core.ResponsesInputElement, msgs []Message) ([]core.ResponsesInputElement, error) { + if len(msgs) != len(elements) { + return nil, core.NewInvalidRequestError("guardrails cannot add or remove responses input items", nil) + } + + result := make([]core.ResponsesInputElement, len(elements)) + for i, original := range elements { + patched, err := applyGuardedResponsesElementToOriginal(original, msgs[i], i) + if err != nil { + return nil, err + } + result[i] = patched + } + return result, nil +} + +func applyGuardedResponsesElementToOriginal(original core.ResponsesInputElement, modified Message, _ int) (core.ResponsesInputElement, error) { + preserved := original + + switch original.Type { + case "function_call": + if modified.Role != "assistant" { + return core.ResponsesInputElement{}, core.NewInvalidRequestError("guardrails cannot reorder or retag responses input items", nil) + } + return preserved, nil + case "function_call_output": + if modified.Role != "tool" { + return core.ResponsesInputElement{}, core.NewInvalidRequestError("guardrails cannot reorder or retag responses input items", nil) + } + preserved.Output = modified.Content + return preserved, nil + default: + role := strings.TrimSpace(original.Role) + if role == "" { + return core.ResponsesInputElement{}, core.NewInvalidRequestError("invalid responses input item: role is required", nil) + } + if modified.Role != role { + return core.ResponsesInputElement{}, core.NewInvalidRequestError("guardrails cannot reorder or retag responses input items", nil) + } + content, err := applyGuardedResponsesContentToOriginal(original.Content, modified.Content, modified.ContentNull) + if err != nil { + return core.ResponsesInputElement{}, core.NewInvalidRequestError("guardrails cannot merge rewritten text into responses input item", err) + } + preserved.Content = content + return preserved, nil + } +} + +func applyGuardedResponsesContentToOriginal(originalContent any, rewrittenText string, contentNull bool) (any, error) { + if isResponsesStructuredContent(originalContent) { + return rewriteStructuredResponsesContentWithTextRewrite(originalContent, rewrittenText) + } + if rewrittenText != "" { + contentNull = false + } + if contentNull { + return nil, nil + } + return rewrittenText, nil +} + +func patchResponsesInputEnvelope(original any, patched []core.ResponsesInputElement) (any, error) { + switch typed := original.(type) { + case []core.ResponsesInputElement: + result := make([]core.ResponsesInputElement, len(patched)) + copy(result, patched) + return result, nil + case []map[string]any: + return patchResponsesInputMapSlice(typed, patched) + case []any: + return patchResponsesInputInterfaceSlice(typed, patched) + default: + return patched, nil + } +} + +func patchResponsesInputMapSlice(original []map[string]any, patched []core.ResponsesInputElement) ([]map[string]any, error) { + if len(original) != len(patched) { + return nil, core.NewInvalidRequestError("guardrails cannot add or remove responses input items", nil) + } + result := make([]map[string]any, len(original)) + for i := range original { + item, err := patchResponsesInputMap(original[i], patched[i]) + if err != nil { + return nil, err + } + result[i] = item + } + return result, nil +} + +func patchResponsesInputInterfaceSlice(original []any, patched []core.ResponsesInputElement) ([]any, error) { + if len(original) != len(patched) { + return nil, core.NewInvalidRequestError("guardrails cannot add or remove responses input items", nil) + } + result := make([]any, len(original)) + for i := range original { + item, err := patchResponsesInputInterfaceElement(original[i], patched[i]) + if err != nil { + return nil, err + } + result[i] = item + } + return result, nil +} + +func patchResponsesInputInterfaceElement(original any, patched core.ResponsesInputElement) (any, error) { + if originalMap, ok := original.(map[string]any); ok { + return patchResponsesInputMap(originalMap, patched) + } + return responsesInputElementAsAny(patched) +} + +func patchResponsesInputMap(original map[string]any, patched core.ResponsesInputElement) (map[string]any, error) { + cloned := cloneStringAnyMap(original) + updated, err := responsesInputElementAsMap(patched) + if err != nil { + return nil, err + } + + for _, key := range []string{"type", "role", "status", "content", "call_id", "id", "name", "arguments", "output"} { + delete(cloned, key) + } + for key, value := range updated { + cloned[key] = value + } + if patched.Type == "function_call_output" { + cloned["output"] = restoreResponsesInputOutputValue(original["output"], patched.Output) + } + return cloned, nil +} + +func restoreResponsesInputOutputValue(original any, rewritten string) any { + if _, ok := original.(string); ok { + return rewritten + } + if strings.TrimSpace(rewritten) == "" { + if original == nil { + return nil + } + return original + } + + var decoded any + if err := json.Unmarshal([]byte(rewritten), &decoded); err == nil { + return decoded + } + if original == nil { + return nil + } + return original +} + +func responsesInputElementAsMap(element core.ResponsesInputElement) (map[string]any, error) { + value, err := responsesInputElementAsAny(element) + if err != nil { + return nil, err + } + itemMap, ok := value.(map[string]any) + if !ok { + return nil, core.NewInvalidRequestError("invalid responses input item", nil) + } + return itemMap, nil +} + +func responsesInputElementAsAny(element core.ResponsesInputElement) (any, error) { + raw, err := json.Marshal(element) + if err != nil { + return nil, core.NewInvalidRequestError("invalid responses input item", err) + } + var value any + if err := json.Unmarshal(raw, &value); err != nil { + return nil, core.NewInvalidRequestError("invalid responses input item", err) + } + return value, nil +} + +func isResponsesStructuredContent(content any) bool { + if content == nil { + return false + } + switch typed := content.(type) { + case []any: + return true + case []core.ContentPart: + return true + case []map[string]any: + return true + default: + _ = typed + } + contentType := reflect.TypeOf(content) + return contentType.Kind() == reflect.Slice || contentType.Kind() == reflect.Array +} + +func rewriteStructuredResponsesContentWithTextRewrite(originalContent any, rewrittenText string) (any, error) { + switch typed := originalContent.(type) { + case []core.ContentPart: + return rewriteStructuredResponsesTypedContentParts(typed, rewrittenText) + case []any: + return rewriteStructuredResponsesInterfaceContentParts(typed, rewrittenText) + case []map[string]any: + return rewriteStructuredResponsesMapContentParts(typed, rewrittenText) + default: + value := reflect.ValueOf(originalContent) + if !value.IsValid() || (value.Kind() != reflect.Slice && value.Kind() != reflect.Array) { + return nil, core.NewInvalidRequestError("unsupported structured responses content", nil) + } + parts := make([]any, value.Len()) + for i := 0; i < value.Len(); i++ { + parts[i] = value.Index(i).Interface() + } + return rewriteStructuredResponsesInterfaceContentParts(parts, rewrittenText) + } +} + +func rewriteStructuredResponsesTypedContentParts(parts []core.ContentPart, rewrittenText string) (any, error) { + textIndexes := make([]int, 0, len(parts)) + originalTexts := make([]string, 0, len(parts)) + for i, part := range parts { + if !isResponsesTextPartType(part.Type) || part.Text == "" { + continue + } + textIndexes = append(textIndexes, i) + originalTexts = append(originalTexts, part.Text) + } + + if len(textIndexes) == 0 { + if rewrittenText == "" { + if len(parts) == 0 { + return []core.ContentPart{}, nil + } + return cloneContentParts(parts), nil + } + prepended := []core.ContentPart{{Type: "input_text", Text: rewrittenText}} + prepended = append(prepended, cloneContentParts(parts)...) + return prepended, nil + } + + if len(textIndexes) == 1 { + merged := cloneContentParts(parts) + textIndex := textIndexes[0] + if rewrittenText == "" { + merged = append(merged[:textIndex], merged[textIndex+1:]...) + } else { + merged[textIndex].Text = rewrittenText + } + if len(merged) == 0 { + return nil, core.NewInvalidRequestError("guardrails produced empty structured responses content after rewrite", nil) + } + return merged, nil + } + + if rewrittenText == strings.Join(originalTexts, " ") { + return cloneContentParts(parts), nil + } + + merged := make([]core.ContentPart, 0, len(parts)) + insertedRewrittenText := false + for _, part := range parts { + if isResponsesTextPartType(part.Type) { + if !insertedRewrittenText && rewrittenText != "" { + rewrittenPart := cloneContentPart(part) + rewrittenPart.Text = rewrittenText + merged = append(merged, rewrittenPart) + insertedRewrittenText = true + } + continue + } + merged = append(merged, cloneContentPart(part)) + } + + if len(merged) == 0 { + return nil, core.NewInvalidRequestError("guardrails produced empty structured responses content after rewrite", nil) + } + return merged, nil +} + +func rewriteStructuredResponsesInterfaceContentParts(parts []any, rewrittenText string) (any, error) { + textIndexes := make([]int, 0, len(parts)) + originalTexts := make([]string, 0, len(parts)) + for i, part := range parts { + partMap, ok := part.(map[string]any) + if !ok { + continue + } + partType, _ := partMap["type"].(string) + if !isResponsesTextPartType(partType) { + continue + } + text, _ := partMap["text"].(string) + if text == "" { + continue + } + textIndexes = append(textIndexes, i) + originalTexts = append(originalTexts, text) + } + + if len(textIndexes) == 0 { + if rewrittenText == "" { + if len(parts) == 0 { + return []any{}, nil + } + return cloneResponsesInterfaceParts(parts), nil + } + prepended := []any{map[string]any{"type": "input_text", "text": rewrittenText}} + prepended = append(prepended, cloneResponsesInterfaceParts(parts)...) + return prepended, nil + } + + if len(textIndexes) == 1 { + merged := make([]any, 0, len(parts)) + textIndex := textIndexes[0] + for i, part := range parts { + if i == textIndex { + if rewrittenText == "" { + continue + } + partMap, ok := part.(map[string]any) + if !ok { + return nil, core.NewInvalidRequestError("guardrails cannot rewrite non-object responses content part", nil) + } + cloned := cloneStringAnyMap(partMap) + cloned["text"] = rewrittenText + merged = append(merged, cloned) + continue + } + merged = append(merged, cloneResponsesInterfacePart(part)) + } + if len(merged) == 0 { + return nil, core.NewInvalidRequestError("guardrails produced empty structured responses content after rewrite", nil) + } + return merged, nil + } + + if rewrittenText == strings.Join(originalTexts, " ") { + return cloneResponsesInterfaceParts(parts), nil + } + + merged := make([]any, 0, len(parts)) + insertedRewrittenText := false + for _, part := range parts { + partMap, ok := part.(map[string]any) + if ok { + partType, _ := partMap["type"].(string) + if isResponsesTextPartType(partType) { + if !insertedRewrittenText && rewrittenText != "" { + cloned := cloneStringAnyMap(partMap) + cloned["text"] = rewrittenText + merged = append(merged, cloned) + insertedRewrittenText = true + } + continue + } + } + merged = append(merged, cloneResponsesInterfacePart(part)) + } + + if len(merged) == 0 { + return nil, core.NewInvalidRequestError("guardrails produced empty structured responses content after rewrite", nil) + } + return merged, nil +} + +func rewriteStructuredResponsesMapContentParts(parts []map[string]any, rewrittenText string) (any, error) { + if len(parts) == 0 && rewrittenText == "" { + return []map[string]any{}, nil + } + + interfaceParts := make([]any, len(parts)) + for i, part := range parts { + interfaceParts[i] = part + } + rewritten, err := rewriteStructuredResponsesInterfaceContentParts(interfaceParts, rewrittenText) + if err != nil { + return nil, err + } + rewrittenParts, ok := rewritten.([]any) + if !ok { + return nil, core.NewInvalidRequestError("unsupported structured responses content", nil) + } + + result := make([]map[string]any, len(rewrittenParts)) + for i, part := range rewrittenParts { + partMap, ok := part.(map[string]any) + if !ok { + return nil, core.NewInvalidRequestError("guardrails cannot rewrite non-object responses content part", nil) + } + result[i] = cloneStringAnyMap(partMap) + } + return result, nil +} + +func isResponsesTextPartType(partType string) bool { + switch partType { + case "text", "input_text", "output_text": + return true + default: + return false + } +} + +func cloneResponsesInterfaceParts(parts []any) []any { + if len(parts) == 0 { + return nil + } + cloned := make([]any, len(parts)) + for i, part := range parts { + cloned[i] = cloneResponsesInterfacePart(part) + } + return cloned +} + +func cloneResponsesInterfacePart(part any) any { + partMap, ok := part.(map[string]any) + if !ok { + return part + } + return cloneStringAnyMap(partMap) +} + +// cloneStringAnyMap performs a shallow copy of the map. Nested maps/slices are +// intentionally shared; callers are expected to either preserve them as-is or +// replace whole top-level values instead of mutating nested structures in place. +func cloneStringAnyMap(src map[string]any) map[string]any { + if src == nil { + return nil + } + cloned := make(map[string]any, len(src)) + for key, value := range src { + cloned[key] = value + } + return cloned +} From 8bec890ca408ca3bb6436a2b641a37e06748bfaa Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Wed, 29 Apr 2026 23:33:08 +0200 Subject: [PATCH 07/15] refactor(admin): split handler.go by domain + extract route registration MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit internal/admin/handler.go was 2400 lines hosting endpoints for ten distinct domains (usage, audit, models, providers, budgets, aliases, auth keys, model overrides, guardrails, workflows). Split methods into peer files in the same package, mirroring the existing per-domain test file layout. handler.go (513) struct, options, NewHandler, runtime refresh types, parse/error helpers errors.go (141) feature/validation/write error helpers, deactivateByID, deleteByName, decode helpers routes.go (68) RegisterRoutes(RouteRegistrar) — admin RouteRegistrar interface so callers can mount under any route group handler_usage.go (398) UsageSummary, DailyUsage, UsageByModel, UsageByUserPath, UsageLog, CacheOverview, RecalculateUsagePricing + recalc helpers/types handler_audit.go (205) AuditLog, AuditConversation, auditLogResponse handler_models.go (129) ListModels, ListCategories, DashboardConfig handler_providers.go (169) ProviderStatus, RefreshRuntime, buildProviderStatusResponse, classifyProviderStatus handler_budgets.go (385) budget endpoints + types + helpers handler_aliases.go (91) alias endpoints + types handler_authkeys.go (81) auth-key endpoints + types handler_model_overrides.go (79) model-override endpoints + types handler_guardrails.go (118) guardrail endpoints + types handler_workflows.go (246) workflow endpoints + types + validation helpers internal/server/http.go: replaced the 40-line inline admin route block with `cfg.AdminHandler.RegisterRoutes(e.Group("/admin/api/v1"))`. Route registration now lives next to the methods it wires up. Pure relocation; full test suite passes. Co-Authored-By: Claude Opus 4.7 --- internal/admin/errors.go | 141 ++ internal/admin/handler.go | 1893 +-------------------- internal/admin/handler_aliases.go | 91 + internal/admin/handler_audit.go | 205 +++ internal/admin/handler_authkeys.go | 81 + internal/admin/handler_budgets.go | 385 +++++ internal/admin/handler_guardrails.go | 118 ++ internal/admin/handler_model_overrides.go | 79 + internal/admin/handler_models.go | 129 ++ internal/admin/handler_providers.go | 169 ++ internal/admin/handler_usage.go | 398 +++++ internal/admin/handler_workflows.go | 246 +++ internal/admin/routes.go | 68 + internal/server/http.go | 41 +- 14 files changed, 2114 insertions(+), 1930 deletions(-) create mode 100644 internal/admin/errors.go create mode 100644 internal/admin/handler_aliases.go create mode 100644 internal/admin/handler_audit.go create mode 100644 internal/admin/handler_authkeys.go create mode 100644 internal/admin/handler_budgets.go create mode 100644 internal/admin/handler_guardrails.go create mode 100644 internal/admin/handler_model_overrides.go create mode 100644 internal/admin/handler_models.go create mode 100644 internal/admin/handler_providers.go create mode 100644 internal/admin/handler_usage.go create mode 100644 internal/admin/handler_workflows.go create mode 100644 internal/admin/routes.go diff --git a/internal/admin/errors.go b/internal/admin/errors.go new file mode 100644 index 000000000..38ce81c63 --- /dev/null +++ b/internal/admin/errors.go @@ -0,0 +1,141 @@ +package admin + +import ( + "context" + "errors" + "net/http" + "net/url" + "strings" + + "github.com/labstack/echo/v5" + + "gomodel/internal/aliases" + "gomodel/internal/authkeys" + "gomodel/internal/budget" + "gomodel/internal/core" + "gomodel/internal/guardrails" + "gomodel/internal/modeloverrides" + "gomodel/internal/workflows" +) + +func budgetServiceError(message string, err error) error { + if errors.Is(err, budget.ErrNotFound) { + return core.NewNotFoundError("budget not found").WithCode("budget_not_found") + } + return core.NewProviderError("budgets", http.StatusServiceUnavailable, message, err) +} + +func featureUnavailableError(message string) error { + return core.NewInvalidRequestErrorWithStatus(http.StatusServiceUnavailable, message, nil). + WithCode("feature_unavailable") +} + +func validationWriter(isValidation func(error) bool) func(error) error { + return func(err error) error { + if err == nil { + return nil + } + if isValidation(err) { + return core.NewInvalidRequestError(err.Error(), err) + } + return err + } +} + +var ( + aliasWriteError = validationWriter(aliases.IsValidationError) + workflowWriteError = validationWriter(workflows.IsValidationError) + authKeyWriteError = validationWriter(authkeys.IsValidationError) + guardrailWriteError = validationWriter(guardrails.IsValidationError) +) + +// modelOverrideWriteError differs from the others: non-validation errors are +// surfaced as 502 so the dashboard distinguishes provider failures from input issues. +func modelOverrideWriteError(err error) error { + if err == nil { + return nil + } + if modeloverrides.IsValidationError(err) { + return core.NewInvalidRequestError(err.Error(), err) + } + return core.NewProviderError("model_overrides", http.StatusBadGateway, err.Error(), err) +} + +func deactivateByID( + c *echo.Context, + unavailableErr error, + idLabel string, + notFoundErr error, + notFoundMessage string, + deactivate func(context.Context, string) error, + writeError func(error) error, +) error { + if unavailableErr != nil { + return handleError(c, unavailableErr) + } + + id := strings.TrimSpace(c.Param("id")) + if id == "" { + return handleError(c, core.NewInvalidRequestError(idLabel+" id is required", nil)) + } + + if err := deactivate(c.Request().Context(), id); err != nil { + if errors.Is(err, notFoundErr) { + return handleError(c, core.NewNotFoundError(notFoundMessage+id)) + } + return handleError(c, writeError(err)) + } + return c.NoContent(http.StatusNoContent) +} + +func deleteByName( + c *echo.Context, + unavailableErr error, + paramName string, + decode func(string) (string, error), + deleteFunc func(context.Context, string) error, + notFoundErr error, + notFoundMessage string, + writeError func(error) error, +) error { + if unavailableErr != nil { + return handleError(c, unavailableErr) + } + + name, err := decode(c.Param(paramName)) + if err != nil { + return handleError(c, err) + } + + if err := deleteFunc(c.Request().Context(), name); err != nil { + if errors.Is(err, notFoundErr) { + return handleError(c, core.NewNotFoundError(notFoundMessage+name)) + } + return handleError(c, writeError(err)) + } + return c.NoContent(http.StatusNoContent) +} + +func decodeAliasPathName(raw string) (string, error) { + name, err := url.PathUnescape(strings.TrimSpace(raw)) + if err != nil { + return "", core.NewInvalidRequestError("invalid alias name", err) + } + name = strings.TrimSpace(name) + if name == "" { + return "", core.NewInvalidRequestError("alias name is required", nil) + } + return name, nil +} + +func decodeModelOverridePathSelector(raw string) (string, error) { + selector, err := url.PathUnescape(strings.TrimSpace(raw)) + if err != nil { + return "", core.NewInvalidRequestError("invalid model override selector", err) + } + selector = strings.TrimSpace(selector) + if selector == "" { + return "", core.NewInvalidRequestError("model override selector is required", nil) + } + return selector, nil +} diff --git a/internal/admin/handler.go b/internal/admin/handler.go index b6e1bcbbf..40f73e651 100644 --- a/internal/admin/handler.go +++ b/internal/admin/handler.go @@ -3,14 +3,9 @@ package admin import ( "context" - "encoding/json" "errors" - "fmt" "log/slog" "net/http" - "net/url" - "slices" - "sort" "strconv" "strings" "sync" @@ -507,1894 +502,12 @@ func requestIDFromAdminContextOrHeader(req *http.Request) string { return strings.TrimSpace(req.Header.Get("X-Request-ID")) } -// UsageSummary handles GET /admin/api/v1/usage/summary -// -// @Summary Get usage summary -// @Tags admin -// @Produce json -// @Security BearerAuth -// @Param days query int false "Number of days (default 30)" -// @Param start_date query string false "Start date (YYYY-MM-DD)" -// @Param end_date query string false "End date (YYYY-MM-DD)" -// @Param user_path query string false "Filter by tracked user path subtree" -// @Param cache_mode query string false "Cache mode filter: uncached, cached, all (default uncached)" -// @Success 200 {object} usage.UsageSummary -// @Failure 400 {object} core.GatewayError -// @Failure 401 {object} core.GatewayError -// @Router /admin/api/v1/usage/summary [get] -func (h *Handler) UsageSummary(c *echo.Context) error { - if h.usageReader == nil { - return c.JSON(http.StatusOK, usage.UsageSummary{}) - } - - params, err := parseUsageParams(c) - if err != nil { - return handleError(c, err) - } - - summary, err := h.usageReader.GetSummary(c.Request().Context(), params) - if err != nil { - return handleError(c, err) - } - - return c.JSON(http.StatusOK, summary) -} - -func usageSliceResponse[T any]( - c *echo.Context, - reader usage.UsageReader, - fetch func(context.Context, usage.UsageQueryParams) ([]T, error), -) error { - if reader == nil { - return c.JSON(http.StatusOK, []T{}) - } - - params, err := parseUsageParams(c) - if err != nil { - return handleError(c, err) - } - - values, err := fetch(c.Request().Context(), params) - if err != nil { - return handleError(c, err) - } - if values == nil { - values = []T{} - } - return c.JSON(http.StatusOK, values) -} - -// DailyUsage handles GET /admin/api/v1/usage/daily -// -// @Summary Get usage breakdown by period -// @Tags admin -// @Produce json -// @Security BearerAuth -// @Param days query int false "Number of days (default 30)" -// @Param start_date query string false "Start date (YYYY-MM-DD)" -// @Param end_date query string false "End date (YYYY-MM-DD)" -// @Param interval query string false "Grouping interval: daily, weekly, monthly, yearly (default daily)" -// @Param user_path query string false "Filter by tracked user path subtree" -// @Param cache_mode query string false "Cache mode filter: uncached, cached, all (default uncached)" -// @Success 200 {array} usage.DailyUsage -// @Failure 400 {object} core.GatewayError -// @Failure 401 {object} core.GatewayError -// @Router /admin/api/v1/usage/daily [get] -func (h *Handler) DailyUsage(c *echo.Context) error { - return usageSliceResponse(c, h.usageReader, func(ctx context.Context, params usage.UsageQueryParams) ([]usage.DailyUsage, error) { - return h.usageReader.GetDailyUsage(ctx, params) - }) -} - -// UsageByModel handles GET /admin/api/v1/usage/models -// -// @Summary Get usage breakdown by model -// @Tags admin -// @Produce json -// @Security BearerAuth -// @Param days query int false "Number of days (default 30)" -// @Param start_date query string false "Start date (YYYY-MM-DD)" -// @Param end_date query string false "End date (YYYY-MM-DD)" -// @Param user_path query string false "Filter by tracked user path subtree" -// @Param cache_mode query string false "Cache mode filter: uncached, cached, all (default uncached)" -// @Success 200 {array} usage.ModelUsage -// @Failure 400 {object} core.GatewayError -// @Failure 401 {object} core.GatewayError -// @Router /admin/api/v1/usage/models [get] -func (h *Handler) UsageByModel(c *echo.Context) error { - return usageSliceResponse(c, h.usageReader, func(ctx context.Context, params usage.UsageQueryParams) ([]usage.ModelUsage, error) { - return h.usageReader.GetUsageByModel(ctx, params) - }) -} - -// UsageByUserPath handles GET /admin/api/v1/usage/user-paths -// -// @Summary Get usage breakdown by user path -// @Tags admin -// @Produce json -// @Security BearerAuth -// @Param days query int false "Number of days (default 30)" -// @Param start_date query string false "Start date (YYYY-MM-DD)" -// @Param end_date query string false "End date (YYYY-MM-DD)" -// @Param user_path query string false "Filter by tracked user path subtree" -// @Param cache_mode query string false "Cache mode filter: uncached, cached, all (default uncached)" -// @Success 200 {array} usage.UserPathUsage -// @Failure 400 {object} core.GatewayError -// @Failure 401 {object} core.GatewayError -// @Router /admin/api/v1/usage/user-paths [get] -func (h *Handler) UsageByUserPath(c *echo.Context) error { - return usageSliceResponse(c, h.usageReader, func(ctx context.Context, params usage.UsageQueryParams) ([]usage.UserPathUsage, error) { - return h.usageReader.GetUsageByUserPath(ctx, params) - }) -} - -// UsageLog handles GET /admin/api/v1/usage/log -// -// @Summary Get paginated usage log entries -// @Tags admin -// @Produce json -// @Security BearerAuth -// @Param days query int false "Number of days (default 30)" -// @Param start_date query string false "Start date (YYYY-MM-DD)" -// @Param end_date query string false "End date (YYYY-MM-DD)" -// @Param model query string false "Filter by model name" -// @Param provider query string false "Filter by provider name or provider type" -// @Param user_path query string false "Filter by tracked user path subtree" -// @Param cache_mode query string false "Cache mode filter: uncached, cached, all (default uncached)" -// @Param search query string false "Search across model, provider, request_id, provider_id" -// @Param limit query int false "Page size (default 50, max 200)" -// @Param offset query int false "Offset for pagination" -// @Success 200 {object} usage.UsageLogResult -// @Failure 400 {object} core.GatewayError -// @Failure 401 {object} core.GatewayError -// @Router /admin/api/v1/usage/log [get] -func (h *Handler) UsageLog(c *echo.Context) error { - if h.usageReader == nil { - return c.JSON(http.StatusOK, usage.UsageLogResult{ - Entries: []usage.UsageLogEntry{}, - }) - } - - baseParams, err := parseUsageParams(c) - if err != nil { - return handleError(c, err) - } - - params := usage.UsageLogParams{ - UsageQueryParams: baseParams, - Model: c.QueryParam("model"), - Provider: c.QueryParam("provider"), - Search: c.QueryParam("search"), - } - - if l := c.QueryParam("limit"); l != "" { - if parsed, err := strconv.Atoi(l); err == nil && parsed > 0 { - params.Limit = parsed - } - } - if o := c.QueryParam("offset"); o != "" { - if parsed, err := strconv.Atoi(o); err == nil && parsed >= 0 { - params.Offset = parsed - } - } - - result, err := h.usageReader.GetUsageLog(c.Request().Context(), params) - if err != nil { - return handleError(c, err) - } - - if result.Entries == nil { - result.Entries = []usage.UsageLogEntry{} - } - - return c.JSON(http.StatusOK, result) -} - -// RecalculateUsagePricing handles POST /admin/api/v1/usage/recalculate-pricing. -// -// @Summary Recalculate stored usage costs from current model pricing metadata -// @Tags admin -// @Accept json -// @Produce json -// @Security BearerAuth -// @Param request body recalculatePricingRequest true "Recalculation filters and confirmation" -// @Success 200 {object} usage.RecalculatePricingResult -// @Failure 400 {object} core.GatewayError -// @Failure 401 {object} core.GatewayError -// @Failure 500 {object} core.GatewayError -// @Failure 503 {object} core.GatewayError -// @Router /admin/api/v1/usage/recalculate-pricing [post] -func (h *Handler) RecalculateUsagePricing(c *echo.Context) error { - if h.usageRecalculator == nil { - return handleError(c, featureUnavailableError("usage pricing recalculation is unavailable")) - } - if h.registry == nil { - return handleError(c, featureUnavailableError("model pricing metadata is unavailable")) - } - - var req recalculatePricingRequest - if err := c.Bind(&req); err != nil { - return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) - } - if strings.TrimSpace(strings.ToLower(req.confirmationValue())) != "recalculate" { - return handleError(c, core.NewInvalidRequestError("confirmation must be recalculate", nil)) - } - - params, err := h.recalculatePricingParams(c, req) - if err != nil { - return handleError(c, err) - } - - h.pricingMu.Lock() - defer h.pricingMu.Unlock() - - result, err := h.usageRecalculator.RecalculatePricing(c.Request().Context(), params, h.registry) - if err != nil { - if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { - return handleError(c, err) - } - if gatewayErr, ok := errors.AsType[*core.GatewayError](err); ok { - return handleError(c, gatewayErr) - } - return handleError(c, core.NewProviderError("usage", http.StatusInternalServerError, "failed to recalculate usage pricing", err)) - } - return c.JSON(http.StatusOK, result) -} - -// CacheOverview handles GET /admin/api/v1/cache/overview -// -// @Summary Get cached-only usage overview -// @Tags admin -// @Produce json -// @Security BearerAuth -// @Param days query int false "Number of days (default 30)" -// @Param start_date query string false "Start date (YYYY-MM-DD)" -// @Param end_date query string false "End date (YYYY-MM-DD)" -// @Param interval query string false "Grouping interval: daily, weekly, monthly, yearly (default daily)" -// @Param user_path query string false "Filter by tracked user path subtree" -// @Param cache_mode query string false "Cache mode filter: uncached, cached, all (cache overview always uses cached mode)" -// @Success 200 {object} usage.CacheOverview -// @Failure 400 {object} core.GatewayError -// @Failure 401 {object} core.GatewayError -// @Failure 503 {object} core.GatewayError -// @Router /admin/api/v1/cache/overview [get] -func (h *Handler) CacheOverview(c *echo.Context) error { - if strings.TrimSpace(h.runtimeConfig.CacheEnabled) != "on" { - return handleError(c, featureUnavailableError("cache analytics is unavailable")) - } - if h.usageReader == nil { - return c.JSON(http.StatusOK, usage.CacheOverview{ - Daily: []usage.CacheOverviewDaily{}, - }) - } - - params, err := parseUsageParams(c) - if err != nil { - return handleError(c, err) - } - params.CacheMode = usage.CacheModeCached - - overview, err := h.usageReader.GetCacheOverview(c.Request().Context(), params) - if err != nil { - return handleError(c, err) - } - if overview == nil { - overview = &usage.CacheOverview{} - } - if overview.Daily == nil { - overview.Daily = []usage.CacheOverviewDaily{} - } - - return c.JSON(http.StatusOK, overview) -} - -// AuditLog handles GET /admin/api/v1/audit/log -// -// @Summary Get paginated audit log entries -// @Tags admin -// @Produce json -// @Security BearerAuth -// @Param days query int false "Number of days (default 30)" -// @Param start_date query string false "Start date (YYYY-MM-DD)" -// @Param end_date query string false "End date (YYYY-MM-DD)" -// @Param requested_model query string false "Filter by requested model selector" -// @Param provider query string false "Filter by provider name or provider type" -// @Param method query string false "Filter by HTTP method" -// @Param path query string false "Filter by request path" -// @Param user_path query string false "Filter by tracked user path subtree" -// @Param error_type query string false "Filter by error type" -// @Param status_code query int false "Filter by status code" -// @Param stream query bool false "Filter by stream mode (true/false)" -// @Param search query string false "Search across request_id/requested_model/provider/method/path/error_type/error_message" -// @Param limit query int false "Page size (default 25, max 100)" -// @Param offset query int false "Offset for pagination" -// @Success 200 {object} auditLogListResponse -// @Failure 400 {object} core.GatewayError -// @Failure 401 {object} core.GatewayError -// @Router /admin/api/v1/audit/log [get] -func (h *Handler) AuditLog(c *echo.Context) error { - if h.auditReader == nil { - return c.JSON(http.StatusOK, auditLogListResponse{ - Entries: []auditLogEntryResponse{}, - }) - } - - dateRange, err := parseDateRangeParams(c) - if err != nil { - return handleError(c, err) - } - userPath, err := normalizeUserPathQueryParam("user_path", c.QueryParam("user_path")) - if err != nil { - return handleError(c, err) - } - - requestedModel := c.QueryParam("requested_model") - if requestedModel == "" { - requestedModel = c.QueryParam("model") - } - - params := auditlog.LogQueryParams{ - QueryParams: auditlog.QueryParams{ - StartDate: dateRange.StartDate, - EndDate: dateRange.EndDate, - }, - RequestedModel: requestedModel, - Provider: c.QueryParam("provider"), - Method: strings.ToUpper(c.QueryParam("method")), - Path: c.QueryParam("path"), - UserPath: userPath, - ErrorType: c.QueryParam("error_type"), - Search: c.QueryParam("search"), - } - - if sc := c.QueryParam("status_code"); sc != "" { - parsed, err := strconv.Atoi(sc) - if err != nil { - return handleError(c, core.NewInvalidRequestError("invalid status_code, expected integer", nil)) - } - params.StatusCode = &parsed - } - - if stream := c.QueryParam("stream"); stream != "" { - parsed, err := strconv.ParseBool(stream) - if err != nil { - return handleError(c, core.NewInvalidRequestError("invalid stream value, expected true or false", nil)) - } - params.Stream = &parsed - } - - if l := c.QueryParam("limit"); l != "" { - if parsed, err := strconv.Atoi(l); err == nil && parsed > 0 { - params.Limit = parsed - } - } - if o := c.QueryParam("offset"); o != "" { - if parsed, err := strconv.Atoi(o); err == nil && parsed >= 0 { - params.Offset = parsed - } - } - - result, err := h.auditReader.GetLogs(c.Request().Context(), params) - if err != nil { - return handleError(c, err) - } - - if result.Entries == nil { - result.Entries = []auditlog.LogEntry{} - } - - response, err := h.auditLogResponse(c.Request().Context(), result) - if err != nil { - return handleError(c, err) - } - return c.JSON(http.StatusOK, response) -} - -func (h *Handler) auditLogResponse(ctx context.Context, result *auditlog.LogListResult) (*auditLogListResponse, error) { - if result == nil { - return &auditLogListResponse{Entries: []auditLogEntryResponse{}}, nil - } - - response := &auditLogListResponse{ - Entries: make([]auditLogEntryResponse, len(result.Entries)), - Total: result.Total, - Limit: result.Limit, - Offset: result.Offset, - } - for i := range result.Entries { - response.Entries[i].LogEntry = result.Entries[i] - } - - if h.usageReader == nil || len(result.Entries) == 0 { - return response, nil - } - - requestIDs := make([]string, 0, len(result.Entries)) - for _, entry := range result.Entries { - requestIDs = append(requestIDs, entry.RequestID) - } - - entriesByRequestID, err := h.usageReader.GetUsageByRequestIDs(ctx, requestIDs) - if err != nil { - slog.Warn("failed to enrich audit log entries with usage", "error", err, "request_count", len(requestIDs)) - return response, nil - } - - summaries := usage.SummarizeUsageByRequestID(entriesByRequestID) - for i := range response.Entries { - requestID := response.Entries[i].RequestID - if summary, ok := summaries[requestID]; ok { - response.Entries[i].Usage = summary - } - } - - return response, nil -} - -// AuditConversation handles GET /admin/api/v1/audit/conversation -// -// @Summary Get conversation thread around an audit log entry -// @Tags admin -// @Produce json -// @Security BearerAuth -// @Param log_id query string true "Anchor audit log entry ID" -// @Param limit query int false "Max entries in thread (default 40, max 200)" -// @Success 200 {object} auditlog.ConversationResult -// @Failure 400 {object} core.GatewayError -// @Failure 401 {object} core.GatewayError -// @Router /admin/api/v1/audit/conversation [get] -func (h *Handler) AuditConversation(c *echo.Context) error { - if h.auditReader == nil { - return c.JSON(http.StatusOK, auditlog.ConversationResult{ - AnchorID: c.QueryParam("log_id"), - Entries: []auditlog.LogEntry{}, - }) - } - - logID := strings.TrimSpace(c.QueryParam("log_id")) - if logID == "" { - return handleError(c, core.NewInvalidRequestError("log_id is required", nil)) - } - - limit := 40 - if l := c.QueryParam("limit"); l != "" { - parsed, err := strconv.Atoi(l) - if err != nil { - return handleError(c, core.NewInvalidRequestError("invalid limit, expected integer", nil)) - } - if parsed < 1 || parsed > 200 { - return handleError(c, core.NewInvalidRequestError("invalid limit parameter: limit must be between 1 and 200", nil)) - } - limit = parsed - } - - result, err := h.auditReader.GetConversation(c.Request().Context(), logID, limit) - if err != nil { - return handleError(c, err) - } - if result == nil { - result = &auditlog.ConversationResult{ - AnchorID: logID, - Entries: []auditlog.LogEntry{}, - } - } - if result.Entries == nil { - result.Entries = []auditlog.LogEntry{} - } - - return c.JSON(http.StatusOK, result) -} - -// ListModels handles GET /admin/api/v1/models -// Supports optional ?category= query param for filtering by model category. -// -// @Summary List all registered models with provider info -// @Tags admin -// @Produce json -// @Security BearerAuth -// @Success 200 {array} providers.ModelWithProvider -// @Failure 401 {object} core.GatewayError -// @Router /admin/api/v1/models [get] -type modelAccessResponse struct { - Selector string `json:"selector"` - DefaultEnabled bool `json:"default_enabled"` - EffectiveEnabled bool `json:"effective_enabled"` - UserPaths []string `json:"user_paths,omitempty"` - Override *modeloverrides.Override `json:"override,omitempty"` -} - -type modelInventoryResponse struct { - providers.ModelWithProvider - Access modelAccessResponse `json:"access"` -} - -func (h *Handler) ListModels(c *echo.Context) error { - if h.registry == nil { - return c.JSON(http.StatusOK, []modelInventoryResponse{}) - } - - cat := core.ModelCategory(c.QueryParam("category")) - if cat != "" && cat != core.CategoryAll { - if !isValidCategory(cat) { - return handleError(c, core.NewInvalidRequestError("invalid category: "+string(cat), nil)) - } - } - - var models []providers.ModelWithProvider - if cat != "" && cat != core.CategoryAll { - models = h.registry.ListModelsWithProviderByCategory(cat) - } else { - models = h.registry.ListModelsWithProvider() - } - - if models == nil { - models = []providers.ModelWithProvider{} - } - if h.modelOverrides == nil { - response := make([]modelInventoryResponse, 0, len(models)) - for _, model := range models { - selector := core.ModelSelector{ - Provider: strings.TrimSpace(model.ProviderName), - Model: strings.TrimSpace(model.Model.ID), - } - response = append(response, modelInventoryResponse{ - ModelWithProvider: model, - Access: modelAccessResponse{ - Selector: selector.QualifiedModel(), - DefaultEnabled: true, - EffectiveEnabled: true, - }, - }) - } - return c.JSON(http.StatusOK, response) - } - - response := make([]modelInventoryResponse, 0, len(models)) - for _, model := range models { - selector := core.ModelSelector{ - Provider: strings.TrimSpace(model.ProviderName), - Model: strings.TrimSpace(model.Model.ID), - } - effective := h.modelOverrides.EffectiveState(selector) - access := modelAccessResponse{ - Selector: effective.Selector, - DefaultEnabled: effective.DefaultEnabled, - EffectiveEnabled: effective.Enabled, - UserPaths: append([]string(nil), effective.UserPaths...), - } - if override, ok := h.modelOverrides.Get(selector.QualifiedModel()); ok && override != nil { - overrideCopy := *override - access.Override = &overrideCopy - } - response = append(response, modelInventoryResponse{ - ModelWithProvider: model, - Access: access, - }) - } - - return c.JSON(http.StatusOK, response) -} - -// isValidCategory returns true if cat is a recognized model category. -func isValidCategory(cat core.ModelCategory) bool { - return slices.Contains(core.AllCategories(), cat) -} - -// ListCategories handles GET /admin/api/v1/models/categories -// -// @Summary List model categories with counts -// @Tags admin -// @Produce json -// @Security BearerAuth -// @Success 200 {array} providers.CategoryCount -// @Failure 401 {object} core.GatewayError -// @Router /admin/api/v1/models/categories [get] -func (h *Handler) ListCategories(c *echo.Context) error { - if h.registry == nil { - return c.JSON(http.StatusOK, []providers.CategoryCount{}) - } - - return c.JSON(http.StatusOK, h.registry.GetCategoryCounts()) -} - -// DashboardConfig handles GET /admin/api/v1/dashboard/config -func (h *Handler) DashboardConfig(c *echo.Context) error { - return c.JSON(http.StatusOK, cloneDashboardRuntimeConfig(h.runtimeConfig)) -} - -// ListBudgets handles GET /admin/api/v1/budgets. -// @Summary List budgets with current status -// @Tags admin -// @Produce json -// @Security BearerAuth -// @Success 200 {object} budgetListResponse -// @Failure 401 {object} core.GatewayError -// @Failure 503 {object} core.GatewayError -// @Router /admin/api/v1/budgets [get] -func (h *Handler) ListBudgets(c *echo.Context) error { - if h.budgets == nil { - return handleError(c, featureUnavailableError("budgets feature is unavailable")) - } - now := time.Now().UTC() - statuses, err := h.budgets.Statuses(c.Request().Context(), now) - if err != nil { - return handleError(c, budgetServiceError("failed to list budgets", err)) - } - return c.JSON(http.StatusOK, budgetListResponse{ - Budgets: budgetStatusResponses(statuses, now), - ServerTime: now, - }) -} - -// UpsertBudget handles PUT /admin/api/v1/budgets/{user_path}/{period}. -// @Summary Create or update one budget -// @Tags admin -// @Accept json -// @Produce json -// @Security BearerAuth -// @Param user_path path string true "URL-encoded budget user path" -// @Param period path string true "Budget period name or seconds" -// @Param budget body upsertBudgetRequest true "Budget amount" -// @Success 200 {object} budgetListResponse -// @Failure 400 {object} core.GatewayError -// @Failure 401 {object} core.GatewayError -// @Failure 503 {object} core.GatewayError -// @Router /admin/api/v1/budgets/{user_path}/{period} [put] -func (h *Handler) UpsertBudget(c *echo.Context) error { - if h.budgets == nil { - return handleError(c, featureUnavailableError("budgets feature is unavailable")) - } - var req upsertBudgetRequest - if err := c.Bind(&req); err != nil { - return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) - } - userPath, periodSeconds, err := budgetRouteKey(c) - if err != nil { - return handleError(c, core.NewInvalidRequestError(err.Error(), err)) - } - item, err := budget.NormalizeBudget(budget.Budget{ - UserPath: userPath, - PeriodSeconds: periodSeconds, - Amount: req.Amount, - Source: budget.SourceManual, - }) - if err != nil { - return handleError(c, core.NewInvalidRequestError(err.Error(), err)) - } - if err := h.budgets.UpsertBudgets(c.Request().Context(), []budget.Budget{item}); err != nil { - return handleError(c, budgetServiceError("failed to save budget", err)) - } - return h.ListBudgets(c) -} - -// DeleteBudget handles DELETE /admin/api/v1/budgets/{user_path}/{period}. -// @Summary Delete one budget -// @Tags admin -// @Produce json -// @Security BearerAuth -// @Param user_path path string true "URL-encoded budget user path" -// @Param period path string true "Budget period name or seconds" -// @Success 200 {object} budgetListResponse -// @Failure 400 {object} core.GatewayError -// @Failure 401 {object} core.GatewayError -// @Failure 503 {object} core.GatewayError -// @Router /admin/api/v1/budgets/{user_path}/{period} [delete] -func (h *Handler) DeleteBudget(c *echo.Context) error { - if h.budgets == nil { - return handleError(c, featureUnavailableError("budgets feature is unavailable")) - } - userPath, periodSeconds, err := budgetRouteKey(c) - if err != nil { - return handleError(c, core.NewInvalidRequestError(err.Error(), err)) - } - if err := h.budgets.DeleteBudget(c.Request().Context(), userPath, periodSeconds); err != nil { - return handleError(c, budgetServiceError("failed to delete budget", err)) - } - return h.ListBudgets(c) -} - -// BudgetSettings handles GET /admin/api/v1/budgets/settings. -// @Summary Get budget reset settings -// @Tags admin -// @Produce json -// @Security BearerAuth -// @Success 200 {object} budget.Settings -// @Failure 401 {object} core.GatewayError -// @Failure 503 {object} core.GatewayError -// @Router /admin/api/v1/budgets/settings [get] -func (h *Handler) BudgetSettings(c *echo.Context) error { - if h.budgets == nil { - return handleError(c, featureUnavailableError("budgets feature is unavailable")) - } - return c.JSON(http.StatusOK, h.budgets.Settings()) -} - -// UpdateBudgetSettings handles PUT /admin/api/v1/budgets/settings. -// @Summary Update budget reset settings -// @Tags admin -// @Accept json -// @Produce json -// @Security BearerAuth -// @Param settings body updateBudgetSettingsRequest true "Budget reset settings" -// @Success 200 {object} budget.Settings -// @Failure 400 {object} core.GatewayError -// @Failure 401 {object} core.GatewayError -// @Failure 503 {object} core.GatewayError -// @Router /admin/api/v1/budgets/settings [put] -func (h *Handler) UpdateBudgetSettings(c *echo.Context) error { - if h.budgets == nil { - return handleError(c, featureUnavailableError("budgets feature is unavailable")) - } - var req updateBudgetSettingsRequest - if err := c.Bind(&req); err != nil { - return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) - } - settings := req.apply(h.budgets.Settings()) - if err := budget.ValidateSettings(settings); err != nil { - return handleError(c, core.NewInvalidRequestError(err.Error(), err)) - } - saved, err := h.budgets.SaveSettings(c.Request().Context(), settings) - if err != nil { - return handleError(c, budgetServiceError("failed to save budget settings", err)) - } - return c.JSON(http.StatusOK, saved) -} - -// ResetBudget handles POST /admin/api/v1/budgets/reset-one. -// @Summary Reset one budget period -// @Tags admin -// @Accept json -// @Produce json -// @Security BearerAuth -// @Param budget body resetBudgetRequest true "Budget key" -// @Success 200 {object} budgetListResponse -// @Failure 400 {object} core.GatewayError -// @Failure 401 {object} core.GatewayError -// @Failure 503 {object} core.GatewayError -// @Router /admin/api/v1/budgets/reset-one [post] -func (h *Handler) ResetBudget(c *echo.Context) error { - if h.budgets == nil { - return handleError(c, featureUnavailableError("budgets feature is unavailable")) - } - var req resetBudgetRequest - if err := c.Bind(&req); err != nil { - return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) - } - periodSeconds, err := budgetRequestPeriodSeconds(req.Period, req.PeriodSeconds) - if err != nil { - return handleError(c, core.NewInvalidRequestError(err.Error(), err)) - } - userPath, err := budget.NormalizeUserPath(req.UserPath) - if err != nil { - return handleError(c, core.NewInvalidRequestError(err.Error(), err)) - } - if err := h.budgets.ResetBudget(c.Request().Context(), userPath, periodSeconds, time.Now().UTC()); err != nil { - return handleError(c, budgetServiceError("failed to reset budget", err)) - } - return h.ListBudgets(c) -} - -// ResetBudgets handles POST /admin/api/v1/budgets/reset. -// @Summary Reset all budget periods -// @Tags admin -// @Accept json -// @Produce json -// @Security BearerAuth -// @Param confirmation body resetBudgetsRequest true "Reset confirmation" -// @Success 200 {object} resetBudgetsResponse -// @Failure 400 {object} core.GatewayError -// @Failure 401 {object} core.GatewayError -// @Failure 503 {object} core.GatewayError -// @Router /admin/api/v1/budgets/reset [post] -func (h *Handler) ResetBudgets(c *echo.Context) error { - if h.budgets == nil { - return handleError(c, featureUnavailableError("budgets feature is unavailable")) - } - var req resetBudgetsRequest - if err := c.Bind(&req); err != nil { - return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) - } - if strings.TrimSpace(strings.ToLower(req.confirmationValue())) != "reset" { - return handleError(c, core.NewInvalidRequestError("confirmation must be reset", nil)) - } - if err := h.budgets.ResetAll(c.Request().Context(), time.Now().UTC()); err != nil { - return handleError(c, budgetServiceError("failed to reset budgets", err)) - } - return c.JSON(http.StatusOK, resetBudgetsResponse{Status: "ok"}) -} - -// ProviderStatus handles GET /admin/api/v1/providers/status -func (h *Handler) ProviderStatus(c *echo.Context) error { - return c.JSON(http.StatusOK, h.buildProviderStatusResponse()) -} - -// RefreshRuntime handles POST /admin/api/v1/runtime/refresh -func (h *Handler) RefreshRuntime(c *echo.Context) error { - if h.runtimeRefresher == nil { - return handleError(c, featureUnavailableError("runtime refresh is unavailable")) - } - - report, err := h.runtimeRefresher.RefreshRuntime(c.Request().Context()) - if err != nil { - if gatewayErr, ok := errors.AsType[*core.GatewayError](err); ok { - return handleError(c, gatewayErr) - } - return handleError(c, core.NewProviderError("runtime_refresh", http.StatusInternalServerError, "runtime refresh failed", err)) - } - if report.Status == "" { - report.Status = RuntimeRefreshStatusOK - } - if report.Steps == nil { - report.Steps = []RuntimeRefreshStep{} - } - return c.JSON(http.StatusOK, report) -} - -func (h *Handler) buildProviderStatusResponse() providerStatusResponse { - configured := cloneConfiguredProviders(h.configuredProviders) - configuredByName := make(map[string]providers.SanitizedProviderConfig, len(configured)) - nameSet := make(map[string]struct{}, len(configured)) - for _, cfg := range configured { - name := strings.TrimSpace(cfg.Name) - if name == "" { - continue - } - configuredByName[name] = cfg - nameSet[name] = struct{}{} - } - - runtimeByName := make(map[string]providers.ProviderRuntimeSnapshot) - if h.registry != nil { - for _, snapshot := range h.registry.ProviderRuntimeSnapshots() { - name := strings.TrimSpace(snapshot.Name) - if name == "" { - continue - } - runtimeByName[name] = snapshot - nameSet[name] = struct{}{} - } - } - - names := make([]string, 0, len(nameSet)) - for name := range nameSet { - names = append(names, name) - } - sort.Strings(names) - - resp := providerStatusResponse{ - Summary: providerStatusSummaryResponse{ - OverallStatus: "degraded", - }, - Providers: make([]providerStatusItemResponse, 0, len(names)), - } - - for _, name := range names { - cfg, hasConfig := configuredByName[name] - runtime, hasRuntime := runtimeByName[name] - if !hasConfig { - cfg = providers.SanitizedProviderConfig{Name: name, Type: strings.TrimSpace(runtime.Type)} - } - if !hasRuntime { - runtime = providers.ProviderRuntimeSnapshot{Name: name, Type: strings.TrimSpace(cfg.Type)} - } - if strings.TrimSpace(cfg.Type) == "" { - cfg.Type = strings.TrimSpace(runtime.Type) - } - if strings.TrimSpace(runtime.Type) == "" { - runtime.Type = strings.TrimSpace(cfg.Type) - } - - status, label, reason, lastError := classifyProviderStatus(cfg, runtime) - resp.Providers = append(resp.Providers, providerStatusItemResponse{ - Name: name, - Type: strings.TrimSpace(cfg.Type), - Status: status, - StatusLabel: label, - StatusReason: reason, - LastError: lastError, - Config: cfg, - Runtime: runtime, - }) - resp.Summary.Total++ - switch status { - case "healthy": - resp.Summary.Healthy++ - case "unhealthy": - resp.Summary.Unhealthy++ - default: - resp.Summary.Degraded++ - } - } - - switch { - case resp.Summary.Total == 0: - resp.Summary.OverallStatus = "degraded" - case resp.Summary.Healthy == resp.Summary.Total: - resp.Summary.OverallStatus = "healthy" - case resp.Summary.Unhealthy == resp.Summary.Total: - resp.Summary.OverallStatus = "unhealthy" - default: - resp.Summary.OverallStatus = "degraded" - } - - if resp.Providers == nil { - resp.Providers = []providerStatusItemResponse{} - } - return resp -} - -func classifyProviderStatus(cfg providers.SanitizedProviderConfig, runtime providers.ProviderRuntimeSnapshot) (status, label, reason, lastError string) { - modelFetchError := strings.TrimSpace(runtime.LastModelFetchError) - availabilityError := strings.TrimSpace(runtime.LastAvailabilityError) - configuredName := strings.TrimSpace(cfg.Name) - usingCachedModels := runtime.Registered && - runtime.DiscoveredModelCount > 0 && - modelFetchError == "" && - runtime.LastModelFetchSuccessAt == nil - - lastError = modelFetchError - if lastError == "" { - lastError = availabilityError - } - - switch { - case runtime.DiscoveredModelCount > 0 && modelFetchError == "": - if usingCachedModels { - return "degraded", "Starting", "serving cached model inventory while live refresh finishes", lastError - } - return "healthy", "Healthy", "configured and model discovery succeeded", lastError - case modelFetchError != "" && runtime.DiscoveredModelCount > 0: - return "degraded", "Degraded", "latest model refresh failed; previous inventory is still available", lastError - case modelFetchError != "": - return "unhealthy", "Unhealthy", "model discovery failed and no provider models are currently available", lastError - case availabilityError != "" && runtime.DiscoveredModelCount == 0: - return "unhealthy", "Unhealthy", "startup availability check failed and no provider models are available", lastError - case runtime.DiscoveredModelCount > 0: - return "healthy", "Healthy", "provider models are currently available", lastError - case !runtime.Registered && configuredName != "": - return "degraded", "Starting", "provider is configured and awaiting live model discovery", lastError - case configuredName != "": - return "degraded", "Configured", "provider is configured but has not exposed models yet", lastError - default: - return "degraded", "Unknown", "provider runtime inventory is unavailable", lastError - } -} - -type upsertAliasRequest struct { - TargetModel string `json:"target_model"` - TargetProvider string `json:"target_provider,omitempty"` - Description string `json:"description,omitempty"` - Enabled *bool `json:"enabled,omitempty"` -} - -type upsertModelOverrideRequest struct { - UserPaths []string `json:"user_paths,omitempty"` -} - -type upsertGuardrailRequest struct { - Type string `json:"type"` - Description string `json:"description,omitempty"` - UserPath string `json:"user_path,omitempty"` - Config json.RawMessage `json:"config"` -} - -type createWorkflowRequest struct { - ScopeProviderName string `json:"scope_provider_name,omitempty"` - LegacyScopeProvider string `json:"scope_provider,omitempty"` - ScopeModel string `json:"scope_model,omitempty"` - ScopeUserPath string `json:"scope_user_path,omitempty"` - Name string `json:"name"` - Description string `json:"description,omitempty"` - Payload workflows.Payload `json:"workflow_payload"` -} - -type createAuthKeyRequest struct { - Name string `json:"name"` - Description string `json:"description,omitempty"` - UserPath string `json:"user_path,omitempty"` - ExpiresAt *time.Time `json:"expires_at,omitempty"` -} - -type budgetListResponse struct { - Budgets []budgetStatusResponse `json:"budgets"` - ServerTime time.Time `json:"server_time"` -} - -type budgetStatusResponse struct { - UserPath string `json:"user_path"` - PeriodSeconds int64 `json:"period_seconds"` - PeriodLabel string `json:"period_label"` - Amount float64 `json:"amount"` - Source string `json:"source,omitempty"` - LastResetAt *time.Time `json:"last_reset_at,omitempty"` - CreatedAt time.Time `json:"created_at"` - UpdatedAt time.Time `json:"updated_at"` - PeriodStart time.Time `json:"period_start"` - PeriodEnd time.Time `json:"period_end"` - Spent float64 `json:"spent"` - HasUsage bool `json:"has_usage"` - Remaining float64 `json:"remaining"` - UsageRatio float64 `json:"usage_ratio"` - PeriodRatio float64 `json:"period_ratio"` -} - -type upsertBudgetRequest struct { - Amount float64 `json:"amount"` -} - -type resetBudgetRequest struct { - UserPath string `json:"user_path"` - Period string `json:"period,omitempty"` - PeriodSeconds int64 `json:"period_seconds,omitempty"` -} - -type updateBudgetSettingsRequest struct { - DailyResetHour *int `json:"daily_reset_hour"` - DailyResetMinute *int `json:"daily_reset_minute"` - WeeklyResetWeekday *int `json:"weekly_reset_weekday"` - WeeklyResetHour *int `json:"weekly_reset_hour"` - WeeklyResetMinute *int `json:"weekly_reset_minute"` - MonthlyResetDay *int `json:"monthly_reset_day"` - MonthlyResetHour *int `json:"monthly_reset_hour"` - MonthlyResetMinute *int `json:"monthly_reset_minute"` -} - -type recalculatePricingRequest struct { - Days int `json:"days,omitempty"` - StartDate string `json:"start_date,omitempty"` - EndDate string `json:"end_date,omitempty"` - UserPath string `json:"user_path,omitempty"` - Selector string `json:"selector,omitempty"` - Confirmation string `json:"confirmation"` - Confirm string `json:"confirm,omitempty"` -} - -func (r recalculatePricingRequest) confirmationValue() string { - if strings.TrimSpace(r.Confirmation) != "" { - return r.Confirmation - } - return r.Confirm -} - -func (r updateBudgetSettingsRequest) apply(settings budget.Settings) budget.Settings { - if r.DailyResetHour != nil { - settings.DailyResetHour = *r.DailyResetHour - } - if r.DailyResetMinute != nil { - settings.DailyResetMinute = *r.DailyResetMinute - } - if r.WeeklyResetWeekday != nil { - settings.WeeklyResetWeekday = *r.WeeklyResetWeekday - } - if r.WeeklyResetHour != nil { - settings.WeeklyResetHour = *r.WeeklyResetHour - } - if r.WeeklyResetMinute != nil { - settings.WeeklyResetMinute = *r.WeeklyResetMinute - } - if r.MonthlyResetDay != nil { - settings.MonthlyResetDay = *r.MonthlyResetDay - } - if r.MonthlyResetHour != nil { - settings.MonthlyResetHour = *r.MonthlyResetHour - } - if r.MonthlyResetMinute != nil { - settings.MonthlyResetMinute = *r.MonthlyResetMinute - } - return settings -} - -type resetBudgetsRequest struct { - Confirmation string `json:"confirmation"` - Confirm string `json:"confirm,omitempty"` -} - -func (r resetBudgetsRequest) confirmationValue() string { - if strings.TrimSpace(r.Confirmation) != "" { - return r.Confirmation - } - return r.Confirm -} - -type resetBudgetsResponse struct { - Status string `json:"status"` -} - -func (h *Handler) recalculatePricingParams(c *echo.Context, req recalculatePricingRequest) (usage.RecalculatePricingParams, error) { - baseParams, err := recalculatePricingDateParams(c, req) - if err != nil { - return usage.RecalculatePricingParams{}, err - } - - userPath, err := normalizeUserPathQueryParam("user_path", req.UserPath) - if err != nil { - return usage.RecalculatePricingParams{}, err - } - baseParams.UserPath = userPath - baseParams.CacheMode = usage.CacheModeAll - - provider, model, err := h.recalculatePricingSelector(req.Selector) - if err != nil { - return usage.RecalculatePricingParams{}, err - } - - return usage.RecalculatePricingParams{ - UsageQueryParams: baseParams, - Provider: provider, - Model: model, - }, nil -} - -func recalculatePricingDateParams(c *echo.Context, req recalculatePricingRequest) (usage.UsageQueryParams, error) { - var params usage.UsageQueryParams - - timeZone, location := dashboardTimeZone(c) - params.TimeZone = timeZone - - now := timeNow().In(location) - today := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, location) - - start, end, err := buildDateRange(strings.TrimSpace(req.StartDate), strings.TrimSpace(req.EndDate), normalizeDateRangeDays(req.Days), location, today) - if err != nil { - return params, err - } - params.StartDate = start - params.EndDate = end - return params, nil -} - -func (h *Handler) recalculatePricingSelector(raw string) (provider, model string, err error) { - raw = strings.TrimSpace(raw) - if raw == "" { - return "", "", nil - } - - if h.aliases != nil { - selector, changed, err := h.aliases.ResolveModel(core.NewRequestedModelSelector(raw, "")) - if err != nil { - return "", "", core.NewInvalidRequestError("invalid selector: "+err.Error(), err) - } - if changed { - return selector.Provider, selector.Model, nil - } - } - - selector, err := core.ParseModelSelector(raw, "") - if err != nil { - return "", "", core.NewInvalidRequestError("invalid selector: "+err.Error(), err) - } - if selector.Provider == "" { - return "", "", core.NewInvalidRequestError("invalid selector: provider/model or alias is required", nil) - } - return selector.Provider, selector.Model, nil -} - -func budgetStatusResponses(statuses []budget.CheckResult, now time.Time) []budgetStatusResponse { - if len(statuses) == 0 { - return []budgetStatusResponse{} - } - responses := make([]budgetStatusResponse, 0, len(statuses)) - for _, status := range statuses { - item := status.Budget - usageRatio := 0.0 - if item.Amount > 0 { - usageRatio = status.Spent / item.Amount - } - periodRatio := 0.0 - periodDuration := status.PeriodEnd.Sub(status.PeriodStart).Seconds() - if periodDuration > 0 { - periodRatio = now.Sub(status.PeriodStart).Seconds() / periodDuration - } - responses = append(responses, budgetStatusResponse{ - UserPath: item.UserPath, - PeriodSeconds: item.PeriodSeconds, - PeriodLabel: budget.PeriodLabel(item.PeriodSeconds), - Amount: item.Amount, - Source: item.Source, - LastResetAt: item.LastResetAt, - CreatedAt: item.CreatedAt, - UpdatedAt: item.UpdatedAt, - PeriodStart: status.PeriodStart, - PeriodEnd: status.PeriodEnd, - Spent: status.Spent, - HasUsage: status.HasUsage, - Remaining: status.Remaining, - UsageRatio: usageRatio, - PeriodRatio: clampBudgetRatio(periodRatio), - }) - } - return responses -} - -func budgetRouteKey(c *echo.Context) (string, int64, error) { - userPathParam := strings.TrimSpace(c.Param("user_path")) - if userPathParam == "" { - return "", 0, errors.New("user_path path parameter is required") - } - userPath, err := url.PathUnescape(userPathParam) - if err != nil { - return "", 0, fmt.Errorf("invalid user_path path parameter: %w", err) - } - userPath, err = budget.NormalizeUserPath(userPath) - if err != nil { - return "", 0, err - } - - periodParam := strings.TrimSpace(c.Param("period")) - if periodParam == "" { - return "", 0, errors.New("period path parameter is required") - } - if seconds, err := strconv.ParseInt(periodParam, 10, 64); err == nil { - if seconds <= 0 { - return "", 0, errors.New("period_seconds must be greater than 0") - } - return userPath, seconds, nil - } - periodSeconds, err := budgetRequestPeriodSeconds(periodParam, 0) - if err != nil { - return "", 0, err - } - return userPath, periodSeconds, nil -} - -func budgetRequestPeriodSeconds(period string, periodSeconds int64) (int64, error) { - if periodSeconds > 0 { - return periodSeconds, nil - } - if parsed, ok := budget.PeriodSeconds(period); ok { - return parsed, nil - } - return 0, errors.New("period must be one of hourly, daily, weekly, monthly or period_seconds must be set") -} - -func clampBudgetRatio(value float64) float64 { - if value < 0 { - return 0 - } - if value > 1 { - return 1 - } - return value -} - -func budgetServiceError(message string, err error) error { - if errors.Is(err, budget.ErrNotFound) { - return core.NewNotFoundError("budget not found").WithCode("budget_not_found") - } - return core.NewProviderError("budgets", http.StatusServiceUnavailable, message, err) -} - -func featureUnavailableError(message string) error { - return core.NewInvalidRequestErrorWithStatus(http.StatusServiceUnavailable, message, nil). - WithCode("feature_unavailable") -} - -func validationWriter(isValidation func(error) bool) func(error) error { - return func(err error) error { - if err == nil { - return nil - } - if isValidation(err) { - return core.NewInvalidRequestError(err.Error(), err) - } - return err - } -} - -var ( - aliasWriteError = validationWriter(aliases.IsValidationError) - workflowWriteError = validationWriter(workflows.IsValidationError) - authKeyWriteError = validationWriter(authkeys.IsValidationError) - guardrailWriteError = validationWriter(guardrails.IsValidationError) -) - -// modelOverrideWriteError differs from the others: non-validation errors are -// surfaced as 502 so the dashboard distinguishes provider failures from input issues. -func modelOverrideWriteError(err error) error { - if err == nil { - return nil - } - if modeloverrides.IsValidationError(err) { - return core.NewInvalidRequestError(err.Error(), err) - } - return core.NewProviderError("model_overrides", http.StatusBadGateway, err.Error(), err) -} - -func deactivateByID( - c *echo.Context, - unavailableErr error, - idLabel string, - notFoundErr error, - notFoundMessage string, - deactivate func(context.Context, string) error, - writeError func(error) error, -) error { - if unavailableErr != nil { - return handleError(c, unavailableErr) - } - - id := strings.TrimSpace(c.Param("id")) - if id == "" { - return handleError(c, core.NewInvalidRequestError(idLabel+" id is required", nil)) - } - - if err := deactivate(c.Request().Context(), id); err != nil { - if errors.Is(err, notFoundErr) { - return handleError(c, core.NewNotFoundError(notFoundMessage+id)) - } - return handleError(c, writeError(err)) - } - return c.NoContent(http.StatusNoContent) -} - -func deleteByName( - c *echo.Context, - unavailableErr error, - paramName string, - decode func(string) (string, error), - deleteFunc func(context.Context, string) error, - notFoundErr error, - notFoundMessage string, - writeError func(error) error, -) error { - if unavailableErr != nil { - return handleError(c, unavailableErr) - } - - name, err := decode(c.Param(paramName)) - if err != nil { - return handleError(c, err) - } - - if err := deleteFunc(c.Request().Context(), name); err != nil { - if errors.Is(err, notFoundErr) { - return handleError(c, core.NewNotFoundError(notFoundMessage+name)) - } - return handleError(c, writeError(err)) - } - return c.NoContent(http.StatusNoContent) -} +// @Param stream query bool false "Filter by stream mode (true/false)" +// @Router /admin/api/v1/models [get] +// @Router /admin/api/v1/budgets [get] // ListModelOverrides handles GET /admin/api/v1/model-overrides. -func (h *Handler) ListModelOverrides(c *echo.Context) error { - if h.modelOverrides == nil { - return handleError(c, featureUnavailableError("model overrides feature is unavailable")) - } - views := h.modelOverrides.ListViews() - if views == nil { - views = []modeloverrides.View{} - } - return c.JSON(http.StatusOK, views) -} - -// UpsertModelOverride handles PUT /admin/api/v1/model-overrides/{selector}. -func (h *Handler) UpsertModelOverride(c *echo.Context) error { - if h.modelOverrides == nil { - return handleError(c, featureUnavailableError("model overrides feature is unavailable")) - } - - selector, err := decodeModelOverridePathSelector(c.Param("selector")) - if err != nil { - return handleError(c, err) - } - - var req upsertModelOverrideRequest - if err := c.Bind(&req); err != nil { - return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) - } - - if err := h.modelOverrides.Upsert(c.Request().Context(), modeloverrides.Override{ - Selector: selector, - UserPaths: req.UserPaths, - }); err != nil { - return handleError(c, modelOverrideWriteError(err)) - } - - override, ok := h.modelOverrides.Get(selector) - if !ok || override == nil { - slog.Error("model override service returned no override after upsert", "selector", selector) - return handleError(c, core.NewProviderError("model_overrides", http.StatusInternalServerError, "model override update failed unexpectedly", nil)) - } - return c.JSON(http.StatusOK, override) -} - -// DeleteModelOverride handles DELETE /admin/api/v1/model-overrides/{selector}. -func (h *Handler) DeleteModelOverride(c *echo.Context) error { - var unavailableErr error - var deleteFunc func(context.Context, string) error - if h.modelOverrides == nil { - unavailableErr = featureUnavailableError("model overrides feature is unavailable") - } else { - deleteFunc = h.modelOverrides.Delete - } - return deleteByName( - c, - unavailableErr, - "selector", - decodeModelOverridePathSelector, - deleteFunc, - modeloverrides.ErrNotFound, - "model override not found: ", - modelOverrideWriteError, - ) -} - // ListAuthKeys handles GET /admin/api/v1/auth-keys -func (h *Handler) ListAuthKeys(c *echo.Context) error { - if h.authKeys == nil { - return handleError(c, featureUnavailableError("auth keys feature is unavailable")) - } - views := h.authKeys.ListViews() - if views == nil { - views = []authkeys.View{} - } - return c.JSON(http.StatusOK, views) -} - -// CreateAuthKey handles POST /admin/api/v1/auth-keys -func (h *Handler) CreateAuthKey(c *echo.Context) error { - if h.authKeys == nil { - return handleError(c, featureUnavailableError("auth keys feature is unavailable")) - } - - var req createAuthKeyRequest - if err := c.Bind(&req); err != nil { - return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) - } - - userPath, err := normalizeUserPathQueryParam("user_path", req.UserPath) - if err != nil { - return handleError(c, err) - } - - issued, err := h.authKeys.Create(c.Request().Context(), authkeys.CreateInput{ - Name: req.Name, - Description: req.Description, - UserPath: userPath, - ExpiresAt: req.ExpiresAt, - }) - if err != nil { - return handleError(c, authKeyWriteError(err)) - } - if issued == nil { - requestID := strings.TrimSpace(core.GetRequestID(c.Request().Context())) - slog.Error("auth key service returned nil issued key", "request_id", requestID, "path", c.Request().URL.Path) - return c.JSON(http.StatusInternalServerError, (&core.GatewayError{ - Type: core.ErrorType("internal_error"), - Message: "auth key creation failed unexpectedly", - StatusCode: http.StatusInternalServerError, - }).WithCode("auth_key_issue_failed").ToJSON()) - } - return c.JSON(http.StatusCreated, issued) -} - -// DeactivateAuthKey handles POST /admin/api/v1/auth-keys/:id/deactivate -func (h *Handler) DeactivateAuthKey(c *echo.Context) error { - var unavailableErr error - var deactivate func(context.Context, string) error - if h.authKeys == nil { - unavailableErr = featureUnavailableError("auth keys feature is unavailable") - } else { - deactivate = h.authKeys.Deactivate - } - return deactivateByID(c, unavailableErr, "auth key", authkeys.ErrNotFound, "auth key not found: ", deactivate, authKeyWriteError) -} - // ListAliases handles GET /admin/api/v1/aliases -func (h *Handler) ListAliases(c *echo.Context) error { - if h.aliases == nil { - return handleError(c, featureUnavailableError("aliases feature is unavailable")) - } - views := h.aliases.ListViews() - if views == nil { - views = []aliases.View{} - } - return c.JSON(http.StatusOK, views) -} - -// UpsertAlias handles PUT /admin/api/v1/aliases/{name} -func (h *Handler) UpsertAlias(c *echo.Context) error { - if h.aliases == nil { - return handleError(c, featureUnavailableError("aliases feature is unavailable")) - } - - name, err := decodeAliasPathName(c.Param("name")) - if err != nil { - return handleError(c, err) - } - - var req upsertAliasRequest - if err := c.Bind(&req); err != nil { - return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) - } - - enabled := true - if existing, ok := h.aliases.Get(name); ok && existing != nil { - enabled = existing.Enabled - } - if req.Enabled != nil { - enabled = *req.Enabled - } - - if err := h.aliases.Upsert(c.Request().Context(), aliases.Alias{ - Name: name, - TargetModel: req.TargetModel, - TargetProvider: req.TargetProvider, - Description: req.Description, - Enabled: enabled, - }); err != nil { - return handleError(c, aliasWriteError(err)) - } - - alias, ok := h.aliases.Get(name) - if !ok { - return c.NoContent(http.StatusNoContent) - } - return c.JSON(http.StatusOK, alias) -} - -// DeleteAlias handles DELETE /admin/api/v1/aliases/{name} -func (h *Handler) DeleteAlias(c *echo.Context) error { - var unavailableErr error - var deleteFunc func(context.Context, string) error - if h.aliases == nil { - unavailableErr = featureUnavailableError("aliases feature is unavailable") - } else { - deleteFunc = h.aliases.Delete - } - return deleteByName( - c, - unavailableErr, - "name", - decodeAliasPathName, - deleteFunc, - aliases.ErrNotFound, - "alias not found: ", - aliasWriteError, - ) -} - // ListGuardrailTypes handles GET /admin/api/v1/guardrails/types -func (h *Handler) ListGuardrailTypes(c *echo.Context) error { - if h.guardrailDefs == nil { - return handleError(c, featureUnavailableError("guardrails feature is unavailable")) - } - return c.JSON(http.StatusOK, h.guardrailDefs.TypeDefinitions()) -} - -// ListGuardrails handles GET /admin/api/v1/guardrails -func (h *Handler) ListGuardrails(c *echo.Context) error { - if h.guardrailDefs == nil { - return handleError(c, featureUnavailableError("guardrails feature is unavailable")) - } - views := h.guardrailDefs.ListViews() - if views == nil { - views = []guardrails.View{} - } - return c.JSON(http.StatusOK, views) -} - -// UpsertGuardrail handles PUT /admin/api/v1/guardrails/{name} -func (h *Handler) UpsertGuardrail(c *echo.Context) error { - if h.guardrailDefs == nil { - return handleError(c, featureUnavailableError("guardrails feature is unavailable")) - } - - name := strings.TrimSpace(c.Param("name")) - if name == "" { - return handleError(c, core.NewInvalidRequestError("guardrail name is required", nil)) - } - - var req upsertGuardrailRequest - if err := c.Bind(&req); err != nil { - return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) - } - - userPath, err := normalizeUserPathQueryParam("user_path", req.UserPath) - if err != nil { - return handleError(c, err) - } - - h.mutationMu.Lock() - defer h.mutationMu.Unlock() - - if err := h.guardrailDefs.Upsert(c.Request().Context(), guardrails.Definition{ - Name: name, - Type: req.Type, - Description: req.Description, - UserPath: userPath, - Config: req.Config, - }); err != nil { - return handleError(c, guardrailWriteError(err)) - } - if err := h.refreshWorkflowsAfterGuardrailChange(c.Request().Context()); err != nil { - return handleError(c, err) - } - - definition, ok := h.guardrailDefs.Get(name) - if !ok { - return c.NoContent(http.StatusNoContent) - } - return c.JSON(http.StatusOK, guardrails.ViewFromDefinition(*definition)) -} - -// DeleteGuardrail handles DELETE /admin/api/v1/guardrails/{name} -func (h *Handler) DeleteGuardrail(c *echo.Context) error { - if h.guardrailDefs == nil { - return handleError(c, featureUnavailableError("guardrails feature is unavailable")) - } - - name := strings.TrimSpace(c.Param("name")) - if name == "" { - return handleError(c, core.NewInvalidRequestError("guardrail name is required", nil)) - } - - h.mutationMu.Lock() - defer h.mutationMu.Unlock() - - referencingWorkflows, err := h.activeWorkflowGuardrailReferences(c.Request().Context(), name) - if err != nil { - return handleError(c, err) - } - if len(referencingWorkflows) > 0 { - return handleError(c, core.NewInvalidRequestError("guardrail is used by active workflows: "+strings.Join(referencingWorkflows, ", "), nil)) - } - - if err := h.guardrailDefs.Delete(c.Request().Context(), name); err != nil { - if errors.Is(err, guardrails.ErrNotFound) { - return handleError(c, core.NewNotFoundError("guardrail not found: "+name)) - } - return handleError(c, guardrailWriteError(err)) - } - if err := h.refreshWorkflowsAfterGuardrailChange(c.Request().Context()); err != nil { - return handleError(c, err) - } - - return c.NoContent(http.StatusNoContent) -} - // ListWorkflows handles GET /admin/api/v1/workflows -func (h *Handler) ListWorkflows(c *echo.Context) error { - if h.workflows == nil { - return handleError(c, featureUnavailableError("workflows feature is unavailable")) - } - - views, err := h.workflows.ListViews(c.Request().Context()) - if err != nil { - return handleError(c, err) - } - if views == nil { - views = []workflows.View{} - } - return c.JSON(http.StatusOK, views) -} - -// GetWorkflow handles GET /admin/api/v1/workflows/:id -func (h *Handler) GetWorkflow(c *echo.Context) error { - if h.workflows == nil { - return handleError(c, featureUnavailableError("workflows feature is unavailable")) - } - - id := strings.TrimSpace(c.Param("id")) - if id == "" { - return handleError(c, core.NewInvalidRequestError("workflow id is required", nil)) - } - - view, err := h.workflows.GetView(c.Request().Context(), id) - if err != nil { - if errors.Is(err, workflows.ErrNotFound) { - return handleError(c, core.NewNotFoundError("workflow not found: "+id)) - } - return handleError(c, err) - } - - return c.JSON(http.StatusOK, view) -} - -// ListWorkflowGuardrails handles GET /admin/api/v1/workflows/guardrails -func (h *Handler) ListWorkflowGuardrails(c *echo.Context) error { - if h.guardrails == nil { - return c.JSON(http.StatusOK, []string{}) - } - - return c.JSON(http.StatusOK, h.guardrails.Names()) -} - -// CreateWorkflow handles POST /admin/api/v1/workflows -func (h *Handler) CreateWorkflow(c *echo.Context) error { - if h.workflows == nil { - return handleError(c, featureUnavailableError("workflows feature is unavailable")) - } - - var req createWorkflowRequest - if err := c.Bind(&req); err != nil { - return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) - } - scopeProviderName := strings.TrimSpace(req.ScopeProviderName) - if scopeProviderName == "" { - scopeProviderName = strings.TrimSpace(req.LegacyScopeProvider) - } - scopeModel := strings.TrimSpace(req.ScopeModel) - - scopeUserPath, err := normalizeUserPathQueryParam("scope_user_path", req.ScopeUserPath) - if err != nil { - return handleError(c, err) - } - - scopeProviderName, err = h.validateWorkflowScope(scopeProviderName, scopeModel) - if err != nil { - return handleError(c, err) - } - - if err := h.validateWorkflowGuardrails(req.Payload); err != nil { - return handleError(c, err) - } - - h.mutationMu.Lock() - defer h.mutationMu.Unlock() - - version, err := h.workflows.Create(c.Request().Context(), workflows.CreateInput{ - Scope: workflows.Scope{ - Provider: scopeProviderName, - Model: scopeModel, - UserPath: scopeUserPath, - }, - Activate: true, - Name: req.Name, - Description: req.Description, - Payload: req.Payload, - }) - if err != nil { - return handleError(c, workflowWriteError(err)) - } - if version == nil { - return c.NoContent(http.StatusNoContent) - } - return c.JSON(http.StatusCreated, version) -} - -// DeactivateWorkflow handles POST /admin/api/v1/workflows/:id/deactivate -func (h *Handler) DeactivateWorkflow(c *echo.Context) error { - if h.workflows == nil { - return handleError(c, featureUnavailableError("workflows feature is unavailable")) - } - - id := strings.TrimSpace(c.Param("id")) - if id == "" { - return handleError(c, core.NewInvalidRequestError("workflow id is required", nil)) - } - - h.mutationMu.Lock() - defer h.mutationMu.Unlock() - - if err := h.workflows.Deactivate(c.Request().Context(), id); err != nil { - if errors.Is(err, workflows.ErrNotFound) { - return handleError(c, core.NewNotFoundError("workflow not found: "+id)) - } - return handleError(c, workflowWriteError(err)) - } - return c.NoContent(http.StatusNoContent) -} - -func (h *Handler) refreshWorkflowsAfterGuardrailChange(ctx context.Context) error { - if h.workflows == nil { - return nil - } - if err := h.workflows.Refresh(ctx); err != nil { - return err - } - return nil -} - -func (h *Handler) activeWorkflowGuardrailReferences(ctx context.Context, name string) ([]string, error) { - if h.workflows == nil { - return nil, nil - } - - name = strings.TrimSpace(name) - if name == "" { - return nil, nil - } - - views, err := h.workflows.ListViews(ctx) - if err != nil { - return nil, err - } - - references := make([]string, 0) - for _, view := range views { - if !view.Payload.Features.Guardrails { - continue - } - for _, step := range view.Payload.Guardrails { - if strings.TrimSpace(step.Ref) != name { - continue - } - references = append(references, view.ScopeDisplay) - break - } - } - sort.Strings(references) - return references, nil -} - -func (h *Handler) validateWorkflowGuardrails(payload workflows.Payload) error { - if !payload.Features.Guardrails || len(payload.Guardrails) == 0 { - return nil - } - if h.guardrails == nil { - return featureUnavailableError("guardrail registry is unavailable for workflow authoring") - } - - known := make(map[string]struct{}, h.guardrails.Len()) - for _, name := range h.guardrails.Names() { - known[name] = struct{}{} - } - for _, step := range payload.Guardrails { - ref := strings.TrimSpace(step.Ref) - if ref == "" { - continue - } - if _, ok := known[ref]; !ok { - return core.NewInvalidRequestError("unknown guardrail ref: "+ref, nil) - } - } - return nil -} - -func (h *Handler) validateWorkflowScope(scopeProviderName, scopeModel string) (string, error) { - scopeProviderName = strings.TrimSpace(scopeProviderName) - scopeModel = strings.TrimSpace(scopeModel) - - if scopeProviderName == "" { - if scopeModel != "" { - return "", core.NewInvalidRequestError("scope_model requires scope_provider_name", nil) - } - return "", nil - } - if h.registry == nil { - return "", core.NewInvalidRequestError("provider registry is unavailable for workflow provider-name validation", nil) - } - if !slices.Contains(h.registry.ProviderNames(), scopeProviderName) { - if resolvedProviderName := strings.TrimSpace(h.registry.GetProviderNameForType(scopeProviderName)); resolvedProviderName != "" { - scopeProviderName = resolvedProviderName - } - } - if !slices.Contains(h.registry.ProviderNames(), scopeProviderName) { - return "", core.NewInvalidRequestError("unknown provider name: "+scopeProviderName, nil) - } - if scopeModel == "" { - return scopeProviderName, nil - } - - for _, model := range h.registry.ListModelsWithProvider() { - if model.ProviderName == scopeProviderName && model.Model.ID == scopeModel { - return scopeProviderName, nil - } - } - return "", core.NewInvalidRequestError("unknown model for provider name "+scopeProviderName+": "+scopeModel, nil) -} - -func decodeAliasPathName(raw string) (string, error) { - name, err := url.PathUnescape(strings.TrimSpace(raw)) - if err != nil { - return "", core.NewInvalidRequestError("invalid alias name", err) - } - name = strings.TrimSpace(name) - if name == "" { - return "", core.NewInvalidRequestError("alias name is required", nil) - } - return name, nil -} - -func decodeModelOverridePathSelector(raw string) (string, error) { - selector, err := url.PathUnescape(strings.TrimSpace(raw)) - if err != nil { - return "", core.NewInvalidRequestError("invalid model override selector", err) - } - selector = strings.TrimSpace(selector) - if selector == "" { - return "", core.NewInvalidRequestError("model override selector is required", nil) - } - return selector, nil -} diff --git a/internal/admin/handler_aliases.go b/internal/admin/handler_aliases.go new file mode 100644 index 000000000..1dfd0ce53 --- /dev/null +++ b/internal/admin/handler_aliases.go @@ -0,0 +1,91 @@ +package admin + +import ( + "context" + "net/http" + + "github.com/labstack/echo/v5" + + "gomodel/internal/aliases" + "gomodel/internal/core" +) + +type upsertAliasRequest struct { + TargetModel string `json:"target_model"` + TargetProvider string `json:"target_provider,omitempty"` + Description string `json:"description,omitempty"` + Enabled *bool `json:"enabled,omitempty"` +} + +func (h *Handler) ListAliases(c *echo.Context) error { + if h.aliases == nil { + return handleError(c, featureUnavailableError("aliases feature is unavailable")) + } + views := h.aliases.ListViews() + if views == nil { + views = []aliases.View{} + } + return c.JSON(http.StatusOK, views) +} + +// UpsertAlias handles PUT /admin/api/v1/aliases/{name} +func (h *Handler) UpsertAlias(c *echo.Context) error { + if h.aliases == nil { + return handleError(c, featureUnavailableError("aliases feature is unavailable")) + } + + name, err := decodeAliasPathName(c.Param("name")) + if err != nil { + return handleError(c, err) + } + + var req upsertAliasRequest + if err := c.Bind(&req); err != nil { + return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) + } + + enabled := true + if existing, ok := h.aliases.Get(name); ok && existing != nil { + enabled = existing.Enabled + } + if req.Enabled != nil { + enabled = *req.Enabled + } + + if err := h.aliases.Upsert(c.Request().Context(), aliases.Alias{ + Name: name, + TargetModel: req.TargetModel, + TargetProvider: req.TargetProvider, + Description: req.Description, + Enabled: enabled, + }); err != nil { + return handleError(c, aliasWriteError(err)) + } + + alias, ok := h.aliases.Get(name) + if !ok { + return c.NoContent(http.StatusNoContent) + } + return c.JSON(http.StatusOK, alias) +} + +// DeleteAlias handles DELETE /admin/api/v1/aliases/{name} +func (h *Handler) DeleteAlias(c *echo.Context) error { + var unavailableErr error + var deleteFunc func(context.Context, string) error + if h.aliases == nil { + unavailableErr = featureUnavailableError("aliases feature is unavailable") + } else { + deleteFunc = h.aliases.Delete + } + return deleteByName( + c, + unavailableErr, + "name", + decodeAliasPathName, + deleteFunc, + aliases.ErrNotFound, + "alias not found: ", + aliasWriteError, + ) +} diff --git a/internal/admin/handler_audit.go b/internal/admin/handler_audit.go new file mode 100644 index 000000000..b4633c7ea --- /dev/null +++ b/internal/admin/handler_audit.go @@ -0,0 +1,205 @@ +package admin + +import ( + "context" + "log/slog" + "net/http" + "strconv" + "strings" + + "github.com/labstack/echo/v5" + + "gomodel/internal/auditlog" + "gomodel/internal/core" + "gomodel/internal/usage" +) + +// @Param search query string false "Search across request_id/requested_model/provider/method/path/error_type/error_message" +// @Param limit query int false "Page size (default 25, max 100)" +// @Param offset query int false "Offset for pagination" +// @Success 200 {object} auditLogListResponse +// @Failure 400 {object} core.GatewayError +// @Failure 401 {object} core.GatewayError +// @Router /admin/api/v1/audit/log [get] +func (h *Handler) AuditLog(c *echo.Context) error { + if h.auditReader == nil { + return c.JSON(http.StatusOK, auditLogListResponse{ + Entries: []auditLogEntryResponse{}, + }) + } + + dateRange, err := parseDateRangeParams(c) + if err != nil { + return handleError(c, err) + } + userPath, err := normalizeUserPathQueryParam("user_path", c.QueryParam("user_path")) + if err != nil { + return handleError(c, err) + } + + requestedModel := c.QueryParam("requested_model") + if requestedModel == "" { + requestedModel = c.QueryParam("model") + } + + params := auditlog.LogQueryParams{ + QueryParams: auditlog.QueryParams{ + StartDate: dateRange.StartDate, + EndDate: dateRange.EndDate, + }, + RequestedModel: requestedModel, + Provider: c.QueryParam("provider"), + Method: strings.ToUpper(c.QueryParam("method")), + Path: c.QueryParam("path"), + UserPath: userPath, + ErrorType: c.QueryParam("error_type"), + Search: c.QueryParam("search"), + } + + if sc := c.QueryParam("status_code"); sc != "" { + parsed, err := strconv.Atoi(sc) + if err != nil { + return handleError(c, core.NewInvalidRequestError("invalid status_code, expected integer", nil)) + } + params.StatusCode = &parsed + } + + if stream := c.QueryParam("stream"); stream != "" { + parsed, err := strconv.ParseBool(stream) + if err != nil { + return handleError(c, core.NewInvalidRequestError("invalid stream value, expected true or false", nil)) + } + params.Stream = &parsed + } + + if l := c.QueryParam("limit"); l != "" { + if parsed, err := strconv.Atoi(l); err == nil && parsed > 0 { + params.Limit = parsed + } + } + if o := c.QueryParam("offset"); o != "" { + if parsed, err := strconv.Atoi(o); err == nil && parsed >= 0 { + params.Offset = parsed + } + } + + result, err := h.auditReader.GetLogs(c.Request().Context(), params) + if err != nil { + return handleError(c, err) + } + + if result.Entries == nil { + result.Entries = []auditlog.LogEntry{} + } + + response, err := h.auditLogResponse(c.Request().Context(), result) + if err != nil { + return handleError(c, err) + } + return c.JSON(http.StatusOK, response) +} + +func (h *Handler) auditLogResponse(ctx context.Context, result *auditlog.LogListResult) (*auditLogListResponse, error) { + if result == nil { + return &auditLogListResponse{Entries: []auditLogEntryResponse{}}, nil + } + + response := &auditLogListResponse{ + Entries: make([]auditLogEntryResponse, len(result.Entries)), + Total: result.Total, + Limit: result.Limit, + Offset: result.Offset, + } + for i := range result.Entries { + response.Entries[i].LogEntry = result.Entries[i] + } + + if h.usageReader == nil || len(result.Entries) == 0 { + return response, nil + } + + requestIDs := make([]string, 0, len(result.Entries)) + for _, entry := range result.Entries { + requestIDs = append(requestIDs, entry.RequestID) + } + + entriesByRequestID, err := h.usageReader.GetUsageByRequestIDs(ctx, requestIDs) + if err != nil { + slog.Warn("failed to enrich audit log entries with usage", "error", err, "request_count", len(requestIDs)) + return response, nil + } + + summaries := usage.SummarizeUsageByRequestID(entriesByRequestID) + for i := range response.Entries { + requestID := response.Entries[i].RequestID + if summary, ok := summaries[requestID]; ok { + response.Entries[i].Usage = summary + } + } + + return response, nil +} + +// AuditConversation handles GET /admin/api/v1/audit/conversation +// +// @Summary Get conversation thread around an audit log entry +// @Tags admin +// @Produce json +// @Security BearerAuth +// @Param log_id query string true "Anchor audit log entry ID" +// @Param limit query int false "Max entries in thread (default 40, max 200)" +// @Success 200 {object} auditlog.ConversationResult +// @Failure 400 {object} core.GatewayError +// @Failure 401 {object} core.GatewayError +// @Router /admin/api/v1/audit/conversation [get] +func (h *Handler) AuditConversation(c *echo.Context) error { + if h.auditReader == nil { + return c.JSON(http.StatusOK, auditlog.ConversationResult{ + AnchorID: c.QueryParam("log_id"), + Entries: []auditlog.LogEntry{}, + }) + } + + logID := strings.TrimSpace(c.QueryParam("log_id")) + if logID == "" { + return handleError(c, core.NewInvalidRequestError("log_id is required", nil)) + } + + limit := 40 + if l := c.QueryParam("limit"); l != "" { + parsed, err := strconv.Atoi(l) + if err != nil { + return handleError(c, core.NewInvalidRequestError("invalid limit, expected integer", nil)) + } + if parsed < 1 || parsed > 200 { + return handleError(c, core.NewInvalidRequestError("invalid limit parameter: limit must be between 1 and 200", nil)) + } + limit = parsed + } + + result, err := h.auditReader.GetConversation(c.Request().Context(), logID, limit) + if err != nil { + return handleError(c, err) + } + if result == nil { + result = &auditlog.ConversationResult{ + AnchorID: logID, + Entries: []auditlog.LogEntry{}, + } + } + if result.Entries == nil { + result.Entries = []auditlog.LogEntry{} + } + + return c.JSON(http.StatusOK, result) +} + +// ListModels handles GET /admin/api/v1/models +// Supports optional ?category= query param for filtering by model category. +// +// @Summary List all registered models with provider info +// @Tags admin +// @Produce json +// @Security BearerAuth +// @Success 200 {array} providers.ModelWithProvider +// @Failure 401 {object} core.GatewayError diff --git a/internal/admin/handler_authkeys.go b/internal/admin/handler_authkeys.go new file mode 100644 index 000000000..7e4b9bf04 --- /dev/null +++ b/internal/admin/handler_authkeys.go @@ -0,0 +1,81 @@ +package admin + +import ( + "context" + "log/slog" + "net/http" + "strings" + "time" + + "github.com/labstack/echo/v5" + + "gomodel/internal/authkeys" + "gomodel/internal/core" +) + +type createAuthKeyRequest struct { + Name string `json:"name"` + Description string `json:"description,omitempty"` + UserPath string `json:"user_path,omitempty"` + ExpiresAt *time.Time `json:"expires_at,omitempty"` +} + +func (h *Handler) ListAuthKeys(c *echo.Context) error { + if h.authKeys == nil { + return handleError(c, featureUnavailableError("auth keys feature is unavailable")) + } + views := h.authKeys.ListViews() + if views == nil { + views = []authkeys.View{} + } + return c.JSON(http.StatusOK, views) +} + +// CreateAuthKey handles POST /admin/api/v1/auth-keys +func (h *Handler) CreateAuthKey(c *echo.Context) error { + if h.authKeys == nil { + return handleError(c, featureUnavailableError("auth keys feature is unavailable")) + } + + var req createAuthKeyRequest + if err := c.Bind(&req); err != nil { + return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) + } + + userPath, err := normalizeUserPathQueryParam("user_path", req.UserPath) + if err != nil { + return handleError(c, err) + } + + issued, err := h.authKeys.Create(c.Request().Context(), authkeys.CreateInput{ + Name: req.Name, + Description: req.Description, + UserPath: userPath, + ExpiresAt: req.ExpiresAt, + }) + if err != nil { + return handleError(c, authKeyWriteError(err)) + } + if issued == nil { + requestID := strings.TrimSpace(core.GetRequestID(c.Request().Context())) + slog.Error("auth key service returned nil issued key", "request_id", requestID, "path", c.Request().URL.Path) + return c.JSON(http.StatusInternalServerError, (&core.GatewayError{ + Type: core.ErrorType("internal_error"), + Message: "auth key creation failed unexpectedly", + StatusCode: http.StatusInternalServerError, + }).WithCode("auth_key_issue_failed").ToJSON()) + } + return c.JSON(http.StatusCreated, issued) +} + +// DeactivateAuthKey handles POST /admin/api/v1/auth-keys/:id/deactivate +func (h *Handler) DeactivateAuthKey(c *echo.Context) error { + var unavailableErr error + var deactivate func(context.Context, string) error + if h.authKeys == nil { + unavailableErr = featureUnavailableError("auth keys feature is unavailable") + } else { + deactivate = h.authKeys.Deactivate + } + return deactivateByID(c, unavailableErr, "auth key", authkeys.ErrNotFound, "auth key not found: ", deactivate, authKeyWriteError) +} diff --git a/internal/admin/handler_budgets.go b/internal/admin/handler_budgets.go new file mode 100644 index 000000000..bd00728c7 --- /dev/null +++ b/internal/admin/handler_budgets.go @@ -0,0 +1,385 @@ +package admin + +import ( + "errors" + "fmt" + "net/http" + "net/url" + "strconv" + "strings" + "time" + + "github.com/labstack/echo/v5" + + "gomodel/internal/budget" + "gomodel/internal/core" +) + +func (h *Handler) ListBudgets(c *echo.Context) error { + if h.budgets == nil { + return handleError(c, featureUnavailableError("budgets feature is unavailable")) + } + now := time.Now().UTC() + statuses, err := h.budgets.Statuses(c.Request().Context(), now) + if err != nil { + return handleError(c, budgetServiceError("failed to list budgets", err)) + } + return c.JSON(http.StatusOK, budgetListResponse{ + Budgets: budgetStatusResponses(statuses, now), + ServerTime: now, + }) +} + +// UpsertBudget handles PUT /admin/api/v1/budgets/{user_path}/{period}. +// @Summary Create or update one budget +// @Tags admin +// @Accept json +// @Produce json +// @Security BearerAuth +// @Param user_path path string true "URL-encoded budget user path" +// @Param period path string true "Budget period name or seconds" +// @Param budget body upsertBudgetRequest true "Budget amount" +// @Success 200 {object} budgetListResponse +// @Failure 400 {object} core.GatewayError +// @Failure 401 {object} core.GatewayError +// @Failure 503 {object} core.GatewayError +// @Router /admin/api/v1/budgets/{user_path}/{period} [put] +func (h *Handler) UpsertBudget(c *echo.Context) error { + if h.budgets == nil { + return handleError(c, featureUnavailableError("budgets feature is unavailable")) + } + var req upsertBudgetRequest + if err := c.Bind(&req); err != nil { + return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) + } + userPath, periodSeconds, err := budgetRouteKey(c) + if err != nil { + return handleError(c, core.NewInvalidRequestError(err.Error(), err)) + } + item, err := budget.NormalizeBudget(budget.Budget{ + UserPath: userPath, + PeriodSeconds: periodSeconds, + Amount: req.Amount, + Source: budget.SourceManual, + }) + if err != nil { + return handleError(c, core.NewInvalidRequestError(err.Error(), err)) + } + if err := h.budgets.UpsertBudgets(c.Request().Context(), []budget.Budget{item}); err != nil { + return handleError(c, budgetServiceError("failed to save budget", err)) + } + return h.ListBudgets(c) +} + +// DeleteBudget handles DELETE /admin/api/v1/budgets/{user_path}/{period}. +// @Summary Delete one budget +// @Tags admin +// @Produce json +// @Security BearerAuth +// @Param user_path path string true "URL-encoded budget user path" +// @Param period path string true "Budget period name or seconds" +// @Success 200 {object} budgetListResponse +// @Failure 400 {object} core.GatewayError +// @Failure 401 {object} core.GatewayError +// @Failure 503 {object} core.GatewayError +// @Router /admin/api/v1/budgets/{user_path}/{period} [delete] +func (h *Handler) DeleteBudget(c *echo.Context) error { + if h.budgets == nil { + return handleError(c, featureUnavailableError("budgets feature is unavailable")) + } + userPath, periodSeconds, err := budgetRouteKey(c) + if err != nil { + return handleError(c, core.NewInvalidRequestError(err.Error(), err)) + } + if err := h.budgets.DeleteBudget(c.Request().Context(), userPath, periodSeconds); err != nil { + return handleError(c, budgetServiceError("failed to delete budget", err)) + } + return h.ListBudgets(c) +} + +// BudgetSettings handles GET /admin/api/v1/budgets/settings. +// @Summary Get budget reset settings +// @Tags admin +// @Produce json +// @Security BearerAuth +// @Success 200 {object} budget.Settings +// @Failure 401 {object} core.GatewayError +// @Failure 503 {object} core.GatewayError +// @Router /admin/api/v1/budgets/settings [get] +func (h *Handler) BudgetSettings(c *echo.Context) error { + if h.budgets == nil { + return handleError(c, featureUnavailableError("budgets feature is unavailable")) + } + return c.JSON(http.StatusOK, h.budgets.Settings()) +} + +// UpdateBudgetSettings handles PUT /admin/api/v1/budgets/settings. +// @Summary Update budget reset settings +// @Tags admin +// @Accept json +// @Produce json +// @Security BearerAuth +// @Param settings body updateBudgetSettingsRequest true "Budget reset settings" +// @Success 200 {object} budget.Settings +// @Failure 400 {object} core.GatewayError +// @Failure 401 {object} core.GatewayError +// @Failure 503 {object} core.GatewayError +// @Router /admin/api/v1/budgets/settings [put] +func (h *Handler) UpdateBudgetSettings(c *echo.Context) error { + if h.budgets == nil { + return handleError(c, featureUnavailableError("budgets feature is unavailable")) + } + var req updateBudgetSettingsRequest + if err := c.Bind(&req); err != nil { + return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) + } + settings := req.apply(h.budgets.Settings()) + if err := budget.ValidateSettings(settings); err != nil { + return handleError(c, core.NewInvalidRequestError(err.Error(), err)) + } + saved, err := h.budgets.SaveSettings(c.Request().Context(), settings) + if err != nil { + return handleError(c, budgetServiceError("failed to save budget settings", err)) + } + return c.JSON(http.StatusOK, saved) +} + +// ResetBudget handles POST /admin/api/v1/budgets/reset-one. +// @Summary Reset one budget period +// @Tags admin +// @Accept json +// @Produce json +// @Security BearerAuth +// @Param budget body resetBudgetRequest true "Budget key" +// @Success 200 {object} budgetListResponse +// @Failure 400 {object} core.GatewayError +// @Failure 401 {object} core.GatewayError +// @Failure 503 {object} core.GatewayError +// @Router /admin/api/v1/budgets/reset-one [post] +func (h *Handler) ResetBudget(c *echo.Context) error { + if h.budgets == nil { + return handleError(c, featureUnavailableError("budgets feature is unavailable")) + } + var req resetBudgetRequest + if err := c.Bind(&req); err != nil { + return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) + } + periodSeconds, err := budgetRequestPeriodSeconds(req.Period, req.PeriodSeconds) + if err != nil { + return handleError(c, core.NewInvalidRequestError(err.Error(), err)) + } + userPath, err := budget.NormalizeUserPath(req.UserPath) + if err != nil { + return handleError(c, core.NewInvalidRequestError(err.Error(), err)) + } + if err := h.budgets.ResetBudget(c.Request().Context(), userPath, periodSeconds, time.Now().UTC()); err != nil { + return handleError(c, budgetServiceError("failed to reset budget", err)) + } + return h.ListBudgets(c) +} + +// ResetBudgets handles POST /admin/api/v1/budgets/reset. +// @Summary Reset all budget periods +// @Tags admin +// @Accept json +// @Produce json +// @Security BearerAuth +// @Param confirmation body resetBudgetsRequest true "Reset confirmation" +// @Success 200 {object} resetBudgetsResponse +// @Failure 400 {object} core.GatewayError +// @Failure 401 {object} core.GatewayError +// @Failure 503 {object} core.GatewayError +// @Router /admin/api/v1/budgets/reset [post] +func (h *Handler) ResetBudgets(c *echo.Context) error { + if h.budgets == nil { + return handleError(c, featureUnavailableError("budgets feature is unavailable")) + } + var req resetBudgetsRequest + if err := c.Bind(&req); err != nil { + return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) + } + if strings.TrimSpace(strings.ToLower(req.confirmationValue())) != "reset" { + return handleError(c, core.NewInvalidRequestError("confirmation must be reset", nil)) + } + if err := h.budgets.ResetAll(c.Request().Context(), time.Now().UTC()); err != nil { + return handleError(c, budgetServiceError("failed to reset budgets", err)) + } + return c.JSON(http.StatusOK, resetBudgetsResponse{Status: "ok"}) +} + +// ProviderStatus handles GET /admin/api/v1/providers/status +type budgetListResponse struct { + Budgets []budgetStatusResponse `json:"budgets"` + ServerTime time.Time `json:"server_time"` +} + +type budgetStatusResponse struct { + UserPath string `json:"user_path"` + PeriodSeconds int64 `json:"period_seconds"` + PeriodLabel string `json:"period_label"` + Amount float64 `json:"amount"` + Source string `json:"source,omitempty"` + LastResetAt *time.Time `json:"last_reset_at,omitempty"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` + PeriodStart time.Time `json:"period_start"` + PeriodEnd time.Time `json:"period_end"` + Spent float64 `json:"spent"` + HasUsage bool `json:"has_usage"` + Remaining float64 `json:"remaining"` + UsageRatio float64 `json:"usage_ratio"` + PeriodRatio float64 `json:"period_ratio"` +} + +type upsertBudgetRequest struct { + Amount float64 `json:"amount"` +} + +type resetBudgetRequest struct { + UserPath string `json:"user_path"` + Period string `json:"period,omitempty"` + PeriodSeconds int64 `json:"period_seconds,omitempty"` +} + +type updateBudgetSettingsRequest struct { + DailyResetHour *int `json:"daily_reset_hour"` + DailyResetMinute *int `json:"daily_reset_minute"` + WeeklyResetWeekday *int `json:"weekly_reset_weekday"` + WeeklyResetHour *int `json:"weekly_reset_hour"` + WeeklyResetMinute *int `json:"weekly_reset_minute"` + MonthlyResetDay *int `json:"monthly_reset_day"` + MonthlyResetHour *int `json:"monthly_reset_hour"` + MonthlyResetMinute *int `json:"monthly_reset_minute"` +} + +func (r updateBudgetSettingsRequest) apply(settings budget.Settings) budget.Settings { + if r.DailyResetHour != nil { + settings.DailyResetHour = *r.DailyResetHour + } + if r.DailyResetMinute != nil { + settings.DailyResetMinute = *r.DailyResetMinute + } + if r.WeeklyResetWeekday != nil { + settings.WeeklyResetWeekday = *r.WeeklyResetWeekday + } + if r.WeeklyResetHour != nil { + settings.WeeklyResetHour = *r.WeeklyResetHour + } + if r.WeeklyResetMinute != nil { + settings.WeeklyResetMinute = *r.WeeklyResetMinute + } + if r.MonthlyResetDay != nil { + settings.MonthlyResetDay = *r.MonthlyResetDay + } + if r.MonthlyResetHour != nil { + settings.MonthlyResetHour = *r.MonthlyResetHour + } + if r.MonthlyResetMinute != nil { + settings.MonthlyResetMinute = *r.MonthlyResetMinute + } + return settings +} + +type resetBudgetsRequest struct { + Confirmation string `json:"confirmation"` + Confirm string `json:"confirm,omitempty"` +} + +func (r resetBudgetsRequest) confirmationValue() string { + if strings.TrimSpace(r.Confirmation) != "" { + return r.Confirmation + } + return r.Confirm +} + +type resetBudgetsResponse struct { + Status string `json:"status"` +} + +func budgetStatusResponses(statuses []budget.CheckResult, now time.Time) []budgetStatusResponse { + if len(statuses) == 0 { + return []budgetStatusResponse{} + } + responses := make([]budgetStatusResponse, 0, len(statuses)) + for _, status := range statuses { + item := status.Budget + usageRatio := 0.0 + if item.Amount > 0 { + usageRatio = status.Spent / item.Amount + } + periodRatio := 0.0 + periodDuration := status.PeriodEnd.Sub(status.PeriodStart).Seconds() + if periodDuration > 0 { + periodRatio = now.Sub(status.PeriodStart).Seconds() / periodDuration + } + responses = append(responses, budgetStatusResponse{ + UserPath: item.UserPath, + PeriodSeconds: item.PeriodSeconds, + PeriodLabel: budget.PeriodLabel(item.PeriodSeconds), + Amount: item.Amount, + Source: item.Source, + LastResetAt: item.LastResetAt, + CreatedAt: item.CreatedAt, + UpdatedAt: item.UpdatedAt, + PeriodStart: status.PeriodStart, + PeriodEnd: status.PeriodEnd, + Spent: status.Spent, + HasUsage: status.HasUsage, + Remaining: status.Remaining, + UsageRatio: usageRatio, + PeriodRatio: clampBudgetRatio(periodRatio), + }) + } + return responses +} + +func budgetRouteKey(c *echo.Context) (string, int64, error) { + userPathParam := strings.TrimSpace(c.Param("user_path")) + if userPathParam == "" { + return "", 0, errors.New("user_path path parameter is required") + } + userPath, err := url.PathUnescape(userPathParam) + if err != nil { + return "", 0, fmt.Errorf("invalid user_path path parameter: %w", err) + } + userPath, err = budget.NormalizeUserPath(userPath) + if err != nil { + return "", 0, err + } + + periodParam := strings.TrimSpace(c.Param("period")) + if periodParam == "" { + return "", 0, errors.New("period path parameter is required") + } + if seconds, err := strconv.ParseInt(periodParam, 10, 64); err == nil { + if seconds <= 0 { + return "", 0, errors.New("period_seconds must be greater than 0") + } + return userPath, seconds, nil + } + periodSeconds, err := budgetRequestPeriodSeconds(periodParam, 0) + if err != nil { + return "", 0, err + } + return userPath, periodSeconds, nil +} + +func budgetRequestPeriodSeconds(period string, periodSeconds int64) (int64, error) { + if periodSeconds > 0 { + return periodSeconds, nil + } + if parsed, ok := budget.PeriodSeconds(period); ok { + return parsed, nil + } + return 0, errors.New("period must be one of hourly, daily, weekly, monthly or period_seconds must be set") +} + +func clampBudgetRatio(value float64) float64 { + if value < 0 { + return 0 + } + if value > 1 { + return 1 + } + return value +} diff --git a/internal/admin/handler_guardrails.go b/internal/admin/handler_guardrails.go new file mode 100644 index 000000000..17feaa515 --- /dev/null +++ b/internal/admin/handler_guardrails.go @@ -0,0 +1,118 @@ +package admin + +import ( + "encoding/json" + "errors" + "net/http" + "strings" + + "github.com/labstack/echo/v5" + + "gomodel/internal/core" + "gomodel/internal/guardrails" +) + +type upsertGuardrailRequest struct { + Type string `json:"type"` + Description string `json:"description,omitempty"` + UserPath string `json:"user_path,omitempty"` + Config json.RawMessage `json:"config"` +} + +func (h *Handler) ListGuardrailTypes(c *echo.Context) error { + if h.guardrailDefs == nil { + return handleError(c, featureUnavailableError("guardrails feature is unavailable")) + } + return c.JSON(http.StatusOK, h.guardrailDefs.TypeDefinitions()) +} + +// ListGuardrails handles GET /admin/api/v1/guardrails +func (h *Handler) ListGuardrails(c *echo.Context) error { + if h.guardrailDefs == nil { + return handleError(c, featureUnavailableError("guardrails feature is unavailable")) + } + views := h.guardrailDefs.ListViews() + if views == nil { + views = []guardrails.View{} + } + return c.JSON(http.StatusOK, views) +} + +// UpsertGuardrail handles PUT /admin/api/v1/guardrails/{name} +func (h *Handler) UpsertGuardrail(c *echo.Context) error { + if h.guardrailDefs == nil { + return handleError(c, featureUnavailableError("guardrails feature is unavailable")) + } + + name := strings.TrimSpace(c.Param("name")) + if name == "" { + return handleError(c, core.NewInvalidRequestError("guardrail name is required", nil)) + } + + var req upsertGuardrailRequest + if err := c.Bind(&req); err != nil { + return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) + } + + userPath, err := normalizeUserPathQueryParam("user_path", req.UserPath) + if err != nil { + return handleError(c, err) + } + + h.mutationMu.Lock() + defer h.mutationMu.Unlock() + + if err := h.guardrailDefs.Upsert(c.Request().Context(), guardrails.Definition{ + Name: name, + Type: req.Type, + Description: req.Description, + UserPath: userPath, + Config: req.Config, + }); err != nil { + return handleError(c, guardrailWriteError(err)) + } + if err := h.refreshWorkflowsAfterGuardrailChange(c.Request().Context()); err != nil { + return handleError(c, err) + } + + definition, ok := h.guardrailDefs.Get(name) + if !ok { + return c.NoContent(http.StatusNoContent) + } + return c.JSON(http.StatusOK, guardrails.ViewFromDefinition(*definition)) +} + +// DeleteGuardrail handles DELETE /admin/api/v1/guardrails/{name} +func (h *Handler) DeleteGuardrail(c *echo.Context) error { + if h.guardrailDefs == nil { + return handleError(c, featureUnavailableError("guardrails feature is unavailable")) + } + + name := strings.TrimSpace(c.Param("name")) + if name == "" { + return handleError(c, core.NewInvalidRequestError("guardrail name is required", nil)) + } + + h.mutationMu.Lock() + defer h.mutationMu.Unlock() + + referencingWorkflows, err := h.activeWorkflowGuardrailReferences(c.Request().Context(), name) + if err != nil { + return handleError(c, err) + } + if len(referencingWorkflows) > 0 { + return handleError(c, core.NewInvalidRequestError("guardrail is used by active workflows: "+strings.Join(referencingWorkflows, ", "), nil)) + } + + if err := h.guardrailDefs.Delete(c.Request().Context(), name); err != nil { + if errors.Is(err, guardrails.ErrNotFound) { + return handleError(c, core.NewNotFoundError("guardrail not found: "+name)) + } + return handleError(c, guardrailWriteError(err)) + } + if err := h.refreshWorkflowsAfterGuardrailChange(c.Request().Context()); err != nil { + return handleError(c, err) + } + + return c.NoContent(http.StatusNoContent) +} diff --git a/internal/admin/handler_model_overrides.go b/internal/admin/handler_model_overrides.go new file mode 100644 index 000000000..5bdb1fb42 --- /dev/null +++ b/internal/admin/handler_model_overrides.go @@ -0,0 +1,79 @@ +package admin + +import ( + "context" + "log/slog" + "net/http" + + "github.com/labstack/echo/v5" + + "gomodel/internal/core" + "gomodel/internal/modeloverrides" +) + +type upsertModelOverrideRequest struct { + UserPaths []string `json:"user_paths,omitempty"` +} + +func (h *Handler) ListModelOverrides(c *echo.Context) error { + if h.modelOverrides == nil { + return handleError(c, featureUnavailableError("model overrides feature is unavailable")) + } + views := h.modelOverrides.ListViews() + if views == nil { + views = []modeloverrides.View{} + } + return c.JSON(http.StatusOK, views) +} + +// UpsertModelOverride handles PUT /admin/api/v1/model-overrides/{selector}. +func (h *Handler) UpsertModelOverride(c *echo.Context) error { + if h.modelOverrides == nil { + return handleError(c, featureUnavailableError("model overrides feature is unavailable")) + } + + selector, err := decodeModelOverridePathSelector(c.Param("selector")) + if err != nil { + return handleError(c, err) + } + + var req upsertModelOverrideRequest + if err := c.Bind(&req); err != nil { + return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) + } + + if err := h.modelOverrides.Upsert(c.Request().Context(), modeloverrides.Override{ + Selector: selector, + UserPaths: req.UserPaths, + }); err != nil { + return handleError(c, modelOverrideWriteError(err)) + } + + override, ok := h.modelOverrides.Get(selector) + if !ok || override == nil { + slog.Error("model override service returned no override after upsert", "selector", selector) + return handleError(c, core.NewProviderError("model_overrides", http.StatusInternalServerError, "model override update failed unexpectedly", nil)) + } + return c.JSON(http.StatusOK, override) +} + +// DeleteModelOverride handles DELETE /admin/api/v1/model-overrides/{selector}. +func (h *Handler) DeleteModelOverride(c *echo.Context) error { + var unavailableErr error + var deleteFunc func(context.Context, string) error + if h.modelOverrides == nil { + unavailableErr = featureUnavailableError("model overrides feature is unavailable") + } else { + deleteFunc = h.modelOverrides.Delete + } + return deleteByName( + c, + unavailableErr, + "selector", + decodeModelOverridePathSelector, + deleteFunc, + modeloverrides.ErrNotFound, + "model override not found: ", + modelOverrideWriteError, + ) +} diff --git a/internal/admin/handler_models.go b/internal/admin/handler_models.go new file mode 100644 index 000000000..1cbcea567 --- /dev/null +++ b/internal/admin/handler_models.go @@ -0,0 +1,129 @@ +package admin + +import ( + "net/http" + "slices" + "strings" + + "github.com/labstack/echo/v5" + + "gomodel/internal/core" + "gomodel/internal/modeloverrides" + "gomodel/internal/providers" +) + +type modelAccessResponse struct { + Selector string `json:"selector"` + DefaultEnabled bool `json:"default_enabled"` + EffectiveEnabled bool `json:"effective_enabled"` + UserPaths []string `json:"user_paths,omitempty"` + Override *modeloverrides.Override `json:"override,omitempty"` +} + +type modelInventoryResponse struct { + providers.ModelWithProvider + Access modelAccessResponse `json:"access"` +} + +func (h *Handler) ListModels(c *echo.Context) error { + if h.registry == nil { + return c.JSON(http.StatusOK, []modelInventoryResponse{}) + } + + cat := core.ModelCategory(c.QueryParam("category")) + if cat != "" && cat != core.CategoryAll { + if !isValidCategory(cat) { + return handleError(c, core.NewInvalidRequestError("invalid category: "+string(cat), nil)) + } + } + + var models []providers.ModelWithProvider + if cat != "" && cat != core.CategoryAll { + models = h.registry.ListModelsWithProviderByCategory(cat) + } else { + models = h.registry.ListModelsWithProvider() + } + + if models == nil { + models = []providers.ModelWithProvider{} + } + if h.modelOverrides == nil { + response := make([]modelInventoryResponse, 0, len(models)) + for _, model := range models { + selector := core.ModelSelector{ + Provider: strings.TrimSpace(model.ProviderName), + Model: strings.TrimSpace(model.Model.ID), + } + response = append(response, modelInventoryResponse{ + ModelWithProvider: model, + Access: modelAccessResponse{ + Selector: selector.QualifiedModel(), + DefaultEnabled: true, + EffectiveEnabled: true, + }, + }) + } + return c.JSON(http.StatusOK, response) + } + + response := make([]modelInventoryResponse, 0, len(models)) + for _, model := range models { + selector := core.ModelSelector{ + Provider: strings.TrimSpace(model.ProviderName), + Model: strings.TrimSpace(model.Model.ID), + } + effective := h.modelOverrides.EffectiveState(selector) + access := modelAccessResponse{ + Selector: effective.Selector, + DefaultEnabled: effective.DefaultEnabled, + EffectiveEnabled: effective.Enabled, + UserPaths: append([]string(nil), effective.UserPaths...), + } + if override, ok := h.modelOverrides.Get(selector.QualifiedModel()); ok && override != nil { + overrideCopy := *override + access.Override = &overrideCopy + } + response = append(response, modelInventoryResponse{ + ModelWithProvider: model, + Access: access, + }) + } + + return c.JSON(http.StatusOK, response) +} + +// isValidCategory returns true if cat is a recognized model category. +func isValidCategory(cat core.ModelCategory) bool { + return slices.Contains(core.AllCategories(), cat) +} + +// ListCategories handles GET /admin/api/v1/models/categories +// +// @Summary List model categories with counts +// @Tags admin +// @Produce json +// @Security BearerAuth +// @Success 200 {array} providers.CategoryCount +// @Failure 401 {object} core.GatewayError +// @Router /admin/api/v1/models/categories [get] +func (h *Handler) ListCategories(c *echo.Context) error { + if h.registry == nil { + return c.JSON(http.StatusOK, []providers.CategoryCount{}) + } + + return c.JSON(http.StatusOK, h.registry.GetCategoryCounts()) +} + +// DashboardConfig handles GET /admin/api/v1/dashboard/config +func (h *Handler) DashboardConfig(c *echo.Context) error { + return c.JSON(http.StatusOK, cloneDashboardRuntimeConfig(h.runtimeConfig)) +} + +// ListBudgets handles GET /admin/api/v1/budgets. +// @Summary List budgets with current status +// @Tags admin +// @Produce json +// @Security BearerAuth +// @Success 200 {object} budgetListResponse +// @Failure 401 {object} core.GatewayError +// @Failure 503 {object} core.GatewayError diff --git a/internal/admin/handler_providers.go b/internal/admin/handler_providers.go new file mode 100644 index 000000000..112f3e0a2 --- /dev/null +++ b/internal/admin/handler_providers.go @@ -0,0 +1,169 @@ +package admin + +import ( + "errors" + "net/http" + "sort" + "strings" + + "github.com/labstack/echo/v5" + + "gomodel/internal/core" + "gomodel/internal/providers" +) + +func (h *Handler) ProviderStatus(c *echo.Context) error { + return c.JSON(http.StatusOK, h.buildProviderStatusResponse()) +} + +// RefreshRuntime handles POST /admin/api/v1/runtime/refresh +func (h *Handler) RefreshRuntime(c *echo.Context) error { + if h.runtimeRefresher == nil { + return handleError(c, featureUnavailableError("runtime refresh is unavailable")) + } + + report, err := h.runtimeRefresher.RefreshRuntime(c.Request().Context()) + if err != nil { + if gatewayErr, ok := errors.AsType[*core.GatewayError](err); ok { + return handleError(c, gatewayErr) + } + return handleError(c, core.NewProviderError("runtime_refresh", http.StatusInternalServerError, "runtime refresh failed", err)) + } + if report.Status == "" { + report.Status = RuntimeRefreshStatusOK + } + if report.Steps == nil { + report.Steps = []RuntimeRefreshStep{} + } + return c.JSON(http.StatusOK, report) +} + +func (h *Handler) buildProviderStatusResponse() providerStatusResponse { + configured := cloneConfiguredProviders(h.configuredProviders) + configuredByName := make(map[string]providers.SanitizedProviderConfig, len(configured)) + nameSet := make(map[string]struct{}, len(configured)) + for _, cfg := range configured { + name := strings.TrimSpace(cfg.Name) + if name == "" { + continue + } + configuredByName[name] = cfg + nameSet[name] = struct{}{} + } + + runtimeByName := make(map[string]providers.ProviderRuntimeSnapshot) + if h.registry != nil { + for _, snapshot := range h.registry.ProviderRuntimeSnapshots() { + name := strings.TrimSpace(snapshot.Name) + if name == "" { + continue + } + runtimeByName[name] = snapshot + nameSet[name] = struct{}{} + } + } + + names := make([]string, 0, len(nameSet)) + for name := range nameSet { + names = append(names, name) + } + sort.Strings(names) + + resp := providerStatusResponse{ + Summary: providerStatusSummaryResponse{ + OverallStatus: "degraded", + }, + Providers: make([]providerStatusItemResponse, 0, len(names)), + } + + for _, name := range names { + cfg, hasConfig := configuredByName[name] + runtime, hasRuntime := runtimeByName[name] + if !hasConfig { + cfg = providers.SanitizedProviderConfig{Name: name, Type: strings.TrimSpace(runtime.Type)} + } + if !hasRuntime { + runtime = providers.ProviderRuntimeSnapshot{Name: name, Type: strings.TrimSpace(cfg.Type)} + } + if strings.TrimSpace(cfg.Type) == "" { + cfg.Type = strings.TrimSpace(runtime.Type) + } + if strings.TrimSpace(runtime.Type) == "" { + runtime.Type = strings.TrimSpace(cfg.Type) + } + + status, label, reason, lastError := classifyProviderStatus(cfg, runtime) + resp.Providers = append(resp.Providers, providerStatusItemResponse{ + Name: name, + Type: strings.TrimSpace(cfg.Type), + Status: status, + StatusLabel: label, + StatusReason: reason, + LastError: lastError, + Config: cfg, + Runtime: runtime, + }) + resp.Summary.Total++ + switch status { + case "healthy": + resp.Summary.Healthy++ + case "unhealthy": + resp.Summary.Unhealthy++ + default: + resp.Summary.Degraded++ + } + } + + switch { + case resp.Summary.Total == 0: + resp.Summary.OverallStatus = "degraded" + case resp.Summary.Healthy == resp.Summary.Total: + resp.Summary.OverallStatus = "healthy" + case resp.Summary.Unhealthy == resp.Summary.Total: + resp.Summary.OverallStatus = "unhealthy" + default: + resp.Summary.OverallStatus = "degraded" + } + + if resp.Providers == nil { + resp.Providers = []providerStatusItemResponse{} + } + return resp +} + +func classifyProviderStatus(cfg providers.SanitizedProviderConfig, runtime providers.ProviderRuntimeSnapshot) (status, label, reason, lastError string) { + modelFetchError := strings.TrimSpace(runtime.LastModelFetchError) + availabilityError := strings.TrimSpace(runtime.LastAvailabilityError) + configuredName := strings.TrimSpace(cfg.Name) + usingCachedModels := runtime.Registered && + runtime.DiscoveredModelCount > 0 && + modelFetchError == "" && + runtime.LastModelFetchSuccessAt == nil + + lastError = modelFetchError + if lastError == "" { + lastError = availabilityError + } + + switch { + case runtime.DiscoveredModelCount > 0 && modelFetchError == "": + if usingCachedModels { + return "degraded", "Starting", "serving cached model inventory while live refresh finishes", lastError + } + return "healthy", "Healthy", "configured and model discovery succeeded", lastError + case modelFetchError != "" && runtime.DiscoveredModelCount > 0: + return "degraded", "Degraded", "latest model refresh failed; previous inventory is still available", lastError + case modelFetchError != "": + return "unhealthy", "Unhealthy", "model discovery failed and no provider models are currently available", lastError + case availabilityError != "" && runtime.DiscoveredModelCount == 0: + return "unhealthy", "Unhealthy", "startup availability check failed and no provider models are available", lastError + case runtime.DiscoveredModelCount > 0: + return "healthy", "Healthy", "provider models are currently available", lastError + case !runtime.Registered && configuredName != "": + return "degraded", "Starting", "provider is configured and awaiting live model discovery", lastError + case configuredName != "": + return "degraded", "Configured", "provider is configured but has not exposed models yet", lastError + default: + return "degraded", "Unknown", "provider runtime inventory is unavailable", lastError + } +} diff --git a/internal/admin/handler_usage.go b/internal/admin/handler_usage.go new file mode 100644 index 000000000..95b98ac12 --- /dev/null +++ b/internal/admin/handler_usage.go @@ -0,0 +1,398 @@ +package admin + +import ( + "context" + "errors" + "net/http" + "strconv" + "strings" + "time" + + "github.com/labstack/echo/v5" + + "gomodel/internal/core" + "gomodel/internal/usage" +) + +// UsageSummary handles GET /admin/api/v1/usage/summary +// +// @Summary Get usage summary +// @Tags admin +// @Produce json +// @Security BearerAuth +// @Param days query int false "Number of days (default 30)" +// @Param start_date query string false "Start date (YYYY-MM-DD)" +// @Param end_date query string false "End date (YYYY-MM-DD)" +// @Param user_path query string false "Filter by tracked user path subtree" +// @Param cache_mode query string false "Cache mode filter: uncached, cached, all (default uncached)" +// @Success 200 {object} usage.UsageSummary +// @Failure 400 {object} core.GatewayError +// @Failure 401 {object} core.GatewayError +// @Router /admin/api/v1/usage/summary [get] +func (h *Handler) UsageSummary(c *echo.Context) error { + if h.usageReader == nil { + return c.JSON(http.StatusOK, usage.UsageSummary{}) + } + + params, err := parseUsageParams(c) + if err != nil { + return handleError(c, err) + } + + summary, err := h.usageReader.GetSummary(c.Request().Context(), params) + if err != nil { + return handleError(c, err) + } + + return c.JSON(http.StatusOK, summary) +} + +func usageSliceResponse[T any]( + c *echo.Context, + reader usage.UsageReader, + fetch func(context.Context, usage.UsageQueryParams) ([]T, error), +) error { + if reader == nil { + return c.JSON(http.StatusOK, []T{}) + } + + params, err := parseUsageParams(c) + if err != nil { + return handleError(c, err) + } + + values, err := fetch(c.Request().Context(), params) + if err != nil { + return handleError(c, err) + } + if values == nil { + values = []T{} + } + return c.JSON(http.StatusOK, values) +} + +// DailyUsage handles GET /admin/api/v1/usage/daily +// +// @Summary Get usage breakdown by period +// @Tags admin +// @Produce json +// @Security BearerAuth +// @Param days query int false "Number of days (default 30)" +// @Param start_date query string false "Start date (YYYY-MM-DD)" +// @Param end_date query string false "End date (YYYY-MM-DD)" +// @Param interval query string false "Grouping interval: daily, weekly, monthly, yearly (default daily)" +// @Param user_path query string false "Filter by tracked user path subtree" +// @Param cache_mode query string false "Cache mode filter: uncached, cached, all (default uncached)" +// @Success 200 {array} usage.DailyUsage +// @Failure 400 {object} core.GatewayError +// @Failure 401 {object} core.GatewayError +// @Router /admin/api/v1/usage/daily [get] +func (h *Handler) DailyUsage(c *echo.Context) error { + return usageSliceResponse(c, h.usageReader, func(ctx context.Context, params usage.UsageQueryParams) ([]usage.DailyUsage, error) { + return h.usageReader.GetDailyUsage(ctx, params) + }) +} + +// UsageByModel handles GET /admin/api/v1/usage/models +// +// @Summary Get usage breakdown by model +// @Tags admin +// @Produce json +// @Security BearerAuth +// @Param days query int false "Number of days (default 30)" +// @Param start_date query string false "Start date (YYYY-MM-DD)" +// @Param end_date query string false "End date (YYYY-MM-DD)" +// @Param user_path query string false "Filter by tracked user path subtree" +// @Param cache_mode query string false "Cache mode filter: uncached, cached, all (default uncached)" +// @Success 200 {array} usage.ModelUsage +// @Failure 400 {object} core.GatewayError +// @Failure 401 {object} core.GatewayError +// @Router /admin/api/v1/usage/models [get] +func (h *Handler) UsageByModel(c *echo.Context) error { + return usageSliceResponse(c, h.usageReader, func(ctx context.Context, params usage.UsageQueryParams) ([]usage.ModelUsage, error) { + return h.usageReader.GetUsageByModel(ctx, params) + }) +} + +// UsageByUserPath handles GET /admin/api/v1/usage/user-paths +// +// @Summary Get usage breakdown by user path +// @Tags admin +// @Produce json +// @Security BearerAuth +// @Param days query int false "Number of days (default 30)" +// @Param start_date query string false "Start date (YYYY-MM-DD)" +// @Param end_date query string false "End date (YYYY-MM-DD)" +// @Param user_path query string false "Filter by tracked user path subtree" +// @Param cache_mode query string false "Cache mode filter: uncached, cached, all (default uncached)" +// @Success 200 {array} usage.UserPathUsage +// @Failure 400 {object} core.GatewayError +// @Failure 401 {object} core.GatewayError +// @Router /admin/api/v1/usage/user-paths [get] +func (h *Handler) UsageByUserPath(c *echo.Context) error { + return usageSliceResponse(c, h.usageReader, func(ctx context.Context, params usage.UsageQueryParams) ([]usage.UserPathUsage, error) { + return h.usageReader.GetUsageByUserPath(ctx, params) + }) +} + +// UsageLog handles GET /admin/api/v1/usage/log +// +// @Summary Get paginated usage log entries +// @Tags admin +// @Produce json +// @Security BearerAuth +// @Param days query int false "Number of days (default 30)" +// @Param start_date query string false "Start date (YYYY-MM-DD)" +// @Param end_date query string false "End date (YYYY-MM-DD)" +// @Param model query string false "Filter by model name" +// @Param provider query string false "Filter by provider name or provider type" +// @Param user_path query string false "Filter by tracked user path subtree" +// @Param cache_mode query string false "Cache mode filter: uncached, cached, all (default uncached)" +// @Param search query string false "Search across model, provider, request_id, provider_id" +// @Param limit query int false "Page size (default 50, max 200)" +// @Param offset query int false "Offset for pagination" +// @Success 200 {object} usage.UsageLogResult +// @Failure 400 {object} core.GatewayError +// @Failure 401 {object} core.GatewayError +// @Router /admin/api/v1/usage/log [get] +func (h *Handler) UsageLog(c *echo.Context) error { + if h.usageReader == nil { + return c.JSON(http.StatusOK, usage.UsageLogResult{ + Entries: []usage.UsageLogEntry{}, + }) + } + + baseParams, err := parseUsageParams(c) + if err != nil { + return handleError(c, err) + } + + params := usage.UsageLogParams{ + UsageQueryParams: baseParams, + Model: c.QueryParam("model"), + Provider: c.QueryParam("provider"), + Search: c.QueryParam("search"), + } + + if l := c.QueryParam("limit"); l != "" { + if parsed, err := strconv.Atoi(l); err == nil && parsed > 0 { + params.Limit = parsed + } + } + if o := c.QueryParam("offset"); o != "" { + if parsed, err := strconv.Atoi(o); err == nil && parsed >= 0 { + params.Offset = parsed + } + } + + result, err := h.usageReader.GetUsageLog(c.Request().Context(), params) + if err != nil { + return handleError(c, err) + } + + if result.Entries == nil { + result.Entries = []usage.UsageLogEntry{} + } + + return c.JSON(http.StatusOK, result) +} + +// RecalculateUsagePricing handles POST /admin/api/v1/usage/recalculate-pricing. +// +// @Summary Recalculate stored usage costs from current model pricing metadata +// @Tags admin +// @Accept json +// @Produce json +// @Security BearerAuth +// @Param request body recalculatePricingRequest true "Recalculation filters and confirmation" +// @Success 200 {object} usage.RecalculatePricingResult +// @Failure 400 {object} core.GatewayError +// @Failure 401 {object} core.GatewayError +// @Failure 500 {object} core.GatewayError +// @Failure 503 {object} core.GatewayError +// @Router /admin/api/v1/usage/recalculate-pricing [post] +func (h *Handler) RecalculateUsagePricing(c *echo.Context) error { + if h.usageRecalculator == nil { + return handleError(c, featureUnavailableError("usage pricing recalculation is unavailable")) + } + if h.registry == nil { + return handleError(c, featureUnavailableError("model pricing metadata is unavailable")) + } + + var req recalculatePricingRequest + if err := c.Bind(&req); err != nil { + return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) + } + if strings.TrimSpace(strings.ToLower(req.confirmationValue())) != "recalculate" { + return handleError(c, core.NewInvalidRequestError("confirmation must be recalculate", nil)) + } + + params, err := h.recalculatePricingParams(c, req) + if err != nil { + return handleError(c, err) + } + + h.pricingMu.Lock() + defer h.pricingMu.Unlock() + + result, err := h.usageRecalculator.RecalculatePricing(c.Request().Context(), params, h.registry) + if err != nil { + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return handleError(c, err) + } + if gatewayErr, ok := errors.AsType[*core.GatewayError](err); ok { + return handleError(c, gatewayErr) + } + return handleError(c, core.NewProviderError("usage", http.StatusInternalServerError, "failed to recalculate usage pricing", err)) + } + return c.JSON(http.StatusOK, result) +} + +// CacheOverview handles GET /admin/api/v1/cache/overview +// +// @Summary Get cached-only usage overview +// @Tags admin +// @Produce json +// @Security BearerAuth +// @Param days query int false "Number of days (default 30)" +// @Param start_date query string false "Start date (YYYY-MM-DD)" +// @Param end_date query string false "End date (YYYY-MM-DD)" +// @Param interval query string false "Grouping interval: daily, weekly, monthly, yearly (default daily)" +// @Param user_path query string false "Filter by tracked user path subtree" +// @Param cache_mode query string false "Cache mode filter: uncached, cached, all (cache overview always uses cached mode)" +// @Success 200 {object} usage.CacheOverview +// @Failure 400 {object} core.GatewayError +// @Failure 401 {object} core.GatewayError +// @Failure 503 {object} core.GatewayError +// @Router /admin/api/v1/cache/overview [get] +func (h *Handler) CacheOverview(c *echo.Context) error { + if strings.TrimSpace(h.runtimeConfig.CacheEnabled) != "on" { + return handleError(c, featureUnavailableError("cache analytics is unavailable")) + } + if h.usageReader == nil { + return c.JSON(http.StatusOK, usage.CacheOverview{ + Daily: []usage.CacheOverviewDaily{}, + }) + } + + params, err := parseUsageParams(c) + if err != nil { + return handleError(c, err) + } + params.CacheMode = usage.CacheModeCached + + overview, err := h.usageReader.GetCacheOverview(c.Request().Context(), params) + if err != nil { + return handleError(c, err) + } + if overview == nil { + overview = &usage.CacheOverview{} + } + if overview.Daily == nil { + overview.Daily = []usage.CacheOverviewDaily{} + } + + return c.JSON(http.StatusOK, overview) +} + +// AuditLog handles GET /admin/api/v1/audit/log +// +// @Summary Get paginated audit log entries +// @Tags admin +// @Produce json +// @Security BearerAuth +// @Param days query int false "Number of days (default 30)" +// @Param start_date query string false "Start date (YYYY-MM-DD)" +// @Param end_date query string false "End date (YYYY-MM-DD)" +// @Param requested_model query string false "Filter by requested model selector" +// @Param provider query string false "Filter by provider name or provider type" +// @Param method query string false "Filter by HTTP method" +// @Param path query string false "Filter by request path" +// @Param user_path query string false "Filter by tracked user path subtree" +// @Param error_type query string false "Filter by error type" +// @Param status_code query int false "Filter by status code" +type recalculatePricingRequest struct { + Days int `json:"days,omitempty"` + StartDate string `json:"start_date,omitempty"` + EndDate string `json:"end_date,omitempty"` + UserPath string `json:"user_path,omitempty"` + Selector string `json:"selector,omitempty"` + Confirmation string `json:"confirmation"` + Confirm string `json:"confirm,omitempty"` +} + +func (r recalculatePricingRequest) confirmationValue() string { + if strings.TrimSpace(r.Confirmation) != "" { + return r.Confirmation + } + return r.Confirm +} + +func (h *Handler) recalculatePricingParams(c *echo.Context, req recalculatePricingRequest) (usage.RecalculatePricingParams, error) { + baseParams, err := recalculatePricingDateParams(c, req) + if err != nil { + return usage.RecalculatePricingParams{}, err + } + + userPath, err := normalizeUserPathQueryParam("user_path", req.UserPath) + if err != nil { + return usage.RecalculatePricingParams{}, err + } + baseParams.UserPath = userPath + baseParams.CacheMode = usage.CacheModeAll + + provider, model, err := h.recalculatePricingSelector(req.Selector) + if err != nil { + return usage.RecalculatePricingParams{}, err + } + + return usage.RecalculatePricingParams{ + UsageQueryParams: baseParams, + Provider: provider, + Model: model, + }, nil +} + +func recalculatePricingDateParams(c *echo.Context, req recalculatePricingRequest) (usage.UsageQueryParams, error) { + var params usage.UsageQueryParams + + timeZone, location := dashboardTimeZone(c) + params.TimeZone = timeZone + + now := timeNow().In(location) + today := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, location) + + start, end, err := buildDateRange(strings.TrimSpace(req.StartDate), strings.TrimSpace(req.EndDate), normalizeDateRangeDays(req.Days), location, today) + if err != nil { + return params, err + } + params.StartDate = start + params.EndDate = end + return params, nil +} + +func (h *Handler) recalculatePricingSelector(raw string) (provider, model string, err error) { + raw = strings.TrimSpace(raw) + if raw == "" { + return "", "", nil + } + + if h.aliases != nil { + selector, changed, err := h.aliases.ResolveModel(core.NewRequestedModelSelector(raw, "")) + if err != nil { + return "", "", core.NewInvalidRequestError("invalid selector: "+err.Error(), err) + } + if changed { + return selector.Provider, selector.Model, nil + } + } + + selector, err := core.ParseModelSelector(raw, "") + if err != nil { + return "", "", core.NewInvalidRequestError("invalid selector: "+err.Error(), err) + } + if selector.Provider == "" { + return "", "", core.NewInvalidRequestError("invalid selector: provider/model or alias is required", nil) + } + return selector.Provider, selector.Model, nil +} diff --git a/internal/admin/handler_workflows.go b/internal/admin/handler_workflows.go new file mode 100644 index 000000000..187b1a2e9 --- /dev/null +++ b/internal/admin/handler_workflows.go @@ -0,0 +1,246 @@ +package admin + +import ( + "context" + "errors" + "net/http" + "slices" + "sort" + "strings" + + "github.com/labstack/echo/v5" + + "gomodel/internal/core" + "gomodel/internal/workflows" +) + +type createWorkflowRequest struct { + ScopeProviderName string `json:"scope_provider_name,omitempty"` + LegacyScopeProvider string `json:"scope_provider,omitempty"` + ScopeModel string `json:"scope_model,omitempty"` + ScopeUserPath string `json:"scope_user_path,omitempty"` + Name string `json:"name"` + Description string `json:"description,omitempty"` + Payload workflows.Payload `json:"workflow_payload"` +} + +func (h *Handler) ListWorkflows(c *echo.Context) error { + if h.workflows == nil { + return handleError(c, featureUnavailableError("workflows feature is unavailable")) + } + + views, err := h.workflows.ListViews(c.Request().Context()) + if err != nil { + return handleError(c, err) + } + if views == nil { + views = []workflows.View{} + } + return c.JSON(http.StatusOK, views) +} + +// GetWorkflow handles GET /admin/api/v1/workflows/:id +func (h *Handler) GetWorkflow(c *echo.Context) error { + if h.workflows == nil { + return handleError(c, featureUnavailableError("workflows feature is unavailable")) + } + + id := strings.TrimSpace(c.Param("id")) + if id == "" { + return handleError(c, core.NewInvalidRequestError("workflow id is required", nil)) + } + + view, err := h.workflows.GetView(c.Request().Context(), id) + if err != nil { + if errors.Is(err, workflows.ErrNotFound) { + return handleError(c, core.NewNotFoundError("workflow not found: "+id)) + } + return handleError(c, err) + } + + return c.JSON(http.StatusOK, view) +} + +// ListWorkflowGuardrails handles GET /admin/api/v1/workflows/guardrails +func (h *Handler) ListWorkflowGuardrails(c *echo.Context) error { + if h.guardrails == nil { + return c.JSON(http.StatusOK, []string{}) + } + + return c.JSON(http.StatusOK, h.guardrails.Names()) +} + +// CreateWorkflow handles POST /admin/api/v1/workflows +func (h *Handler) CreateWorkflow(c *echo.Context) error { + if h.workflows == nil { + return handleError(c, featureUnavailableError("workflows feature is unavailable")) + } + + var req createWorkflowRequest + if err := c.Bind(&req); err != nil { + return handleError(c, core.NewInvalidRequestError("invalid request body: "+err.Error(), err)) + } + scopeProviderName := strings.TrimSpace(req.ScopeProviderName) + if scopeProviderName == "" { + scopeProviderName = strings.TrimSpace(req.LegacyScopeProvider) + } + scopeModel := strings.TrimSpace(req.ScopeModel) + + scopeUserPath, err := normalizeUserPathQueryParam("scope_user_path", req.ScopeUserPath) + if err != nil { + return handleError(c, err) + } + + scopeProviderName, err = h.validateWorkflowScope(scopeProviderName, scopeModel) + if err != nil { + return handleError(c, err) + } + + if err := h.validateWorkflowGuardrails(req.Payload); err != nil { + return handleError(c, err) + } + + h.mutationMu.Lock() + defer h.mutationMu.Unlock() + + version, err := h.workflows.Create(c.Request().Context(), workflows.CreateInput{ + Scope: workflows.Scope{ + Provider: scopeProviderName, + Model: scopeModel, + UserPath: scopeUserPath, + }, + Activate: true, + Name: req.Name, + Description: req.Description, + Payload: req.Payload, + }) + if err != nil { + return handleError(c, workflowWriteError(err)) + } + if version == nil { + return c.NoContent(http.StatusNoContent) + } + return c.JSON(http.StatusCreated, version) +} + +// DeactivateWorkflow handles POST /admin/api/v1/workflows/:id/deactivate +func (h *Handler) DeactivateWorkflow(c *echo.Context) error { + if h.workflows == nil { + return handleError(c, featureUnavailableError("workflows feature is unavailable")) + } + + id := strings.TrimSpace(c.Param("id")) + if id == "" { + return handleError(c, core.NewInvalidRequestError("workflow id is required", nil)) + } + + h.mutationMu.Lock() + defer h.mutationMu.Unlock() + + if err := h.workflows.Deactivate(c.Request().Context(), id); err != nil { + if errors.Is(err, workflows.ErrNotFound) { + return handleError(c, core.NewNotFoundError("workflow not found: "+id)) + } + return handleError(c, workflowWriteError(err)) + } + return c.NoContent(http.StatusNoContent) +} + +func (h *Handler) refreshWorkflowsAfterGuardrailChange(ctx context.Context) error { + if h.workflows == nil { + return nil + } + if err := h.workflows.Refresh(ctx); err != nil { + return err + } + return nil +} + +func (h *Handler) activeWorkflowGuardrailReferences(ctx context.Context, name string) ([]string, error) { + if h.workflows == nil { + return nil, nil + } + + name = strings.TrimSpace(name) + if name == "" { + return nil, nil + } + + views, err := h.workflows.ListViews(ctx) + if err != nil { + return nil, err + } + + references := make([]string, 0) + for _, view := range views { + if !view.Payload.Features.Guardrails { + continue + } + for _, step := range view.Payload.Guardrails { + if strings.TrimSpace(step.Ref) != name { + continue + } + references = append(references, view.ScopeDisplay) + break + } + } + sort.Strings(references) + return references, nil +} + +func (h *Handler) validateWorkflowGuardrails(payload workflows.Payload) error { + if !payload.Features.Guardrails || len(payload.Guardrails) == 0 { + return nil + } + if h.guardrails == nil { + return featureUnavailableError("guardrail registry is unavailable for workflow authoring") + } + + known := make(map[string]struct{}, h.guardrails.Len()) + for _, name := range h.guardrails.Names() { + known[name] = struct{}{} + } + for _, step := range payload.Guardrails { + ref := strings.TrimSpace(step.Ref) + if ref == "" { + continue + } + if _, ok := known[ref]; !ok { + return core.NewInvalidRequestError("unknown guardrail ref: "+ref, nil) + } + } + return nil +} + +func (h *Handler) validateWorkflowScope(scopeProviderName, scopeModel string) (string, error) { + scopeProviderName = strings.TrimSpace(scopeProviderName) + scopeModel = strings.TrimSpace(scopeModel) + + if scopeProviderName == "" { + if scopeModel != "" { + return "", core.NewInvalidRequestError("scope_model requires scope_provider_name", nil) + } + return "", nil + } + if h.registry == nil { + return "", core.NewInvalidRequestError("provider registry is unavailable for workflow provider-name validation", nil) + } + if !slices.Contains(h.registry.ProviderNames(), scopeProviderName) { + if resolvedProviderName := strings.TrimSpace(h.registry.GetProviderNameForType(scopeProviderName)); resolvedProviderName != "" { + scopeProviderName = resolvedProviderName + } + } + if !slices.Contains(h.registry.ProviderNames(), scopeProviderName) { + return "", core.NewInvalidRequestError("unknown provider name: "+scopeProviderName, nil) + } + if scopeModel == "" { + return scopeProviderName, nil + } + + for _, model := range h.registry.ListModelsWithProvider() { + if model.ProviderName == scopeProviderName && model.Model.ID == scopeModel { + return scopeProviderName, nil + } + } + return "", core.NewInvalidRequestError("unknown model for provider name "+scopeProviderName+": "+scopeModel, nil) +} diff --git a/internal/admin/routes.go b/internal/admin/routes.go new file mode 100644 index 000000000..7f618cd50 --- /dev/null +++ b/internal/admin/routes.go @@ -0,0 +1,68 @@ +package admin + +import "github.com/labstack/echo/v5" + +// RouteRegistrar is the subset of *echo.Group / *echo.Echo that RegisterRoutes +// uses. Decoupling from a concrete echo type keeps the admin package useful for +// callers that want to mount the API under a different path prefix or wrap the +// routes with extra middleware. +type RouteRegistrar interface { + GET(path string, h echo.HandlerFunc, m ...echo.MiddlewareFunc) echo.RouteInfo + POST(path string, h echo.HandlerFunc, m ...echo.MiddlewareFunc) echo.RouteInfo + PUT(path string, h echo.HandlerFunc, m ...echo.MiddlewareFunc) echo.RouteInfo + DELETE(path string, h echo.HandlerFunc, m ...echo.MiddlewareFunc) echo.RouteInfo +} + +// RegisterRoutes mounts the admin REST API on the given route group. +// Callers typically pass an *echo.Group rooted at /admin/api/v1. +func (h *Handler) RegisterRoutes(g RouteRegistrar) { + g.GET("/dashboard/config", h.DashboardConfig) + g.GET("/cache/overview", h.CacheOverview) + + g.GET("/usage/summary", h.UsageSummary) + g.GET("/usage/daily", h.DailyUsage) + g.GET("/usage/models", h.UsageByModel) + g.GET("/usage/user-paths", h.UsageByUserPath) + g.GET("/usage/log", h.UsageLog) + g.POST("/usage/recalculate-pricing", h.RecalculateUsagePricing) + + g.GET("/audit/log", h.AuditLog) + g.GET("/audit/conversation", h.AuditConversation) + + g.GET("/providers/status", h.ProviderStatus) + g.POST("/runtime/refresh", h.RefreshRuntime) + + g.GET("/budgets", h.ListBudgets) + g.PUT("/budgets/:user_path/:period", h.UpsertBudget) + g.DELETE("/budgets/:user_path/:period", h.DeleteBudget) + g.GET("/budgets/settings", h.BudgetSettings) + g.PUT("/budgets/settings", h.UpdateBudgetSettings) + g.POST("/budgets/reset-one", h.ResetBudget) + g.POST("/budgets/reset", h.ResetBudgets) + + g.GET("/models", h.ListModels) + g.GET("/models/categories", h.ListCategories) + + g.GET("/model-overrides", h.ListModelOverrides) + g.PUT("/model-overrides/:selector", h.UpsertModelOverride) + g.DELETE("/model-overrides/:selector", h.DeleteModelOverride) + + g.GET("/auth-keys", h.ListAuthKeys) + g.POST("/auth-keys", h.CreateAuthKey) + g.POST("/auth-keys/:id/deactivate", h.DeactivateAuthKey) + + g.GET("/aliases", h.ListAliases) + g.PUT("/aliases/:name", h.UpsertAlias) + g.DELETE("/aliases/:name", h.DeleteAlias) + + g.GET("/guardrails/types", h.ListGuardrailTypes) + g.GET("/guardrails", h.ListGuardrails) + g.PUT("/guardrails/:name", h.UpsertGuardrail) + g.DELETE("/guardrails/:name", h.DeleteGuardrail) + + g.GET("/workflows", h.ListWorkflows) + g.GET("/workflows/guardrails", h.ListWorkflowGuardrails) + g.GET("/workflows/:id", h.GetWorkflow) + g.POST("/workflows", h.CreateWorkflow) + g.POST("/workflows/:id/deactivate", h.DeactivateWorkflow) +} diff --git a/internal/server/http.go b/internal/server/http.go index 1e7161164..84cd74783 100644 --- a/internal/server/http.go +++ b/internal/server/http.go @@ -315,46 +315,7 @@ func New(provider core.RoutableProvider, cfg *Config) *Server { // Admin API routes (behind ADMIN_ENDPOINTS_ENABLED flag) if cfg != nil && cfg.AdminEndpointsEnabled && cfg.AdminHandler != nil { - adminAPI := e.Group("/admin/api/v1") - adminAPI.GET("/dashboard/config", cfg.AdminHandler.DashboardConfig) - adminAPI.GET("/cache/overview", cfg.AdminHandler.CacheOverview) - adminAPI.GET("/usage/summary", cfg.AdminHandler.UsageSummary) - adminAPI.GET("/usage/daily", cfg.AdminHandler.DailyUsage) - adminAPI.GET("/usage/models", cfg.AdminHandler.UsageByModel) - adminAPI.GET("/usage/user-paths", cfg.AdminHandler.UsageByUserPath) - adminAPI.GET("/usage/log", cfg.AdminHandler.UsageLog) - adminAPI.POST("/usage/recalculate-pricing", cfg.AdminHandler.RecalculateUsagePricing) - adminAPI.GET("/audit/log", cfg.AdminHandler.AuditLog) - adminAPI.GET("/audit/conversation", cfg.AdminHandler.AuditConversation) - adminAPI.GET("/providers/status", cfg.AdminHandler.ProviderStatus) - adminAPI.POST("/runtime/refresh", cfg.AdminHandler.RefreshRuntime) - adminAPI.GET("/budgets", cfg.AdminHandler.ListBudgets) - adminAPI.PUT("/budgets/:user_path/:period", cfg.AdminHandler.UpsertBudget) - adminAPI.DELETE("/budgets/:user_path/:period", cfg.AdminHandler.DeleteBudget) - adminAPI.GET("/budgets/settings", cfg.AdminHandler.BudgetSettings) - adminAPI.PUT("/budgets/settings", cfg.AdminHandler.UpdateBudgetSettings) - adminAPI.POST("/budgets/reset-one", cfg.AdminHandler.ResetBudget) - adminAPI.POST("/budgets/reset", cfg.AdminHandler.ResetBudgets) - adminAPI.GET("/models", cfg.AdminHandler.ListModels) - adminAPI.GET("/models/categories", cfg.AdminHandler.ListCategories) - adminAPI.GET("/model-overrides", cfg.AdminHandler.ListModelOverrides) - adminAPI.PUT("/model-overrides/:selector", cfg.AdminHandler.UpsertModelOverride) - adminAPI.DELETE("/model-overrides/:selector", cfg.AdminHandler.DeleteModelOverride) - adminAPI.GET("/auth-keys", cfg.AdminHandler.ListAuthKeys) - adminAPI.POST("/auth-keys", cfg.AdminHandler.CreateAuthKey) - adminAPI.POST("/auth-keys/:id/deactivate", cfg.AdminHandler.DeactivateAuthKey) - adminAPI.GET("/aliases", cfg.AdminHandler.ListAliases) - adminAPI.PUT("/aliases/:name", cfg.AdminHandler.UpsertAlias) - adminAPI.DELETE("/aliases/:name", cfg.AdminHandler.DeleteAlias) - adminAPI.GET("/guardrails/types", cfg.AdminHandler.ListGuardrailTypes) - adminAPI.GET("/guardrails", cfg.AdminHandler.ListGuardrails) - adminAPI.PUT("/guardrails/:name", cfg.AdminHandler.UpsertGuardrail) - adminAPI.DELETE("/guardrails/:name", cfg.AdminHandler.DeleteGuardrail) - adminAPI.GET("/workflows", cfg.AdminHandler.ListWorkflows) - adminAPI.GET("/workflows/guardrails", cfg.AdminHandler.ListWorkflowGuardrails) - adminAPI.GET("/workflows/:id", cfg.AdminHandler.GetWorkflow) - adminAPI.POST("/workflows", cfg.AdminHandler.CreateWorkflow) - adminAPI.POST("/workflows/:id/deactivate", cfg.AdminHandler.DeactivateWorkflow) + cfg.AdminHandler.RegisterRoutes(e.Group("/admin/api/v1")) } // Admin dashboard UI routes (behind ADMIN_UI_ENABLED flag) From 9a165e6b8b2e26ff840ddc1f999e7c9fdf96120c Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Thu, 30 Apr 2026 08:58:13 +0200 Subject: [PATCH 08/15] fix(admin): reattach swagger doc blocks split across files The handler.go split mid-doc-comment for three handlers, leaving orphaned annotation fragments swaggo/swag silently ignores: - AuditLog Swagger header (\`@Summary\`/\`@Tags\`/\`@Param days...\`) ended up in handler_usage.go above recalculatePricingRequest while its tail (\`@Param search/limit/offset\`/\`@Success\`/\`@Router\`) sat above the AuditLog func in handler_audit.go. - ListModels header was at the tail of handler_audit.go with no function below it, while ListModels in handler_models.go had no comment. - ListBudgets block was orphaned at the tail of handler_models.go while ListBudgets in handler_budgets.go had no comment. - handler.go retained dangling \`@Router\`/\`@Param\` lines and one-line handler stubs ("// ListAliases handles ...") with no following code. - handler_budgets.go had a stray "ProviderStatus handles ..." comment above budgetListResponse. Reassemble each Swagger block above its handler and drop the orphans. Co-Authored-By: Claude Opus 4.7 --- internal/admin/handler.go | 10 ---------- internal/admin/handler_audit.go | 27 +++++++++++++++++---------- internal/admin/handler_budgets.go | 11 ++++++++++- internal/admin/handler_models.go | 19 ++++++++++--------- internal/admin/handler_usage.go | 16 ---------------- 5 files changed, 37 insertions(+), 46 deletions(-) diff --git a/internal/admin/handler.go b/internal/admin/handler.go index 40f73e651..138b97a38 100644 --- a/internal/admin/handler.go +++ b/internal/admin/handler.go @@ -501,13 +501,3 @@ func requestIDFromAdminContextOrHeader(req *http.Request) string { } return strings.TrimSpace(req.Header.Get("X-Request-ID")) } - -// @Param stream query bool false "Filter by stream mode (true/false)" -// @Router /admin/api/v1/models [get] -// @Router /admin/api/v1/budgets [get] - -// ListModelOverrides handles GET /admin/api/v1/model-overrides. -// ListAuthKeys handles GET /admin/api/v1/auth-keys -// ListAliases handles GET /admin/api/v1/aliases -// ListGuardrailTypes handles GET /admin/api/v1/guardrails/types -// ListWorkflows handles GET /admin/api/v1/workflows diff --git a/internal/admin/handler_audit.go b/internal/admin/handler_audit.go index b4633c7ea..4062fd0f4 100644 --- a/internal/admin/handler_audit.go +++ b/internal/admin/handler_audit.go @@ -14,6 +14,23 @@ import ( "gomodel/internal/usage" ) +// AuditLog handles GET /admin/api/v1/audit/log +// +// @Summary Get paginated audit log entries +// @Tags admin +// @Produce json +// @Security BearerAuth +// @Param days query int false "Number of days (default 30)" +// @Param start_date query string false "Start date (YYYY-MM-DD)" +// @Param end_date query string false "End date (YYYY-MM-DD)" +// @Param requested_model query string false "Filter by requested model selector" +// @Param provider query string false "Filter by provider name or provider type" +// @Param method query string false "Filter by HTTP method" +// @Param path query string false "Filter by request path" +// @Param user_path query string false "Filter by tracked user path subtree" +// @Param error_type query string false "Filter by error type" +// @Param status_code query int false "Filter by status code" +// @Param stream query bool false "Filter by stream mode (true/false)" // @Param search query string false "Search across request_id/requested_model/provider/method/path/error_type/error_message" // @Param limit query int false "Page size (default 25, max 100)" // @Param offset query int false "Offset for pagination" @@ -193,13 +210,3 @@ func (h *Handler) AuditConversation(c *echo.Context) error { return c.JSON(http.StatusOK, result) } - -// ListModels handles GET /admin/api/v1/models -// Supports optional ?category= query param for filtering by model category. -// -// @Summary List all registered models with provider info -// @Tags admin -// @Produce json -// @Security BearerAuth -// @Success 200 {array} providers.ModelWithProvider -// @Failure 401 {object} core.GatewayError diff --git a/internal/admin/handler_budgets.go b/internal/admin/handler_budgets.go index bd00728c7..444bab5d7 100644 --- a/internal/admin/handler_budgets.go +++ b/internal/admin/handler_budgets.go @@ -15,6 +15,16 @@ import ( "gomodel/internal/core" ) +// ListBudgets handles GET /admin/api/v1/budgets. +// +// @Summary List budgets with current status +// @Tags admin +// @Produce json +// @Security BearerAuth +// @Success 200 {object} budgetListResponse +// @Failure 401 {object} core.GatewayError +// @Failure 503 {object} core.GatewayError +// @Router /admin/api/v1/budgets [get] func (h *Handler) ListBudgets(c *echo.Context) error { if h.budgets == nil { return handleError(c, featureUnavailableError("budgets feature is unavailable")) @@ -207,7 +217,6 @@ func (h *Handler) ResetBudgets(c *echo.Context) error { return c.JSON(http.StatusOK, resetBudgetsResponse{Status: "ok"}) } -// ProviderStatus handles GET /admin/api/v1/providers/status type budgetListResponse struct { Budgets []budgetStatusResponse `json:"budgets"` ServerTime time.Time `json:"server_time"` diff --git a/internal/admin/handler_models.go b/internal/admin/handler_models.go index 1cbcea567..4b57924e5 100644 --- a/internal/admin/handler_models.go +++ b/internal/admin/handler_models.go @@ -25,6 +25,16 @@ type modelInventoryResponse struct { Access modelAccessResponse `json:"access"` } +// ListModels handles GET /admin/api/v1/models +// Supports optional ?category= query param for filtering by model category. +// +// @Summary List all registered models with provider info +// @Tags admin +// @Produce json +// @Security BearerAuth +// @Success 200 {array} providers.ModelWithProvider +// @Failure 401 {object} core.GatewayError +// @Router /admin/api/v1/models [get] func (h *Handler) ListModels(c *echo.Context) error { if h.registry == nil { return c.JSON(http.StatusOK, []modelInventoryResponse{}) @@ -118,12 +128,3 @@ func (h *Handler) ListCategories(c *echo.Context) error { func (h *Handler) DashboardConfig(c *echo.Context) error { return c.JSON(http.StatusOK, cloneDashboardRuntimeConfig(h.runtimeConfig)) } - -// ListBudgets handles GET /admin/api/v1/budgets. -// @Summary List budgets with current status -// @Tags admin -// @Produce json -// @Security BearerAuth -// @Success 200 {object} budgetListResponse -// @Failure 401 {object} core.GatewayError -// @Failure 503 {object} core.GatewayError diff --git a/internal/admin/handler_usage.go b/internal/admin/handler_usage.go index 95b98ac12..8369cf6ec 100644 --- a/internal/admin/handler_usage.go +++ b/internal/admin/handler_usage.go @@ -295,22 +295,6 @@ func (h *Handler) CacheOverview(c *echo.Context) error { return c.JSON(http.StatusOK, overview) } -// AuditLog handles GET /admin/api/v1/audit/log -// -// @Summary Get paginated audit log entries -// @Tags admin -// @Produce json -// @Security BearerAuth -// @Param days query int false "Number of days (default 30)" -// @Param start_date query string false "Start date (YYYY-MM-DD)" -// @Param end_date query string false "End date (YYYY-MM-DD)" -// @Param requested_model query string false "Filter by requested model selector" -// @Param provider query string false "Filter by provider name or provider type" -// @Param method query string false "Filter by HTTP method" -// @Param path query string false "Filter by request path" -// @Param user_path query string false "Filter by tracked user path subtree" -// @Param error_type query string false "Filter by error type" -// @Param status_code query int false "Filter by status code" type recalculatePricingRequest struct { Days int `json:"days,omitempty"` StartDate string `json:"start_date,omitempty"` From d74a6c6002c26f91fb5b6d70d96d5d80060dc532 Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Thu, 30 Apr 2026 08:59:56 +0200 Subject: [PATCH 09/15] docs(swagger): regenerate after reattaching admin handler annotations Re-runs \`make swagger\` after the previous commit reattached the orphaned ListModels Swagger block to its handler. The previously- orphaned ListModels comment lived in a file with no following function so swaggo silently dropped it, which left GET /admin/api/v1/models missing from the generated docs. Now included along with the providers.ModelWithProvider schema it references. Co-Authored-By: Claude Opus 4.7 --- cmd/gomodel/docs/docs.go | 50 ++++++++++++++++++++++++++++++++++++ docs/openapi.json | 55 ++++++++++++++++++++++++++++++++++++++++ 2 files changed, 105 insertions(+) diff --git a/cmd/gomodel/docs/docs.go b/cmd/gomodel/docs/docs.go index 09a12c130..597a763b7 100644 --- a/cmd/gomodel/docs/docs.go +++ b/cmd/gomodel/docs/docs.go @@ -633,6 +633,39 @@ const docTemplate = `{ ] } }, + "/admin/api/v1/models": { + "get": { + "produces": [ + "application/json" + ], + "tags": [ + "admin" + ], + "summary": "List all registered models with provider info", + "responses": { + "200": { + "description": "OK", + "schema": { + "type": "array", + "items": { + "$ref": "#/definitions/providers.ModelWithProvider" + } + } + }, + "401": { + "description": "Unauthorized", + "schema": { + "$ref": "#/definitions/core.GatewayError" + } + } + }, + "security": [ + { + "BearerAuth": [] + } + ] + } + }, "/admin/api/v1/models/categories": { "get": { "produces": [ @@ -4783,6 +4816,23 @@ const docTemplate = `{ } } }, + "providers.ModelWithProvider": { + "type": "object", + "properties": { + "model": { + "$ref": "#/definitions/core.Model" + }, + "provider_name": { + "type": "string" + }, + "provider_type": { + "type": "string" + }, + "selector": { + "type": "string" + } + } + }, "usage.CacheOverview": { "type": "object", "properties": { diff --git a/docs/openapi.json b/docs/openapi.json index 4e4074803..c9fb99834 100644 --- a/docs/openapi.json +++ b/docs/openapi.json @@ -780,6 +780,44 @@ ] } }, + "/admin/api/v1/models": { + "get": { + "tags": [ + "admin" + ], + "summary": "List all registered models with provider info", + "responses": { + "200": { + "description": "OK", + "content": { + "application/json": { + "schema": { + "type": "array", + "items": { + "$ref": "#/components/schemas/providers.ModelWithProvider" + } + } + } + } + }, + "401": { + "description": "Unauthorized", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.GatewayError" + } + } + } + } + }, + "security": [ + { + "BearerAuth": [] + } + ] + } + }, "/admin/api/v1/models/categories": { "get": { "tags": [ @@ -6283,6 +6321,23 @@ } } }, + "providers.ModelWithProvider": { + "type": "object", + "properties": { + "model": { + "$ref": "#/components/schemas/core.Model" + }, + "provider_name": { + "type": "string" + }, + "provider_type": { + "type": "string" + }, + "selector": { + "type": "string" + } + } + }, "usage.CacheOverview": { "type": "object", "properties": { From a63e16d15cfe91e7f4834d0b282f9f571f326f2e Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Thu, 30 Apr 2026 09:05:39 +0200 Subject: [PATCH 10/15] fix: reliability cleanups surfaced during tier 2 review MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Three small, scope-bounded reliability improvements pulled out of review comments on the tier 2 split PR. Each is independently testable and behavior-preserving in the success path. - audit log limit/offset (handler_audit.go): malformed or out-of-range values silently fell back to defaults while status_code/stream returned 400. Mirror the status_code/stream pattern so all paginated audit query params surface bad input the same way. limit=0 and negative values now error too — they were no-ops before. Adds TestAuditLog_InvalidLimit/InvalidOffset table tests. - guarded chat batch body rewrite (guardrails/batch_rewrite.go): rewriteGuardedChatBatchBody fell through to json.Marshal(modified) on any patch failure including modified == nil, which would serialize "null" as the batch item body. Add an explicit nil check matching the inline check in patchGuardedChatBatchBody so the fallback only runs for genuine raw-body preservation failures. - registry InitializeAsync (providers/registry_init.go): the background goroutine pinned itself to context.Background(), so caller cancellation (e.g., shutdown) couldn't propagate and the worker ran for up to 60s after the parent ctx died. Derive the timeout from the caller's ctx instead. Also re-evaluated the reviewer's "guardrails round-trip drops function_call_output type" finding and rejected it: ResponsesInputElement.Output is typed as string at this layer, so preserved.Output = modified.Content is type-correct. Type preservation for non-string raw-map outputs already happens correctly one level up via patchResponsesInputMap → restoreResponsesInputOutputValue. Co-Authored-By: Claude Opus 4.7 --- internal/admin/handler_audit.go | 12 +++++--- internal/admin/handler_test.go | 46 ++++++++++++++++++++++++++++ internal/guardrails/batch_rewrite.go | 5 +++ internal/providers/registry_init.go | 6 ++-- 4 files changed, 63 insertions(+), 6 deletions(-) diff --git a/internal/admin/handler_audit.go b/internal/admin/handler_audit.go index 4062fd0f4..1b2ed9fe9 100644 --- a/internal/admin/handler_audit.go +++ b/internal/admin/handler_audit.go @@ -90,14 +90,18 @@ func (h *Handler) AuditLog(c *echo.Context) error { } if l := c.QueryParam("limit"); l != "" { - if parsed, err := strconv.Atoi(l); err == nil && parsed > 0 { - params.Limit = parsed + parsed, err := strconv.Atoi(l) + if err != nil || parsed <= 0 { + return handleError(c, core.NewInvalidRequestError("invalid limit, expected positive integer", nil)) } + params.Limit = parsed } if o := c.QueryParam("offset"); o != "" { - if parsed, err := strconv.Atoi(o); err == nil && parsed >= 0 { - params.Offset = parsed + parsed, err := strconv.Atoi(o) + if err != nil || parsed < 0 { + return handleError(c, core.NewInvalidRequestError("invalid offset, expected non-negative integer", nil)) } + params.Offset = parsed } result, err := h.auditReader.GetLogs(c.Request().Context(), params) diff --git a/internal/admin/handler_test.go b/internal/admin/handler_test.go index 7654ade19..774316a07 100644 --- a/internal/admin/handler_test.go +++ b/internal/admin/handler_test.go @@ -1074,6 +1074,52 @@ func TestAuditLog_InvalidStream(t *testing.T) { } } +func TestAuditLog_InvalidLimit(t *testing.T) { + cases := []string{"abc", "0", "-1"} + for _, q := range cases { + t.Run(q, func(t *testing.T) { + reader := &mockAuditReader{ + logResult: &auditlog.LogListResult{Entries: []auditlog.LogEntry{}}, + } + h := NewHandler(nil, nil, WithAuditReader(reader)) + c, rec := newHandlerContext("/admin/api/v1/audit/log?limit=" + q) + + if err := h.AuditLog(c); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if rec.Code != http.StatusBadRequest { + t.Errorf("expected 400 for limit=%q, got %d", q, rec.Code) + } + if !containsString(rec.Body.String(), "invalid_request_error") { + t.Errorf("expected invalid_request_error in body for limit=%q, got: %s", q, rec.Body.String()) + } + }) + } +} + +func TestAuditLog_InvalidOffset(t *testing.T) { + cases := []string{"abc", "-1"} + for _, q := range cases { + t.Run(q, func(t *testing.T) { + reader := &mockAuditReader{ + logResult: &auditlog.LogListResult{Entries: []auditlog.LogEntry{}}, + } + h := NewHandler(nil, nil, WithAuditReader(reader)) + c, rec := newHandlerContext("/admin/api/v1/audit/log?offset=" + q) + + if err := h.AuditLog(c); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if rec.Code != http.StatusBadRequest { + t.Errorf("expected 400 for offset=%q, got %d", q, rec.Code) + } + if !containsString(rec.Body.String(), "invalid_request_error") { + t.Errorf("expected invalid_request_error in body for offset=%q, got: %s", q, rec.Body.String()) + } + }) + } +} + func TestAuditLog_Error(t *testing.T) { reader := &mockAuditReader{ logErr: core.NewProviderError("test", http.StatusBadGateway, "upstream failed", nil), diff --git a/internal/guardrails/batch_rewrite.go b/internal/guardrails/batch_rewrite.go index a18d792af..1703acc27 100644 --- a/internal/guardrails/batch_rewrite.go +++ b/internal/guardrails/batch_rewrite.go @@ -7,10 +7,15 @@ import ( ) func rewriteGuardedChatBatchBody(originalBody json.RawMessage, original *core.ChatRequest, modified *core.ChatRequest) (json.RawMessage, error) { + if modified == nil { + return nil, core.NewInvalidRequestError("missing guarded chat request", nil) + } body, err := patchGuardedChatBatchBody(originalBody, original, modified) if err == nil { return body, nil } + // Fallback: serialize the modified request directly when raw-body + // preservation fails (e.g., when the original body is malformed JSON). return json.Marshal(modified) } diff --git a/internal/providers/registry_init.go b/internal/providers/registry_init.go index 55a28788e..c309e889e 100644 --- a/internal/providers/registry_init.go +++ b/internal/providers/registry_init.go @@ -332,9 +332,11 @@ func (r *ModelRegistry) InitializeAsync(ctx context.Context) { slog.Info("serving traffic with cached models while refreshing", "cached_models", cached) } - // Start background initialization + // Start background initialization. Derive the timeout from the caller's + // ctx so shutdown cancellation propagates instead of leaving the goroutine + // running until the 60s timeout fires on its own. go func() { - initCtx, cancel := context.WithTimeout(context.Background(), 60*time.Second) + initCtx, cancel := context.WithTimeout(ctx, 60*time.Second) defer cancel() if err := r.Initialize(initCtx); err != nil { From a856b6e90264c76bf147fd9a41837ae5b82dec8e Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Thu, 30 Apr 2026 09:12:33 +0200 Subject: [PATCH 11/15] refactor: collapse duplicated model-inventory loop and lazy-init default output Two small cleanups: - handler_models.go: ListModels had near-identical loops for the modelOverrides==nil and !=nil branches, differing only in how the per-model access view was populated. Extract modelAccessResolver() that returns the appropriate per-selector callback once and run a single iteration. Removes ~25 lines and the slight risk of the two branches drifting apart on future field additions to modelAccessResponse. - responses_output.go: ConvertChatResponseToResponses unconditionally built a default empty ResponsesOutputItem and immediately replaced it whenever resp.Choices was non-empty (the common case). Switch to if/else so the placeholder is only allocated on the empty-choices path. No behavior change. Co-Authored-By: Claude Opus 4.7 --- internal/admin/handler_models.go | 51 +++++++++++++------------- internal/providers/responses_output.go | 32 ++++++++-------- 2 files changed, 43 insertions(+), 40 deletions(-) diff --git a/internal/admin/handler_models.go b/internal/admin/handler_models.go index 4b57924e5..3d3d2c119 100644 --- a/internal/admin/handler_models.go +++ b/internal/admin/handler_models.go @@ -57,31 +57,37 @@ func (h *Handler) ListModels(c *echo.Context) error { if models == nil { models = []providers.ModelWithProvider{} } - if h.modelOverrides == nil { - response := make([]modelInventoryResponse, 0, len(models)) - for _, model := range models { - selector := core.ModelSelector{ - Provider: strings.TrimSpace(model.ProviderName), - Model: strings.TrimSpace(model.Model.ID), - } - response = append(response, modelInventoryResponse{ - ModelWithProvider: model, - Access: modelAccessResponse{ - Selector: selector.QualifiedModel(), - DefaultEnabled: true, - EffectiveEnabled: true, - }, - }) - } - return c.JSON(http.StatusOK, response) - } - + access := h.modelAccessResolver() response := make([]modelInventoryResponse, 0, len(models)) for _, model := range models { selector := core.ModelSelector{ Provider: strings.TrimSpace(model.ProviderName), Model: strings.TrimSpace(model.Model.ID), } + response = append(response, modelInventoryResponse{ + ModelWithProvider: model, + Access: access(selector), + }) + } + + return c.JSON(http.StatusOK, response) +} + +// modelAccessResolver returns a function that produces the access view for a +// given selector. When model overrides are configured the resolver consults +// the service for effective state; otherwise every model is reported as +// default-on. +func (h *Handler) modelAccessResolver() func(core.ModelSelector) modelAccessResponse { + if h.modelOverrides == nil { + return func(selector core.ModelSelector) modelAccessResponse { + return modelAccessResponse{ + Selector: selector.QualifiedModel(), + DefaultEnabled: true, + EffectiveEnabled: true, + } + } + } + return func(selector core.ModelSelector) modelAccessResponse { effective := h.modelOverrides.EffectiveState(selector) access := modelAccessResponse{ Selector: effective.Selector, @@ -93,13 +99,8 @@ func (h *Handler) ListModels(c *echo.Context) error { overrideCopy := *override access.Override = &overrideCopy } - response = append(response, modelInventoryResponse{ - ModelWithProvider: model, - Access: access, - }) + return access } - - return c.JSON(http.StatusOK, response) } // isValidCategory returns true if cat is a recognized model category. diff --git a/internal/providers/responses_output.go b/internal/providers/responses_output.go index 284e63de0..de7ccb46a 100644 --- a/internal/providers/responses_output.go +++ b/internal/providers/responses_output.go @@ -146,23 +146,25 @@ func BuildResponsesOutputItems(msg core.ResponseMessage) []core.ResponsesOutputI // ConvertChatResponseToResponses converts a ChatResponse to a ResponsesResponse. func ConvertChatResponseToResponses(resp *core.ChatResponse) *core.ResponsesResponse { - output := []core.ResponsesOutputItem{ - { - ID: "msg_" + uuid.New().String(), - Type: "message", - Role: "assistant", - Status: "completed", - Content: []core.ResponsesContentItem{ - { - Type: "output_text", - Text: "", - Annotations: []json.RawMessage{}, - }, - }, - }, - } + var output []core.ResponsesOutputItem if len(resp.Choices) > 0 { output = BuildResponsesOutputItems(resp.Choices[0].Message) + } else { + output = []core.ResponsesOutputItem{ + { + ID: "msg_" + uuid.New().String(), + Type: "message", + Role: "assistant", + Status: "completed", + Content: []core.ResponsesContentItem{ + { + Type: "output_text", + Text: "", + Annotations: []json.RawMessage{}, + }, + }, + }, + } } return &core.ResponsesResponse{ From 7b1433067b90d5b3749d2f34e95f23f5cc95f623 Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Thu, 30 Apr 2026 09:18:45 +0200 Subject: [PATCH 12/15] fix: review-driven reliability and doc fixes MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Verified each comment from the second review pass against current code and applied the ones that are real correctness, reliability, or doc issues. Skipped pure style/perf and false positives. Admin handlers (internal/admin): - handler_audit.go: defensive nil-result check after auditReader.GetLogs so a (nil, nil) return does not panic in mocks/future drivers. - handler_audit.go: move AuditLog and AuditConversation request-shape validation ahead of the disabled-reader fast path, so callers always get 400 for malformed params regardless of whether audit logging is configured. - handler_models.go: ListModels @Success now references the real response type (admin.modelInventoryResponse) instead of providers.ModelWithProvider; adds @Param category and @Failure 400. - handler_models.go: trim whitespace from the category query param so inputs like "chat " are accepted. - handler_usage.go: UsageLog now rejects malformed/<=0 limit, enforces the documented max (200), and rejects negative offset — same shape as AuditLog rather than silently falling back to defaults. Background loops (internal/providers): - registry_init.go: StartBackgroundRefresh guards interval <= 0 and returns a no-op stop instead of panicking inside time.NewTicker. Matches the guard already present in workflows/authkeys/aliases StartBackgroundRefresh implementations. Guardrails (internal/guardrails): - batch_rewrite.go: rewriteGuardedChatBatchBody now propagates invalid_request_error from patchGuardedChatBatchBody instead of swallowing them. Falling back to Marshal(modified) on a validation error (e.g., guardrails inserted/reordered messages) would have silently published the invalid rewrite. Only fall back for raw-body preservation failures. Docs: - Regenerated cmd/gomodel/docs/docs.go and docs/openapi.json. The /admin/api/v1/models response is now properly modeled (with the embedded provider info plus the access object) and exposes admin.modelInventoryResponse + admin.modelAccessResponse as schemas. Co-Authored-By: Claude Opus 4.7 --- cmd/gomodel/docs/docs.go | 100 ++++++++++++++++++++----- docs/openapi.json | 106 ++++++++++++++++++++++----- internal/admin/handler_audit.go | 36 +++++---- internal/admin/handler_models.go | 8 +- internal/admin/handler_usage.go | 19 ++++- internal/guardrails/batch_rewrite.go | 10 ++- internal/providers/registry_init.go | 7 ++ 7 files changed, 229 insertions(+), 57 deletions(-) diff --git a/cmd/gomodel/docs/docs.go b/cmd/gomodel/docs/docs.go index 597a763b7..9892bb655 100644 --- a/cmd/gomodel/docs/docs.go +++ b/cmd/gomodel/docs/docs.go @@ -641,17 +641,31 @@ const docTemplate = `{ "tags": [ "admin" ], - "summary": "List all registered models with provider info", + "summary": "List all registered models with provider info and access state", + "parameters": [ + { + "type": "string", + "description": "Filter by model category", + "name": "category", + "in": "query" + } + ], "responses": { "200": { "description": "OK", "schema": { "type": "array", "items": { - "$ref": "#/definitions/providers.ModelWithProvider" + "$ref": "#/definitions/admin.modelInventoryResponse" } } }, + "400": { + "description": "Bad Request", + "schema": { + "$ref": "#/definitions/core.GatewayError" + } + }, "401": { "description": "Unauthorized", "schema": { @@ -3241,6 +3255,49 @@ const docTemplate = `{ } } }, + "admin.modelAccessResponse": { + "type": "object", + "properties": { + "default_enabled": { + "type": "boolean" + }, + "effective_enabled": { + "type": "boolean" + }, + "override": { + "$ref": "#/definitions/modeloverrides.Override" + }, + "selector": { + "type": "string" + }, + "user_paths": { + "type": "array", + "items": { + "type": "string" + } + } + } + }, + "admin.modelInventoryResponse": { + "type": "object", + "properties": { + "access": { + "$ref": "#/definitions/admin.modelAccessResponse" + }, + "model": { + "$ref": "#/definitions/core.Model" + }, + "provider_name": { + "type": "string" + }, + "provider_type": { + "type": "string" + }, + "selector": { + "type": "string" + } + } + }, "admin.recalculatePricingRequest": { "type": "object", "properties": { @@ -4802,33 +4859,42 @@ const docTemplate = `{ } } }, - "providers.CategoryCount": { + "modeloverrides.Override": { "type": "object", "properties": { - "category": { - "$ref": "#/definitions/core.ModelCategory" + "created_at": { + "type": "string" }, - "count": { - "type": "integer" + "model": { + "type": "string" }, - "display_name": { + "provider_name": { "type": "string" + }, + "selector": { + "type": "string" + }, + "updated_at": { + "type": "string" + }, + "user_paths": { + "type": "array", + "items": { + "type": "string" + } } } }, - "providers.ModelWithProvider": { + "providers.CategoryCount": { "type": "object", "properties": { - "model": { - "$ref": "#/definitions/core.Model" - }, - "provider_name": { - "type": "string" + "category": { + "$ref": "#/definitions/core.ModelCategory" }, - "provider_type": { - "type": "string" + "count": { + "type": "integer" }, - "selector": { + "display_name": { "type": "string" } } diff --git a/docs/openapi.json b/docs/openapi.json index c9fb99834..502476f06 100644 --- a/docs/openapi.json +++ b/docs/openapi.json @@ -785,7 +785,17 @@ "tags": [ "admin" ], - "summary": "List all registered models with provider info", + "summary": "List all registered models with provider info and access state", + "parameters": [ + { + "description": "Filter by model category", + "name": "category", + "in": "query", + "schema": { + "type": "string" + } + } + ], "responses": { "200": { "description": "OK", @@ -794,12 +804,22 @@ "schema": { "type": "array", "items": { - "$ref": "#/components/schemas/providers.ModelWithProvider" + "$ref": "#/components/schemas/admin.modelInventoryResponse" } } } } }, + "400": { + "description": "Bad Request", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.GatewayError" + } + } + } + }, "401": { "description": "Unauthorized", "content": { @@ -4710,6 +4730,49 @@ } } }, + "admin.modelAccessResponse": { + "type": "object", + "properties": { + "default_enabled": { + "type": "boolean" + }, + "effective_enabled": { + "type": "boolean" + }, + "override": { + "$ref": "#/components/schemas/modeloverrides.Override" + }, + "selector": { + "type": "string" + }, + "user_paths": { + "type": "array", + "items": { + "type": "string" + } + } + } + }, + "admin.modelInventoryResponse": { + "type": "object", + "properties": { + "access": { + "$ref": "#/components/schemas/admin.modelAccessResponse" + }, + "model": { + "$ref": "#/components/schemas/core.Model" + }, + "provider_name": { + "type": "string" + }, + "provider_type": { + "type": "string" + }, + "selector": { + "type": "string" + } + } + }, "admin.recalculatePricingRequest": { "type": "object", "properties": { @@ -6307,33 +6370,42 @@ } } }, - "providers.CategoryCount": { + "modeloverrides.Override": { "type": "object", "properties": { - "category": { - "$ref": "#/components/schemas/core.ModelCategory" + "created_at": { + "type": "string" }, - "count": { - "type": "integer" + "model": { + "type": "string" }, - "display_name": { + "provider_name": { + "type": "string" + }, + "selector": { "type": "string" + }, + "updated_at": { + "type": "string" + }, + "user_paths": { + "type": "array", + "items": { + "type": "string" + } } } }, - "providers.ModelWithProvider": { + "providers.CategoryCount": { "type": "object", "properties": { - "model": { - "$ref": "#/components/schemas/core.Model" - }, - "provider_name": { - "type": "string" + "category": { + "$ref": "#/components/schemas/core.ModelCategory" }, - "provider_type": { - "type": "string" + "count": { + "type": "integer" }, - "selector": { + "display_name": { "type": "string" } } diff --git a/internal/admin/handler_audit.go b/internal/admin/handler_audit.go index 1b2ed9fe9..59de5033f 100644 --- a/internal/admin/handler_audit.go +++ b/internal/admin/handler_audit.go @@ -39,12 +39,9 @@ import ( // @Failure 401 {object} core.GatewayError // @Router /admin/api/v1/audit/log [get] func (h *Handler) AuditLog(c *echo.Context) error { - if h.auditReader == nil { - return c.JSON(http.StatusOK, auditLogListResponse{ - Entries: []auditLogEntryResponse{}, - }) - } - + // Validate request shape before the disabled-reader fast path so callers + // always get a 400 for malformed inputs, regardless of whether audit + // logging is configured. dateRange, err := parseDateRangeParams(c) if err != nil { return handleError(c, err) @@ -104,11 +101,19 @@ func (h *Handler) AuditLog(c *echo.Context) error { params.Offset = parsed } + if h.auditReader == nil { + return c.JSON(http.StatusOK, auditLogListResponse{ + Entries: []auditLogEntryResponse{}, + }) + } + result, err := h.auditReader.GetLogs(c.Request().Context(), params) if err != nil { return handleError(c, err) } - + if result == nil { + result = &auditlog.LogListResult{Entries: []auditlog.LogEntry{}} + } if result.Entries == nil { result.Entries = []auditlog.LogEntry{} } @@ -174,13 +179,9 @@ func (h *Handler) auditLogResponse(ctx context.Context, result *auditlog.LogList // @Failure 401 {object} core.GatewayError // @Router /admin/api/v1/audit/conversation [get] func (h *Handler) AuditConversation(c *echo.Context) error { - if h.auditReader == nil { - return c.JSON(http.StatusOK, auditlog.ConversationResult{ - AnchorID: c.QueryParam("log_id"), - Entries: []auditlog.LogEntry{}, - }) - } - + // Validate request shape before the disabled-reader fast path so callers + // always get a 400 for missing/invalid params, regardless of whether + // audit logging is configured. logID := strings.TrimSpace(c.QueryParam("log_id")) if logID == "" { return handleError(c, core.NewInvalidRequestError("log_id is required", nil)) @@ -198,6 +199,13 @@ func (h *Handler) AuditConversation(c *echo.Context) error { limit = parsed } + if h.auditReader == nil { + return c.JSON(http.StatusOK, auditlog.ConversationResult{ + AnchorID: logID, + Entries: []auditlog.LogEntry{}, + }) + } + result, err := h.auditReader.GetConversation(c.Request().Context(), logID, limit) if err != nil { return handleError(c, err) diff --git a/internal/admin/handler_models.go b/internal/admin/handler_models.go index 3d3d2c119..d7856676d 100644 --- a/internal/admin/handler_models.go +++ b/internal/admin/handler_models.go @@ -28,11 +28,13 @@ type modelInventoryResponse struct { // ListModels handles GET /admin/api/v1/models // Supports optional ?category= query param for filtering by model category. // -// @Summary List all registered models with provider info +// @Summary List all registered models with provider info and access state // @Tags admin // @Produce json // @Security BearerAuth -// @Success 200 {array} providers.ModelWithProvider +// @Param category query string false "Filter by model category" +// @Success 200 {array} modelInventoryResponse +// @Failure 400 {object} core.GatewayError // @Failure 401 {object} core.GatewayError // @Router /admin/api/v1/models [get] func (h *Handler) ListModels(c *echo.Context) error { @@ -40,7 +42,7 @@ func (h *Handler) ListModels(c *echo.Context) error { return c.JSON(http.StatusOK, []modelInventoryResponse{}) } - cat := core.ModelCategory(c.QueryParam("category")) + cat := core.ModelCategory(strings.TrimSpace(c.QueryParam("category"))) if cat != "" && cat != core.CategoryAll { if !isValidCategory(cat) { return handleError(c, core.NewInvalidRequestError("invalid category: "+string(cat), nil)) diff --git a/internal/admin/handler_usage.go b/internal/admin/handler_usage.go index 8369cf6ec..3fe436b52 100644 --- a/internal/admin/handler_usage.go +++ b/internal/admin/handler_usage.go @@ -14,6 +14,10 @@ import ( "gomodel/internal/usage" ) +// maxUsageLogLimit caps the page size accepted by the usage log endpoint and +// matches the value documented in the @Param limit annotation below. +const maxUsageLogLimit = 200 + // UsageSummary handles GET /admin/api/v1/usage/summary // // @Summary Get usage summary @@ -175,14 +179,21 @@ func (h *Handler) UsageLog(c *echo.Context) error { } if l := c.QueryParam("limit"); l != "" { - if parsed, err := strconv.Atoi(l); err == nil && parsed > 0 { - params.Limit = parsed + parsed, err := strconv.Atoi(l) + if err != nil || parsed <= 0 { + return handleError(c, core.NewInvalidRequestError("invalid limit, expected positive integer", nil)) + } + if parsed > maxUsageLogLimit { + return handleError(c, core.NewInvalidRequestError("invalid limit parameter: limit must be between 1 and 200", nil)) } + params.Limit = parsed } if o := c.QueryParam("offset"); o != "" { - if parsed, err := strconv.Atoi(o); err == nil && parsed >= 0 { - params.Offset = parsed + parsed, err := strconv.Atoi(o) + if err != nil || parsed < 0 { + return handleError(c, core.NewInvalidRequestError("invalid offset, expected non-negative integer", nil)) } + params.Offset = parsed } result, err := h.usageReader.GetUsageLog(c.Request().Context(), params) diff --git a/internal/guardrails/batch_rewrite.go b/internal/guardrails/batch_rewrite.go index 1703acc27..3ce87a164 100644 --- a/internal/guardrails/batch_rewrite.go +++ b/internal/guardrails/batch_rewrite.go @@ -2,6 +2,7 @@ package guardrails import ( "encoding/json" + "errors" "gomodel/internal/core" ) @@ -14,8 +15,13 @@ func rewriteGuardedChatBatchBody(originalBody json.RawMessage, original *core.Ch if err == nil { return body, nil } - // Fallback: serialize the modified request directly when raw-body - // preservation fails (e.g., when the original body is malformed JSON). + // Validation errors (e.g. guardrails tried to insert/reorder messages) + // must propagate — falling back to Marshal(modified) would silently + // publish the invalid rewrite. Only fall back for raw-body preservation + // failures (malformed original JSON, etc.). + if gwErr, ok := errors.AsType[*core.GatewayError](err); ok && gwErr.Type == core.ErrorTypeInvalidRequest { + return nil, err + } return json.Marshal(modified) } diff --git a/internal/providers/registry_init.go b/internal/providers/registry_init.go index c309e889e..770dc9cd7 100644 --- a/internal/providers/registry_init.go +++ b/internal/providers/registry_init.go @@ -365,6 +365,13 @@ func (r *ModelRegistry) IsInitialized() bool { // for the goroutine to exit before returning, so callers should expect it to // block during shutdown until any in-flight refresh work unwinds. func (r *ModelRegistry) StartBackgroundRefresh(interval time.Duration, modelListURL string) func() { + if interval <= 0 { + // time.NewTicker panics on non-positive durations and a refresh loop + // with a zero interval would be meaningless. Skip the goroutine and + // hand back a no-op stop so callers can still defer it safely. + slog.Debug("model registry background refresh disabled", "interval", interval) + return func() {} + } ctx, cancel := context.WithCancel(context.Background()) done := make(chan struct{}) var stopOnce sync.Once From b2a3ab38b2581eacf443b1f62d9924a823d3bd9f Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Thu, 30 Apr 2026 10:56:17 +0200 Subject: [PATCH 13/15] fix: review-driven hardening for admin handlers and guardrails MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Six fixes across admin handlers, guardrails batch rewriting, registry initialization, and OpenAPI generation. Each verified as a real bug or compliance gap against the current code. Admin (internal/admin): - handler_audit.go: AuditLog now caps the limit query param at the documented maximum of 100 (named const maxAuditLogLimit) so callers asking for 1000 rows are not silently allowed past the doc. - handler_models.go: DashboardConfig had no Swagger annotations, so the endpoint never appeared in docs.go / openapi.json. Added @Summary, @Tags, @Produce, @Security, @Success (DashboardConfigResponse), @Failure 401, and @Router. The endpoint is now documented. - handler_usage.go: defensive nil-result guards after usageReader.GetSummary and usageReader.GetUsageLog so a (nil, nil) return never panics; matches the pattern used in handler_audit.go. Guardrails (internal/guardrails): - batch_rewrite.go: rewriteGuardedChatBatchBody and patchGuardedChatBatchBody now reject a nil original chat request with invalid_request_error before any *original.Messages dereference. - batch_rewrite_test.go (new): table-driven coverage for nil modified, nil original, successful raw-body patch, validation-error propagation (the new InvalidRequest check), and raw-body parse failure falling back to Marshal(modified). Locks the propagation behavior added in the prior commit. Providers (internal/providers): - registry_init.go: Initialize, RefreshModelList, and InitializeAsync all normalize a nil ctx to context.Background() at entry. Previously the local normalize in acquireRefresh was discarded, so initialize / refreshModelListLocked / context.WithTimeout(nil, …) could panic on a nil caller ctx. OpenAPI tooling (tools, docs): - openapi-postprocess.mjs: helper to attach maxItems on chosen array responses. Bound /admin/api/v1/models to maxItems=10000 so unbounded- array scanners (CKV_OPENAPI_21) stop flagging it. Runtime size is bounded by configured providers + the model list registry, well inside that ceiling. - Regenerated cmd/gomodel/docs/docs.go and docs/openapi.json. The /admin/api/v1/dashboard/config path is now present, the /models array carries the bound, and admin.DashboardConfigResponse is exposed as a schema. Co-Authored-By: Claude Opus 4.7 --- cmd/gomodel/docs/docs.go | 62 +++++++++++++ docs/openapi.json | 71 ++++++++++++++- internal/admin/handler_audit.go | 7 ++ internal/admin/handler_models.go | 8 ++ internal/admin/handler_usage.go | 7 +- internal/guardrails/batch_rewrite.go | 6 ++ internal/guardrails/batch_rewrite_test.go | 106 ++++++++++++++++++++++ internal/providers/registry_init.go | 9 ++ tools/openapi-postprocess.mjs | 25 +++++ 9 files changed, 299 insertions(+), 2 deletions(-) create mode 100644 internal/guardrails/batch_rewrite_test.go diff --git a/cmd/gomodel/docs/docs.go b/cmd/gomodel/docs/docs.go index 9892bb655..989f84ded 100644 --- a/cmd/gomodel/docs/docs.go +++ b/cmd/gomodel/docs/docs.go @@ -633,6 +633,36 @@ const docTemplate = `{ ] } }, + "/admin/api/v1/dashboard/config": { + "get": { + "produces": [ + "application/json" + ], + "tags": [ + "admin" + ], + "summary": "Get dashboard runtime configuration", + "responses": { + "200": { + "description": "OK", + "schema": { + "$ref": "#/definitions/admin.DashboardConfigResponse" + } + }, + "401": { + "description": "Unauthorized", + "schema": { + "$ref": "#/definitions/core.GatewayError" + } + } + }, + "security": [ + { + "BearerAuth": [] + } + ] + } + }, "/admin/api/v1/models": { "get": { "produces": [ @@ -3089,6 +3119,38 @@ const docTemplate = `{ } }, "definitions": { + "admin.DashboardConfigResponse": { + "type": "object", + "properties": { + "BUDGETS_ENABLED": { + "type": "string" + }, + "CACHE_ENABLED": { + "type": "string" + }, + "FEATURE_FALLBACK_MODE": { + "type": "string" + }, + "GUARDRAILS_ENABLED": { + "type": "string" + }, + "LOGGING_ENABLED": { + "type": "string" + }, + "REDIS_URL": { + "type": "string" + }, + "SEMANTIC_CACHE_ENABLED": { + "type": "string" + }, + "USAGE_ENABLED": { + "type": "string" + }, + "USAGE_PRICING_RECALCULATION_ENABLED": { + "type": "string" + } + } + }, "admin.auditLogEntryResponse": { "type": "object", "properties": { diff --git a/docs/openapi.json b/docs/openapi.json index 502476f06..21dda063f 100644 --- a/docs/openapi.json +++ b/docs/openapi.json @@ -780,6 +780,41 @@ ] } }, + "/admin/api/v1/dashboard/config": { + "get": { + "tags": [ + "admin" + ], + "summary": "Get dashboard runtime configuration", + "responses": { + "200": { + "description": "OK", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/admin.DashboardConfigResponse" + } + } + } + }, + "401": { + "description": "Unauthorized", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/core.GatewayError" + } + } + } + } + }, + "security": [ + { + "BearerAuth": [] + } + ] + } + }, "/admin/api/v1/models": { "get": { "tags": [ @@ -805,7 +840,9 @@ "type": "array", "items": { "$ref": "#/components/schemas/admin.modelInventoryResponse" - } + }, + "maxItems": 10000, + "description": "Bounded by maxItems=10000." } } } @@ -4564,6 +4601,38 @@ } }, "schemas": { + "admin.DashboardConfigResponse": { + "type": "object", + "properties": { + "BUDGETS_ENABLED": { + "type": "string" + }, + "CACHE_ENABLED": { + "type": "string" + }, + "FEATURE_FALLBACK_MODE": { + "type": "string" + }, + "GUARDRAILS_ENABLED": { + "type": "string" + }, + "LOGGING_ENABLED": { + "type": "string" + }, + "REDIS_URL": { + "type": "string" + }, + "SEMANTIC_CACHE_ENABLED": { + "type": "string" + }, + "USAGE_ENABLED": { + "type": "string" + }, + "USAGE_PRICING_RECALCULATION_ENABLED": { + "type": "string" + } + } + }, "admin.auditLogEntryResponse": { "type": "object", "properties": { diff --git a/internal/admin/handler_audit.go b/internal/admin/handler_audit.go index 59de5033f..85091a61a 100644 --- a/internal/admin/handler_audit.go +++ b/internal/admin/handler_audit.go @@ -14,6 +14,10 @@ import ( "gomodel/internal/usage" ) +// maxAuditLogLimit caps the page size accepted by the audit log endpoint and +// matches the value documented in the @Param limit annotation below. +const maxAuditLogLimit = 100 + // AuditLog handles GET /admin/api/v1/audit/log // // @Summary Get paginated audit log entries @@ -91,6 +95,9 @@ func (h *Handler) AuditLog(c *echo.Context) error { if err != nil || parsed <= 0 { return handleError(c, core.NewInvalidRequestError("invalid limit, expected positive integer", nil)) } + if parsed > maxAuditLogLimit { + return handleError(c, core.NewInvalidRequestError("invalid limit parameter: limit must be between 1 and 100", nil)) + } params.Limit = parsed } if o := c.QueryParam("offset"); o != "" { diff --git a/internal/admin/handler_models.go b/internal/admin/handler_models.go index d7856676d..1ae5cdc17 100644 --- a/internal/admin/handler_models.go +++ b/internal/admin/handler_models.go @@ -128,6 +128,14 @@ func (h *Handler) ListCategories(c *echo.Context) error { } // DashboardConfig handles GET /admin/api/v1/dashboard/config +// +// @Summary Get dashboard runtime configuration +// @Tags admin +// @Produce json +// @Security BearerAuth +// @Success 200 {object} DashboardConfigResponse +// @Failure 401 {object} core.GatewayError +// @Router /admin/api/v1/dashboard/config [get] func (h *Handler) DashboardConfig(c *echo.Context) error { return c.JSON(http.StatusOK, cloneDashboardRuntimeConfig(h.runtimeConfig)) } diff --git a/internal/admin/handler_usage.go b/internal/admin/handler_usage.go index 3fe436b52..a106e3e1d 100644 --- a/internal/admin/handler_usage.go +++ b/internal/admin/handler_usage.go @@ -47,6 +47,9 @@ func (h *Handler) UsageSummary(c *echo.Context) error { if err != nil { return handleError(c, err) } + if summary == nil { + summary = &usage.UsageSummary{} + } return c.JSON(http.StatusOK, summary) } @@ -200,7 +203,9 @@ func (h *Handler) UsageLog(c *echo.Context) error { if err != nil { return handleError(c, err) } - + if result == nil { + result = &usage.UsageLogResult{Entries: []usage.UsageLogEntry{}} + } if result.Entries == nil { result.Entries = []usage.UsageLogEntry{} } diff --git a/internal/guardrails/batch_rewrite.go b/internal/guardrails/batch_rewrite.go index 3ce87a164..5d8f0f3a8 100644 --- a/internal/guardrails/batch_rewrite.go +++ b/internal/guardrails/batch_rewrite.go @@ -11,6 +11,9 @@ func rewriteGuardedChatBatchBody(originalBody json.RawMessage, original *core.Ch if modified == nil { return nil, core.NewInvalidRequestError("missing guarded chat request", nil) } + if original == nil { + return nil, core.NewInvalidRequestError("missing original chat request", nil) + } body, err := patchGuardedChatBatchBody(originalBody, original, modified) if err == nil { return body, nil @@ -29,6 +32,9 @@ func patchGuardedChatBatchBody(originalBody json.RawMessage, original *core.Chat if modified == nil { return nil, core.NewInvalidRequestError("missing guarded chat request", nil) } + if original == nil { + return nil, core.NewInvalidRequestError("missing original chat request", nil) + } var raw map[string]json.RawMessage if err := json.Unmarshal(originalBody, &raw); err != nil { diff --git a/internal/guardrails/batch_rewrite_test.go b/internal/guardrails/batch_rewrite_test.go new file mode 100644 index 000000000..da1d8fe2f --- /dev/null +++ b/internal/guardrails/batch_rewrite_test.go @@ -0,0 +1,106 @@ +package guardrails + +import ( + "encoding/json" + "errors" + "strings" + "testing" + + "gomodel/internal/core" +) + +func TestRewriteGuardedChatBatchBody(t *testing.T) { + makeReq := func(role, content string) *core.ChatRequest { + return &core.ChatRequest{ + Model: "gpt-4", + Messages: []core.Message{{Role: role, Content: content}}, + } + } + + originalBody := func(req *core.ChatRequest) json.RawMessage { + body, err := json.Marshal(req) + if err != nil { + t.Fatalf("marshal helper: %v", err) + } + return body + } + + tests := []struct { + name string + originalBody func(orig *core.ChatRequest) json.RawMessage + original *core.ChatRequest + modified *core.ChatRequest + wantErrIs core.ErrorType // empty = expect success + wantBodyHas string // substring assertion when no error + }{ + { + name: "nil modified rejected with invalid_request_error", + originalBody: originalBody, + original: makeReq("user", "hello"), + modified: nil, + wantErrIs: core.ErrorTypeInvalidRequest, + }, + { + name: "nil original rejected with invalid_request_error", + originalBody: originalBody, + original: nil, + modified: makeReq("user", "hello"), + wantErrIs: core.ErrorTypeInvalidRequest, + }, + { + name: "successful raw-body patch returns patched body", + originalBody: originalBody, + original: makeReq("user", "hello"), + modified: makeReq("user", "rewritten"), + wantBodyHas: `"rewritten"`, + }, + { + name: "validation error from message reorder propagates as invalid_request_error", + originalBody: originalBody, + original: makeReq("user", "hello"), + modified: &core.ChatRequest{ + Model: "gpt-4", + // guardrails inserted a non-system message — patcher returns InvalidRequest. + Messages: []core.Message{ + {Role: "user", Content: "hello"}, + {Role: "user", Content: "extra-injected"}, + }, + }, + wantErrIs: core.ErrorTypeInvalidRequest, + }, + { + name: "raw-body parse failure falls back to Marshal(modified)", + originalBody: func(orig *core.ChatRequest) json.RawMessage { + return json.RawMessage(`not valid json`) + }, + original: makeReq("user", "hello"), + modified: makeReq("user", "rewritten"), + wantBodyHas: `"rewritten"`, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + body, err := rewriteGuardedChatBatchBody(tt.originalBody(tt.original), tt.original, tt.modified) + if tt.wantErrIs != "" { + if err == nil { + t.Fatalf("expected error type %q, got nil", tt.wantErrIs) + } + var gwErr *core.GatewayError + if !errors.As(err, &gwErr) { + t.Fatalf("expected *core.GatewayError, got %T: %v", err, err) + } + if gwErr.Type != tt.wantErrIs { + t.Fatalf("expected error type %q, got %q", tt.wantErrIs, gwErr.Type) + } + return + } + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if !strings.Contains(string(body), tt.wantBodyHas) { + t.Fatalf("expected body to contain %q, got %s", tt.wantBodyHas, body) + } + }) + } +} diff --git a/internal/providers/registry_init.go b/internal/providers/registry_init.go index 770dc9cd7..262eb58ef 100644 --- a/internal/providers/registry_init.go +++ b/internal/providers/registry_init.go @@ -19,6 +19,9 @@ import ( // Initialize fetches models from all registered providers and populates the registry. // This should be called on application startup. func (r *ModelRegistry) Initialize(ctx context.Context) error { + if ctx == nil { + ctx = context.Background() + } release, err := r.acquireRefresh(ctx) if err != nil { return err @@ -324,6 +327,9 @@ func registryRefreshAcquireError(err error) *core.GatewayError { // Returns immediately after loading cache. The background goroutine will update models // and save to cache when network fetch completes. func (r *ModelRegistry) InitializeAsync(ctx context.Context) { + if ctx == nil { + ctx = context.Background() + } // First, try to load from cache for instant startup cached, err := r.LoadFromCache(ctx) if err != nil { @@ -428,6 +434,9 @@ func (r *ModelRegistry) RefreshModelList(ctx context.Context, url string) (int, if strings.TrimSpace(url) == "" { return 0, nil } + if ctx == nil { + ctx = context.Background() + } release, err := r.acquireRefresh(ctx) if err != nil { diff --git a/tools/openapi-postprocess.mjs b/tools/openapi-postprocess.mjs index 40acb93ca..2b6604d14 100644 --- a/tools/openapi-postprocess.mjs +++ b/tools/openapi-postprocess.mjs @@ -119,11 +119,36 @@ function ensureRequiredProperty(schemaName, propertyName) { target.required = Array.from(required).sort(); } +function applyArrayMaxItems(operationPath, method, statusCode, maxItems) { + const op = spec.paths?.[operationPath]?.[method]; + if (!op) { + throw new Error(`missing OpenAPI operation: ${method.toUpperCase()} ${operationPath}`); + } + const response = op.responses?.[statusCode]; + if (!response) { + throw new Error(`missing response ${statusCode} on ${method.toUpperCase()} ${operationPath}`); + } + const schemaRef = response.content?.["application/json"]?.schema || response.schema; + if (!schemaRef || schemaRef.type !== "array") { + throw new Error(`expected array schema on ${method.toUpperCase()} ${operationPath} ${statusCode}`); + } + schemaRef.maxItems = maxItems; + if (!schemaRef.description) { + schemaRef.description = `Bounded by maxItems=${maxItems}.`; + } +} + spec.servers = parseServers(process.env.DOCS_API_SERVERS); ensureResponsesInputElementSchema(); ensureBearerAuthSecurityScheme(); ensureRequiredProperty("admin.recalculatePricingRequest", "confirmation"); +// Bound the registry-backed admin model listing so OpenAPI consumers (and +// security scanners like CKV_OPENAPI_21) see an explicit upper limit. The +// runtime registry is bounded by configured providers and the backing +// model list; 10000 leaves substantial headroom for that worst case. +applyArrayMaxItems("/admin/api/v1/models", "get", "200", 10000); + for (const name of [ "core.ResponsesRequest", "core.ResponseInputTokensRequest", From 5796979067110c821e6d6bc1dc7e424268479510 Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Thu, 30 Apr 2026 11:24:29 +0200 Subject: [PATCH 14/15] fix: validate usage queries early; clear stale provider runtime errors Two review-driven fixes. handler_usage.go: UsageSummary, usageSliceResponse (DailyUsage, UsageByModel, UsageByUserPath), UsageLog, and CacheOverview all returned an empty success payload when the usage reader was nil without first validating query params. Move parseUsageParams (and the limit/offset parsing in UsageLog) ahead of the disabled-reader fast path so a malformed request gets a 400 regardless of whether usage tracking is wired up. CacheOverview keeps its feature-gate check first; the gate is conceptually higher-priority than param validation since the endpoint is unavailable when cache analytics is off. registry_init.go (applyProviderRuntimeUpdatesLocked): a successful allowlist-mode refresh produces an update with empty lastModelFetchError but no lastModelFetchSuccessAt bump (semantics: SuccessAt tracks genuine upstream success, which never happens in allowlist mode). The previous apply logic only cleared current.error when SuccessAt was set, leaving a stale error from a prior failed refresh visible in runtime status forever. Treat any non-zero update.lastModelFetchAt as authoritative: overwrite current.lastModelFetchError unconditionally on a refresh attempt (empty = success, non-empty = failure). This matches every update construction path in initialize(), each of which already sets lastModelFetchError to either the failure message or configuredUpstreamError. Adds TestApplyProviderRuntimeUpdates_ClearsStaleErrorOnSuccessful- Refresh to lock the new behavior. Existing tests still pass: ConfiguredModelsFallbackModeUsesConfiguredWhenUpstreamFails still records the upstream error (its update has non-empty configuredUpstreamError), and ConfiguredModelsAllowlistModeSkipsUpstream still leaves SuccessAt nil. Co-Authored-By: Claude Opus 4.7 --- internal/admin/handler_usage.go | 50 ++++++++++++++++++----------- internal/providers/registry_init.go | 13 ++++---- internal/providers/registry_test.go | 31 ++++++++++++++++++ 3 files changed, 69 insertions(+), 25 deletions(-) diff --git a/internal/admin/handler_usage.go b/internal/admin/handler_usage.go index a106e3e1d..cb196ac4e 100644 --- a/internal/admin/handler_usage.go +++ b/internal/admin/handler_usage.go @@ -34,15 +34,17 @@ const maxUsageLogLimit = 200 // @Failure 401 {object} core.GatewayError // @Router /admin/api/v1/usage/summary [get] func (h *Handler) UsageSummary(c *echo.Context) error { - if h.usageReader == nil { - return c.JSON(http.StatusOK, usage.UsageSummary{}) - } - + // Validate request shape before the disabled-reader fast path so callers + // always get a 400 for malformed inputs, regardless of wiring. params, err := parseUsageParams(c) if err != nil { return handleError(c, err) } + if h.usageReader == nil { + return c.JSON(http.StatusOK, usage.UsageSummary{}) + } + summary, err := h.usageReader.GetSummary(c.Request().Context(), params) if err != nil { return handleError(c, err) @@ -59,15 +61,17 @@ func usageSliceResponse[T any]( reader usage.UsageReader, fetch func(context.Context, usage.UsageQueryParams) ([]T, error), ) error { - if reader == nil { - return c.JSON(http.StatusOK, []T{}) - } - + // Validate before the disabled-reader fast path so malformed query + // params produce a 400 even when usage tracking is disabled. params, err := parseUsageParams(c) if err != nil { return handleError(c, err) } + if reader == nil { + return c.JSON(http.StatusOK, []T{}) + } + values, err := fetch(c.Request().Context(), params) if err != nil { return handleError(c, err) @@ -163,12 +167,8 @@ func (h *Handler) UsageByUserPath(c *echo.Context) error { // @Failure 401 {object} core.GatewayError // @Router /admin/api/v1/usage/log [get] func (h *Handler) UsageLog(c *echo.Context) error { - if h.usageReader == nil { - return c.JSON(http.StatusOK, usage.UsageLogResult{ - Entries: []usage.UsageLogEntry{}, - }) - } - + // Validate request shape before the disabled-reader fast path so callers + // always get a 400 for malformed inputs, regardless of wiring. baseParams, err := parseUsageParams(c) if err != nil { return handleError(c, err) @@ -199,6 +199,12 @@ func (h *Handler) UsageLog(c *echo.Context) error { params.Offset = parsed } + if h.usageReader == nil { + return c.JSON(http.StatusOK, usage.UsageLogResult{ + Entries: []usage.UsageLogEntry{}, + }) + } + result, err := h.usageReader.GetUsageLog(c.Request().Context(), params) if err != nil { return handleError(c, err) @@ -282,21 +288,27 @@ func (h *Handler) RecalculateUsagePricing(c *echo.Context) error { // @Failure 503 {object} core.GatewayError // @Router /admin/api/v1/cache/overview [get] func (h *Handler) CacheOverview(c *echo.Context) error { + // Feature-gate check stays first: this endpoint is conceptually unavailable + // when cache analytics is off, and the response shape (503) communicates + // that to the dashboard. if strings.TrimSpace(h.runtimeConfig.CacheEnabled) != "on" { return handleError(c, featureUnavailableError("cache analytics is unavailable")) } - if h.usageReader == nil { - return c.JSON(http.StatusOK, usage.CacheOverview{ - Daily: []usage.CacheOverviewDaily{}, - }) - } + // Validate request shape before the disabled-reader fast path so callers + // always get a 400 for malformed inputs, regardless of wiring. params, err := parseUsageParams(c) if err != nil { return handleError(c, err) } params.CacheMode = usage.CacheModeCached + if h.usageReader == nil { + return c.JSON(http.StatusOK, usage.CacheOverview{ + Daily: []usage.CacheOverviewDaily{}, + }) + } + overview, err := h.usageReader.GetCacheOverview(c.Request().Context(), params) if err != nil { return handleError(c, err) diff --git a/internal/providers/registry_init.go b/internal/providers/registry_init.go index 262eb58ef..cb9443135 100644 --- a/internal/providers/registry_init.go +++ b/internal/providers/registry_init.go @@ -270,15 +270,16 @@ func (r *ModelRegistry) applyProviderRuntimeUpdatesLocked(updates map[string]pro current.registered = update.registered || current.registered if !update.lastModelFetchAt.IsZero() { current.lastModelFetchAt = update.lastModelFetchAt + // A non-zero fetchAt represents a refresh attempt whose outcome + // is captured authoritatively in lastModelFetchError (empty = + // success, non-empty = failure). Overwrite unconditionally so an + // old error doesn't survive a subsequent successful refresh — + // this matters in particular for allowlist-mode refreshes which + // don't bump SuccessAt but still produce usable models. + current.lastModelFetchError = strings.TrimSpace(update.lastModelFetchError) } if !update.lastModelFetchSuccessAt.IsZero() { current.lastModelFetchSuccessAt = update.lastModelFetchSuccessAt - if strings.TrimSpace(update.lastModelFetchError) == "" { - current.lastModelFetchError = "" - } - } - if strings.TrimSpace(update.lastModelFetchError) != "" { - current.lastModelFetchError = update.lastModelFetchError } r.providerRuntime[providerName] = current } diff --git a/internal/providers/registry_test.go b/internal/providers/registry_test.go index abf2ae249..cd67b31fa 100644 --- a/internal/providers/registry_test.go +++ b/internal/providers/registry_test.go @@ -1345,6 +1345,37 @@ func (c *countingRegistryMockProvider) ListModels(ctx context.Context) (*core.Mo return c.registryMockProvider.ListModels(ctx) } +// TestApplyProviderRuntimeUpdates_ClearsStaleErrorOnSuccessfulRefresh locks the +// behavior that a successful refresh (non-zero fetchAt + empty fetch error) +// clears any error left over from a previous failed refresh, regardless of +// whether the success bumps lastModelFetchSuccessAt. The allowlist case is +// the original motivator: SuccessAt stays nil because upstream is never +// called, but a stale error must not survive into runtime status. +func TestApplyProviderRuntimeUpdates_ClearsStaleErrorOnSuccessfulRefresh(t *testing.T) { + registry := NewModelRegistry() + + // Seed runtime state with a prior error. + registry.providerRuntime["test"] = providerRuntimeState{ + registered: true, + lastModelFetchAt: time.Now().Add(-time.Hour), + lastModelFetchError: "previous upstream failure", + } + + // Apply a successful refresh that produced usable models without + // touching upstream — mimics allowlist mode. + registry.applyProviderRuntimeUpdatesLocked(map[string]providerRuntimeState{ + "test": { + registered: true, + lastModelFetchAt: time.Now(), + // lastModelFetchError intentionally empty; SuccessAt deliberately zero. + }, + }) + + if got := registry.providerRuntime["test"].lastModelFetchError; got != "" { + t.Fatalf("lastModelFetchError = %q, want empty after successful refresh", got) + } +} + func TestStartBackgroundRefresh(t *testing.T) { t.Run("RefreshesAtInterval", func(t *testing.T) { var refreshCount atomic.Int32 From 1180539a99e7925f3221a7231c815e3974889d0c Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Thu, 30 Apr 2026 11:32:57 +0200 Subject: [PATCH 15/15] test(admin): coverage gaps surfaced during review MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Two follow-up tests recommended in the PR self-review: handler_test.go (3 new tests): - TestAuditLog_NilReaderStillValidatesParams asserts that ?status_code=not-an-int returns 400 even when the audit reader is nil. - TestAuditConversation_NilReaderStillValidatesParams asserts that a request missing the required log_id returns 400 even when the audit reader is nil. - TestUsageLog_NilReaderStillValidatesParams asserts that ?start_date=not-a-date returns 400 even when the usage reader is nil. These lock the validation-before-fast-path ordering introduced earlier in the branch — the existing *_NilReader tests only exercised the happy path through the disabled-reader fast path. routes_test.go (new file): - TestRegisterRoutes_RegistersExpectedPaths is a smoke test that mounts the admin handler on a real *echo.Echo group, asserts RegisterRoutes doesn't panic with a zero-value handler, and walks the resulting router to verify every expected method+path is wired and that no extra routes were registered. Catches typos and missing wires when endpoints are added or renamed without updating routes.go (or vice versa). Co-Authored-By: Claude Opus 4.7 --- internal/admin/handler_test.go | 53 +++++++++++++++ internal/admin/routes_test.go | 115 +++++++++++++++++++++++++++++++++ 2 files changed, 168 insertions(+) create mode 100644 internal/admin/routes_test.go diff --git a/internal/admin/handler_test.go b/internal/admin/handler_test.go index 774316a07..ea373efae 100644 --- a/internal/admin/handler_test.go +++ b/internal/admin/handler_test.go @@ -1260,6 +1260,59 @@ func TestAuditConversation_Error(t *testing.T) { } } +// --- Validation-before-fast-path tests --- +// +// AuditLog, AuditConversation, and UsageLog all short-circuit to an empty +// success payload when their reader is nil. These tests assert that request- +// shape validation runs *before* that fast path, so callers get a 400 for +// missing/malformed required params regardless of whether the underlying +// reader is wired up. + +func TestAuditLog_NilReaderStillValidatesParams(t *testing.T) { + h := NewHandler(nil, nil) // no audit reader configured + c, rec := newHandlerContext("/admin/api/v1/audit/log?status_code=not-an-int") + + if err := h.AuditLog(c); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if rec.Code != http.StatusBadRequest { + t.Errorf("expected 400, got %d", rec.Code) + } + if !containsString(rec.Body.String(), "invalid_request_error") { + t.Errorf("expected invalid_request_error, got: %s", rec.Body.String()) + } +} + +func TestAuditConversation_NilReaderStillValidatesParams(t *testing.T) { + h := NewHandler(nil, nil) // no audit reader configured + c, rec := newHandlerContext("/admin/api/v1/audit/conversation") // missing required log_id + + if err := h.AuditConversation(c); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if rec.Code != http.StatusBadRequest { + t.Errorf("expected 400, got %d", rec.Code) + } + if !containsString(rec.Body.String(), "log_id is required") { + t.Errorf("expected log_id-is-required, got: %s", rec.Body.String()) + } +} + +func TestUsageLog_NilReaderStillValidatesParams(t *testing.T) { + h := NewHandler(nil, nil) // no usage reader configured + c, rec := newHandlerContext("/admin/api/v1/usage/log?start_date=not-a-date") + + if err := h.UsageLog(c); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if rec.Code != http.StatusBadRequest { + t.Errorf("expected 400, got %d", rec.Code) + } + if !containsString(rec.Body.String(), "invalid_request_error") { + t.Errorf("expected invalid_request_error, got: %s", rec.Body.String()) + } +} + // --- ListModels handler tests --- func TestListModels_NilRegistry(t *testing.T) { diff --git a/internal/admin/routes_test.go b/internal/admin/routes_test.go new file mode 100644 index 000000000..4f10a5080 --- /dev/null +++ b/internal/admin/routes_test.go @@ -0,0 +1,115 @@ +package admin + +import ( + "sort" + "testing" + + "github.com/labstack/echo/v5" +) + +// TestRegisterRoutes_RegistersExpectedPaths is a smoke test for the admin +// RouteRegistrar plumbing. It mounts the handler on a real echo router and +// verifies that every method+path the route table claims to register is +// actually known to the router after RegisterRoutes returns. +// +// The intent is to catch regressions when handlers are added or renamed +// without updating routes.go (or vice-versa) — including typos and missing +// wires that would otherwise only surface in production traffic. +func TestRegisterRoutes_RegistersExpectedPaths(t *testing.T) { + h := &Handler{} + e := echo.New() + g := e.Group("/admin/api/v1") + + // RegisterRoutes must not panic with a zero-value handler — every endpoint + // reads its own dependencies inside the handler body, so route mounting + // itself must remain side-effect-free. + defer func() { + if r := recover(); r != nil { + t.Fatalf("RegisterRoutes panicked: %v", r) + } + }() + h.RegisterRoutes(g) + + want := []string{ + "GET /admin/api/v1/dashboard/config", + "GET /admin/api/v1/cache/overview", + + "GET /admin/api/v1/usage/summary", + "GET /admin/api/v1/usage/daily", + "GET /admin/api/v1/usage/models", + "GET /admin/api/v1/usage/user-paths", + "GET /admin/api/v1/usage/log", + "POST /admin/api/v1/usage/recalculate-pricing", + + "GET /admin/api/v1/audit/log", + "GET /admin/api/v1/audit/conversation", + + "GET /admin/api/v1/providers/status", + "POST /admin/api/v1/runtime/refresh", + + "GET /admin/api/v1/budgets", + "PUT /admin/api/v1/budgets/:user_path/:period", + "DELETE /admin/api/v1/budgets/:user_path/:period", + "GET /admin/api/v1/budgets/settings", + "PUT /admin/api/v1/budgets/settings", + "POST /admin/api/v1/budgets/reset-one", + "POST /admin/api/v1/budgets/reset", + + "GET /admin/api/v1/models", + "GET /admin/api/v1/models/categories", + + "GET /admin/api/v1/model-overrides", + "PUT /admin/api/v1/model-overrides/:selector", + "DELETE /admin/api/v1/model-overrides/:selector", + + "GET /admin/api/v1/auth-keys", + "POST /admin/api/v1/auth-keys", + "POST /admin/api/v1/auth-keys/:id/deactivate", + + "GET /admin/api/v1/aliases", + "PUT /admin/api/v1/aliases/:name", + "DELETE /admin/api/v1/aliases/:name", + + "GET /admin/api/v1/guardrails/types", + "GET /admin/api/v1/guardrails", + "PUT /admin/api/v1/guardrails/:name", + "DELETE /admin/api/v1/guardrails/:name", + + "GET /admin/api/v1/workflows", + "GET /admin/api/v1/workflows/guardrails", + "GET /admin/api/v1/workflows/:id", + "POST /admin/api/v1/workflows", + "POST /admin/api/v1/workflows/:id/deactivate", + } + + registered := make(map[string]struct{}) + for _, route := range e.Router().Routes() { + registered[route.Method+" "+route.Path] = struct{}{} + } + + sort.Strings(want) + missing := make([]string, 0) + for _, key := range want { + if _, ok := registered[key]; !ok { + missing = append(missing, key) + } + } + if len(missing) != 0 { + t.Fatalf("RegisterRoutes did not register %d route(s):\n %s", len(missing), missing) + } + + if got, expected := len(registered), len(want); got != expected { + extras := make([]string, 0) + wantSet := make(map[string]struct{}, len(want)) + for _, k := range want { + wantSet[k] = struct{}{} + } + for k := range registered { + if _, ok := wantSet[k]; !ok { + extras = append(extras, k) + } + } + sort.Strings(extras) + t.Fatalf("RegisterRoutes registered %d route(s), want %d; extras: %v", got, expected, extras) + } +}