From b4dc659573d8c2ad7fec06e9017b74025cec4f34 Mon Sep 17 00:00:00 2001 From: "Jakub A. W" Date: Mon, 6 Jul 2026 15:29:33 +0200 Subject: [PATCH] chore(cleanup): remove dead code and relocate test-only seams Systematic dead-code sweep driven by golang.org/x/tools/cmd/deadcode (binary and -test modes) and staticcheck U1000. Net -2,148 lines. Removed: - The unused responsecache stream-reconstruction subsystem (stream_cache_chat.go, stream_cache_responses.go, and the dead half of stream_cache.go), superseded by verbatim raw-SSE replay, plus the tests that only exercised it. - Dead thin wrappers, migrating their test call sites to the live variants: workflows.NewCompiler, dashboard.New, gateway.ResolveRequestModel, server resolveRequestModel, DetermineBatchExecutionSelection(WithAuthorizer), ensureTranslatedRequestWorkflow, doOpenAICompatibleFileIDRequest, and the batchrewrite RecordPreparation/CleanupSuperseded* family (the guardrails harness now mirrors production via RecordResult/CleanupFile). - Dead functions together with their dead tests: observability.GetMetrics/HealthCheck, core.HasNonTextContent, core.NormalizeModelSelector, the DecodedBatchItemRequest accessors, BatchItemModelSelector, auditlog.DefaultConfig, conversationstore.WithUnboundedRetention, responsecache guardrails-hash and cache-skip wrappers, and the unused MockEmbedder test type. Relocated production-dead but test-used seams into test files: server NewHandler/newHandler/AuthMiddleware/WorkflowResolution( WithResolver)/GetProviderType/handleStreamingResponse (new seams_test.go), vecstore_map.go, NewRedisModelCacheWithStore, failover.NewResolver, admin.WithGuardrailsRegistry, derivedEnvNames. Deliberately kept: the ext package public extension API, the provider NewWithHTTPClient constructor family (contract tests and cross-provider delegation), and cross-package test seams (cache.MapStore, llmclient.DefaultConfig, ResetMetrics, NewResponseCacheMiddlewareWithStore, memory-store retention options). deadcode -test now reports only the intentional ext API; staticcheck U1000 (including tests) is clean; all packages pass, and contract/ integration suites compile under their build tags. Co-Authored-By: Claude Fable 5 --- internal/admin/dashboard/dashboard.go | 6 +- internal/admin/dashboard/dashboard_test.go | 54 +- internal/admin/handler.go | 7 - internal/admin/handler_guardrails_test.go | 4 +- internal/admin/handler_workflows_test.go | 11 +- internal/auditlog/auditlog.go | 14 - internal/batchrewrite/helpers.go | 62 -- internal/batchrewrite/helpers_test.go | 32 - internal/cache/modelcache/redis.go | 11 - internal/cache/modelcache/redis_test.go | 12 + internal/conversationstore/store_memory.go | 11 +- internal/core/batch_semantic.go | 53 -- internal/core/batch_semantic_test.go | 27 +- internal/core/chat_content.go | 14 - internal/core/chat_content_test.go | 10 - internal/core/semantic.go | 7 +- internal/core/semantic_canonical.go | 26 - internal/core/semantic_canonical_test.go | 31 - internal/embedding/embedding_test.go | 15 - internal/failover/resolver.go | 11 +- internal/failover/resolver_test.go | 7 + internal/gateway/batch_selection.go | 20 - internal/gateway/edge_cases_test.go | 4 +- internal/gateway/request_model_resolution.go | 5 - internal/guardrails/provider_harness_test.go | 7 +- internal/observability/metrics.go | 37 -- internal/observability/metrics_test.go | 36 - internal/providers/config.go | 17 - internal/providers/config_test.go | 19 + .../providers/file_adapter_openai_compat.go | 4 - .../file_adapter_openai_compat_test.go | 4 +- internal/responsecache/handle_request_test.go | 477 ------------- internal/responsecache/semantic.go | 24 +- internal/responsecache/semantic_test.go | 12 +- internal/responsecache/stream_cache.go | 164 ----- internal/responsecache/stream_cache_chat.go | 396 ----------- .../responsecache/stream_cache_responses.go | 628 ------------------ .../{vecstore_map.go => vecstore_map_test.go} | 0 internal/server/auth.go | 13 +- internal/server/handlers.go | 30 +- internal/server/http_test.go | 2 +- internal/server/model_validation.go | 34 +- internal/server/model_validation_test.go | 21 - internal/server/request_model_resolution.go | 4 - .../server/request_model_resolution_test.go | 12 +- internal/server/seams_test.go | 90 +++ .../server/translated_inference_service.go | 13 - internal/server/workflow_helpers.go | 11 - internal/server/workflow_helpers_test.go | 2 +- internal/server/workflow_policy_test.go | 4 +- internal/workflows/compiler.go | 5 - internal/workflows/compiler_test.go | 6 +- internal/workflows/service_test.go | 34 +- tests/e2e/setup_test.go | 2 +- 54 files changed, 252 insertions(+), 2310 deletions(-) delete mode 100644 internal/responsecache/stream_cache_chat.go delete mode 100644 internal/responsecache/stream_cache_responses.go rename internal/responsecache/{vecstore_map.go => vecstore_map_test.go} (100%) create mode 100644 internal/server/seams_test.go diff --git a/internal/admin/dashboard/dashboard.go b/internal/admin/dashboard/dashboard.go index 14546cc5b..7893fbc6e 100644 --- a/internal/admin/dashboard/dashboard.go +++ b/internal/admin/dashboard/dashboard.go @@ -28,12 +28,8 @@ type Handler struct { basePath string } -// New creates a new dashboard handler with parsed templates and static file server. -func New() (*Handler, error) { - return NewWithBasePath("/") -} - // NewWithBasePath creates a dashboard handler for an app mounted under basePath. +// It parses templates and sets up the static file server. func NewWithBasePath(basePath string) (*Handler, error) { basePath = config.NormalizeBasePath(basePath) assetVersions, err := buildFrontendAssetVersions() diff --git a/internal/admin/dashboard/dashboard_test.go b/internal/admin/dashboard/dashboard_test.go index 8c90808a5..710d00ec0 100644 --- a/internal/admin/dashboard/dashboard_test.go +++ b/internal/admin/dashboard/dashboard_test.go @@ -11,19 +11,19 @@ import ( ) func TestNew(t *testing.T) { - h, err := New() + h, err := NewWithBasePath("/") if err != nil { - t.Fatalf("New() returned error: %v", err) + t.Fatalf("NewWithBasePath() returned error: %v", err) } if h == nil { - t.Fatal("New() returned nil handler") + t.Fatalf("NewWithBasePath() returned nil handler") } } func TestIndex_ReturnsHTML(t *testing.T) { - h, err := New() + h, err := NewWithBasePath("/") if err != nil { - t.Fatalf("New() returned error: %v", err) + t.Fatalf("NewWithBasePath() returned error: %v", err) } e := echo.New() @@ -115,9 +115,9 @@ func TestIndex_UsesBasePathForGeneratedURLs(t *testing.T) { } func TestStatic_ServesCSS(t *testing.T) { - h, err := New() + h, err := NewWithBasePath("/") if err != nil { - t.Fatalf("New() returned error: %v", err) + t.Fatalf("NewWithBasePath() returned error: %v", err) } e := echo.New() @@ -138,9 +138,9 @@ func TestStatic_ServesCSS(t *testing.T) { } func TestStatic_ServesJS(t *testing.T) { - h, err := New() + h, err := NewWithBasePath("/") if err != nil { - t.Fatalf("New() returned error: %v", err) + t.Fatalf("NewWithBasePath() returned error: %v", err) } e := echo.New() @@ -161,9 +161,9 @@ func TestStatic_ServesJS(t *testing.T) { } func TestStatic_ServesModuleJS(t *testing.T) { - h, err := New() + h, err := NewWithBasePath("/") if err != nil { - t.Fatalf("New() returned error: %v", err) + t.Fatalf("NewWithBasePath() returned error: %v", err) } e := echo.New() @@ -184,9 +184,9 @@ func TestStatic_ServesModuleJS(t *testing.T) { } func TestStatic_ServesProvidersModuleJS(t *testing.T) { - h, err := New() + h, err := NewWithBasePath("/") if err != nil { - t.Fatalf("New() returned error: %v", err) + t.Fatalf("NewWithBasePath() returned error: %v", err) } e := echo.New() @@ -207,9 +207,9 @@ func TestStatic_ServesProvidersModuleJS(t *testing.T) { } func TestStatic_ServesVirtualModelsModuleJS(t *testing.T) { - h, err := New() + h, err := NewWithBasePath("/") if err != nil { - t.Fatalf("New() returned error: %v", err) + t.Fatalf("NewWithBasePath() returned error: %v", err) } e := echo.New() @@ -230,9 +230,9 @@ func TestStatic_ServesVirtualModelsModuleJS(t *testing.T) { } func TestStatic_ServesWorkflowsModuleJS(t *testing.T) { - h, err := New() + h, err := NewWithBasePath("/") if err != nil { - t.Fatalf("New() returned error: %v", err) + t.Fatalf("NewWithBasePath() returned error: %v", err) } e := echo.New() @@ -253,9 +253,9 @@ func TestStatic_ServesWorkflowsModuleJS(t *testing.T) { } func TestStatic_ServesGuardrailsModuleJS(t *testing.T) { - h, err := New() + h, err := NewWithBasePath("/") if err != nil { - t.Fatalf("New() returned error: %v", err) + t.Fatalf("NewWithBasePath() returned error: %v", err) } e := echo.New() @@ -276,9 +276,9 @@ func TestStatic_ServesGuardrailsModuleJS(t *testing.T) { } func TestStatic_ServesFavicon(t *testing.T) { - h, err := New() + h, err := NewWithBasePath("/") if err != nil { - t.Fatalf("New() returned error: %v", err) + t.Fatalf("NewWithBasePath() returned error: %v", err) } e := echo.New() @@ -299,9 +299,9 @@ func TestStatic_ServesFavicon(t *testing.T) { } func TestStatic_NotFound(t *testing.T) { - h, err := New() + h, err := NewWithBasePath("/") if err != nil { - t.Fatalf("New() returned error: %v", err) + t.Fatalf("NewWithBasePath() returned error: %v", err) } e := echo.New() @@ -322,9 +322,9 @@ func TestStatic_NotFound(t *testing.T) { // page must load every script, style, and font from the embedded /admin/static // tree, never from a CDN or remote font host. func TestIndex_HasNoExternalResources(t *testing.T) { - h, err := New() + h, err := NewWithBasePath("/") if err != nil { - t.Fatalf("New() returned error: %v", err) + t.Fatalf("NewWithBasePath() returned error: %v", err) } e := echo.New() @@ -358,9 +358,9 @@ func TestIndex_HasNoExternalResources(t *testing.T) { // TestStatic_ServesVendoredAssets confirms the vendored libraries and font // files are embedded and served, so the dashboard renders without network access. func TestStatic_ServesVendoredAssets(t *testing.T) { - h, err := New() + h, err := NewWithBasePath("/") if err != nil { - t.Fatalf("New() returned error: %v", err) + t.Fatalf("NewWithBasePath() returned error: %v", err) } paths := []string{ diff --git a/internal/admin/handler.go b/internal/admin/handler.go index bb8d99716..a5fb1fc5c 100644 --- a/internal/admin/handler.go +++ b/internal/admin/handler.go @@ -224,13 +224,6 @@ func WithTagging(service *tagging.Service) Option { } } -// WithGuardrailsRegistry enables listing valid guardrail references for workflow authoring. -func WithGuardrailsRegistry(registry guardrails.Catalog) Option { - return func(h *Handler) { - h.guardrails = registry - } -} - // WithGuardrailService enables full guardrail definition administration endpoints. func WithGuardrailService(service *guardrails.Service) Option { return func(h *Handler) { diff --git a/internal/admin/handler_guardrails_test.go b/internal/admin/handler_guardrails_test.go index 7bfdd57f0..e06a687c8 100644 --- a/internal/admin/handler_guardrails_test.go +++ b/internal/admin/handler_guardrails_test.go @@ -353,7 +353,7 @@ func TestDeleteGuardrailRejectsActiveWorkflowReference(t *testing.T) { }, }, } - planService, err := workflows.NewService(planStore, workflows.NewCompiler(guardrailService)) + planService, err := workflows.NewService(planStore, workflows.NewCompilerWithFeatureCaps(guardrailService, core.DefaultWorkflowFeatures())) if err != nil { t.Fatalf("workflows.NewService() error = %v", err) } @@ -408,7 +408,7 @@ func TestDeleteGuardrailIgnoresDisabledWorkflowGuardrailRefs(t *testing.T) { }, }, } - planService, err := workflows.NewService(planStore, workflows.NewCompiler(guardrailService)) + planService, err := workflows.NewService(planStore, workflows.NewCompilerWithFeatureCaps(guardrailService, core.DefaultWorkflowFeatures())) if err != nil { t.Fatalf("workflows.NewService() error = %v", err) } diff --git a/internal/admin/handler_workflows_test.go b/internal/admin/handler_workflows_test.go index d7d825979..a3de0f95e 100644 --- a/internal/admin/handler_workflows_test.go +++ b/internal/admin/handler_workflows_test.go @@ -18,6 +18,15 @@ import ( "gomodel/internal/workflows" ) +// WithGuardrailsRegistry enables listing valid guardrail references for +// workflow authoring. Test-only seam: production wires the full guardrail +// service via WithGuardrailService. +func WithGuardrailsRegistry(registry guardrails.Catalog) Option { + return func(h *Handler) { + h.guardrails = registry + } +} + type workflowTestStore struct { versions []workflows.Version } @@ -192,7 +201,7 @@ func newWorkflowHandler(t *testing.T, store workflows.Store, registry *guardrail func newWorkflowHandlerWithModelRegistry(t *testing.T, store workflows.Store, modelRegistry *providers.ModelRegistry, guardrailRegistry *guardrails.Registry) *Handler { t.Helper() - service, err := workflows.NewService(store, workflows.NewCompiler(guardrailRegistry)) + service, err := workflows.NewService(store, workflows.NewCompilerWithFeatureCaps(guardrailRegistry, core.DefaultWorkflowFeatures())) if err != nil { t.Fatalf("NewService() error = %v", err) } diff --git a/internal/auditlog/auditlog.go b/internal/auditlog/auditlog.go index 9bef72237..81076cfb9 100644 --- a/internal/auditlog/auditlog.go +++ b/internal/auditlog/auditlog.go @@ -369,17 +369,3 @@ type Config struct { // When true, only /v1/chat/completions, /v1/responses, /v1/embeddings, /v1/files, and /v1/batches are logged OnlyModelInteractions bool } - -// DefaultConfig returns a Config with sensible defaults -func DefaultConfig() Config { - return Config{ - Enabled: false, - LogBodies: false, - LogAudioBodies: false, - LogHeaders: false, - BufferSize: 1000, - FlushInterval: 5 * time.Second, - RetentionDays: 30, - OnlyModelInteractions: true, - } -} diff --git a/internal/batchrewrite/helpers.go b/internal/batchrewrite/helpers.go index ea09fe351..0a3440bbd 100644 --- a/internal/batchrewrite/helpers.go +++ b/internal/batchrewrite/helpers.go @@ -17,22 +17,6 @@ type FileDeleter interface { DeleteFile(ctx context.Context, providerType, id string) (*core.FileDeleteResponse, error) } -// FileRouter resolves a routed native file provider lazily. -type FileRouter func() (core.NativeFileRoutableProvider, error) - -// RecordPreparation stores request-scoped rewrite metadata for persistence and -// later cleanup. -func RecordPreparation(ctx context.Context, original, rewritten *core.BatchRequest) { - if ctx == nil || original == nil || rewritten == nil { - return - } - metadata := core.GetBatchPreparationMetadata(ctx) - if metadata == nil { - return - } - metadata.RecordInputFileRewrite(original.InputFileID, rewritten.InputFileID) -} - // RecordResult stores rewrite metadata produced by an explicit batch preparer. func RecordResult(ctx context.Context, result *core.BatchRewriteResult) { if ctx == nil || result == nil { @@ -65,52 +49,6 @@ func CleanupFile(ctx context.Context, files FileDeleter, providerType, fileID, l return true } -// CleanupFileFromRouter resolves the native file API only when there is a file -// id to delete. -func CleanupFileFromRouter(ctx context.Context, router FileRouter, providerType, fileID, logMessage string, attrs ...any) bool { - fileID = strings.TrimSpace(fileID) - if router == nil || fileID == "" { - return false - } - files, err := router() - if err != nil { - return false - } - return CleanupFile(ctx, files, providerType, fileID, logMessage, attrs...) -} - -// CleanupSupersededFileFromRouter deletes a local rewrite artifact only when a -// later rewrite has replaced it in the request-scoped batch metadata. -func CleanupSupersededFileFromRouter(ctx context.Context, router FileRouter, providerType, fileID, logMessage string, attrs ...any) bool { - if !ShouldCleanupSupersededFile(ctx, fileID) { - return false - } - return CleanupFileFromRouter(ctx, router, providerType, fileID, logMessage, attrs...) -} - -// CleanupSupersededFile deletes a local rewrite artifact only when a later -// rewrite has replaced it in the request-scoped batch metadata. -func CleanupSupersededFile(ctx context.Context, files FileDeleter, providerType, fileID, logMessage string, attrs ...any) bool { - if !ShouldCleanupSupersededFile(ctx, fileID) { - return false - } - return CleanupFile(ctx, files, providerType, fileID, logMessage, attrs...) -} - -// ShouldCleanupSupersededFile reports whether fileID is a temporary rewrite -// artifact that has been superseded by a later rewrite stage. -func ShouldCleanupSupersededFile(ctx context.Context, fileID string) bool { - fileID = strings.TrimSpace(fileID) - if fileID == "" { - return false - } - metadata := core.GetBatchPreparationMetadata(ctx) - if metadata == nil { - return false - } - return strings.TrimSpace(metadata.RewrittenInputFileID) != fileID -} - // MergeEndpointHints returns a fresh map containing left hints overwritten by // right hints. It preserves nil when both inputs are empty. func MergeEndpointHints(left, right map[string]string) map[string]string { diff --git a/internal/batchrewrite/helpers_test.go b/internal/batchrewrite/helpers_test.go index 6ef8e35cb..e844ce14b 100644 --- a/internal/batchrewrite/helpers_test.go +++ b/internal/batchrewrite/helpers_test.go @@ -27,20 +27,6 @@ func (d *recordingDeleter) DeleteFile(_ context.Context, providerType, id string return &core.FileDeleteResponse{ID: id, Deleted: true}, nil } -func TestRecordPreparation(t *testing.T) { - metadata := &core.BatchPreparationMetadata{} - ctx := core.WithBatchPreparationMetadata(context.Background(), metadata) - - RecordPreparation(ctx, &core.BatchRequest{InputFileID: " file_original "}, &core.BatchRequest{InputFileID: " file_rewritten "}) - - if metadata.OriginalInputFileID != "file_original" { - t.Fatalf("OriginalInputFileID = %q, want file_original", metadata.OriginalInputFileID) - } - if metadata.RewrittenInputFileID != "file_rewritten" { - t.Fatalf("RewrittenInputFileID = %q, want file_rewritten", metadata.RewrittenInputFileID) - } -} - func TestRecordResult(t *testing.T) { metadata := &core.BatchPreparationMetadata{} ctx := core.WithBatchPreparationMetadata(context.Background(), metadata) @@ -79,24 +65,6 @@ func TestCleanupFileReturnsFalseOnDeleteError(t *testing.T) { } } -func TestCleanupSupersededFile(t *testing.T) { - metadata := &core.BatchPreparationMetadata{RewrittenInputFileID: "file_current"} - ctx := core.WithBatchPreparationMetadata(context.Background(), metadata) - deleter := &recordingDeleter{} - - if CleanupSupersededFile(ctx, deleter, "openai", "file_current", "") { - t.Fatal("CleanupSupersededFile deleted current file") - } - if !CleanupSupersededFile(ctx, deleter, "openai", "file_old", "") { - t.Fatal("CleanupSupersededFile returned false for superseded file") - } - - want := []deleteCall{{providerType: "openai", fileID: "file_old"}} - if !reflect.DeepEqual(deleter.calls, want) { - t.Fatalf("calls = %#v, want %#v", deleter.calls, want) - } -} - func TestMergeEndpointHints(t *testing.T) { left := map[string]string{"a": "/v1/chat/completions", "b": "/v1/responses"} right := map[string]string{"b": "/v1/chat/completions", "c": "/v1/embeddings"} diff --git a/internal/cache/modelcache/redis.go b/internal/cache/modelcache/redis.go index 36a85f893..9d7deff7b 100644 --- a/internal/cache/modelcache/redis.go +++ b/internal/cache/modelcache/redis.go @@ -56,17 +56,6 @@ func NewRedisModelCache(cfg RedisModelCacheConfig) (Cache, error) { return &redisModelCache{store: store, key: key, ttl: ttl, owned: true}, nil } -// NewRedisModelCacheWithStore creates a Cache from an existing Store (for testing). -func NewRedisModelCacheWithStore(store cache.Store, key string, ttl time.Duration) Cache { - if key == "" { - key = DefaultRedisKey - } - if ttl == 0 { - ttl = cache.DefaultRedisTTL - } - return &redisModelCache{store: store, key: key, ttl: ttl, owned: false} -} - type redisModelCache struct { store cache.Store key string diff --git a/internal/cache/modelcache/redis_test.go b/internal/cache/modelcache/redis_test.go index b3212bee4..65d284202 100644 --- a/internal/cache/modelcache/redis_test.go +++ b/internal/cache/modelcache/redis_test.go @@ -8,6 +8,18 @@ import ( "gomodel/internal/cache" ) +// NewRedisModelCacheWithStore creates a Cache from an existing Store, letting +// tests exercise redisModelCache without a real Redis connection. +func NewRedisModelCacheWithStore(store cache.Store, key string, ttl time.Duration) Cache { + if key == "" { + key = DefaultRedisKey + } + if ttl == 0 { + ttl = cache.DefaultRedisTTL + } + return &redisModelCache{store: store, key: key, ttl: ttl, owned: false} +} + func TestRedisModelCache_GetSet(t *testing.T) { store := cache.NewMapStore() defer store.Close() diff --git a/internal/conversationstore/store_memory.go b/internal/conversationstore/store_memory.go index 8551edb25..2e25300d5 100644 --- a/internal/conversationstore/store_memory.go +++ b/internal/conversationstore/store_memory.go @@ -65,17 +65,8 @@ func WithMaxBytes(maxBytes int64) MemoryStoreOption { } } -// WithUnboundedRetention disables default in-memory retention bounds. -func WithUnboundedRetention() MemoryStoreOption { - return func(s *MemoryStore) { - s.ttl = 0 - s.maxEntries = 0 - s.maxBytes = 0 - } -} - // NewMemoryStore creates an empty in-memory conversation store. -// By default retention is bounded; pass WithUnboundedRetention to opt out. +// Retention is bounded by default; options can adjust or disable the bounds. func NewMemoryStore(options ...MemoryStoreOption) *MemoryStore { store := &MemoryStore{ items: make(map[string]*StoredConversation), diff --git a/internal/core/batch_semantic.go b/internal/core/batch_semantic.go index 13bffc0d9..25ecf8b82 100644 --- a/internal/core/batch_semantic.go +++ b/internal/core/batch_semantic.go @@ -26,50 +26,6 @@ type DecodedBatchItemHandlers[T any] struct { Default func(*DecodedBatchItemRequest) (T, error) } -// ChatRequest returns the decoded ChatRequest when the receiver is non-nil and -// the underlying Request is a *ChatRequest. It returns nil for a nil receiver -// or for non-chat batch items. -func (decoded *DecodedBatchItemRequest) ChatRequest() *ChatRequest { - if decoded == nil { - return nil - } - req, _ := decoded.Request.(*ChatRequest) - return req -} - -// ResponsesRequest returns the decoded ResponsesRequest when the receiver is -// non-nil and the underlying Request is a *ResponsesRequest. It returns nil for -// a nil receiver or for non-responses batch items. -func (decoded *DecodedBatchItemRequest) ResponsesRequest() *ResponsesRequest { - if decoded == nil { - return nil - } - req, _ := decoded.Request.(*ResponsesRequest) - return req -} - -// EmbeddingRequest returns the decoded EmbeddingRequest when the receiver is -// non-nil and the underlying Request is an *EmbeddingRequest. It returns nil -// for a nil receiver or for non-embedding batch items. -func (decoded *DecodedBatchItemRequest) EmbeddingRequest() *EmbeddingRequest { - if decoded == nil { - return nil - } - req, _ := decoded.Request.(*EmbeddingRequest) - return req -} - -// ModelSelector returns the selected model/provider pair for the decoded batch -// item. It returns an error when the receiver is nil, when the decoded request -// type is unsupported, or when the canonical selector cannot be parsed. -func (decoded *DecodedBatchItemRequest) ModelSelector() (ModelSelector, error) { - requested, err := decoded.RequestedModelSelector() - if err != nil { - return ModelSelector{}, err - } - return requested.Normalize() -} - // RequestedModelSelector returns the raw selector requested by the decoded batch // item, preserving whether the provider came from the explicit field. func (decoded *DecodedBatchItemRequest) RequestedModelSelector() (RequestedModelSelector, error) { @@ -210,15 +166,6 @@ func DecodeKnownBatchItemRequest(defaultEndpoint string, item BatchRequestItem) return decoded, nil } -// BatchItemModelSelector derives the model selector for a known JSON batch subrequest. -func BatchItemModelSelector(defaultEndpoint string, item BatchRequestItem) (ModelSelector, error) { - decoded, err := DecodeKnownBatchItemRequest(defaultEndpoint, item) - if err != nil { - return ModelSelector{}, err - } - return decoded.ModelSelector() -} - // BatchItemRequestedModelSelector derives the raw requested selector for a // known JSON batch subrequest. func BatchItemRequestedModelSelector(defaultEndpoint string, item BatchRequestItem) (RequestedModelSelector, error) { diff --git a/internal/core/batch_semantic_test.go b/internal/core/batch_semantic_test.go index 65e69eb63..af2ec780f 100644 --- a/internal/core/batch_semantic_test.go +++ b/internal/core/batch_semantic_test.go @@ -14,7 +14,7 @@ func TestNormalizeOperationPath(t *testing.T) { } } -func TestBatchItemModelSelector(t *testing.T) { +func TestBatchItemRequestedModelSelector(t *testing.T) { t.Parallel() tests := []struct { @@ -56,26 +56,30 @@ func TestBatchItemModelSelector(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - selector, err := BatchItemModelSelector(tt.defaultEndpoint, tt.item) + requested, err := BatchItemRequestedModelSelector(tt.defaultEndpoint, tt.item) if err != nil { - t.Fatalf("BatchItemModelSelector() error = %v", err) + t.Fatalf("BatchItemRequestedModelSelector() error = %v", err) + } + selector, err := requested.Normalize() + if err != nil { + t.Fatalf("Normalize() error = %v", err) } if got := selector.QualifiedModel(); got != tt.want { - t.Fatalf("BatchItemModelSelector() = %q, want %q", got, tt.want) + t.Fatalf("BatchItemRequestedModelSelector() = %q, want %q", got, tt.want) } }) } } -func TestBatchItemModelSelectorRejectsUnsupportedEndpoint(t *testing.T) { +func TestBatchItemRequestedModelSelectorRejectsUnsupportedEndpoint(t *testing.T) { t.Parallel() - _, err := BatchItemModelSelector("/v1/files", BatchRequestItem{ + _, err := BatchItemRequestedModelSelector("/v1/files", BatchRequestItem{ URL: "/v1/files", Body: json.RawMessage(`{"purpose":"batch"}`), }) if err == nil { - t.Fatal("BatchItemModelSelector() error = nil, want unsupported endpoint error") + t.Fatal("BatchItemRequestedModelSelector() error = nil, want unsupported endpoint error") } } @@ -95,11 +99,12 @@ func TestDecodeKnownBatchItemRequest_NormalizesFullURLAndDecodesCanonicalRequest if decoded.Operation != OperationResponses { t.Fatalf("Operation = %q, want responses", decoded.Operation) } - if decoded.ResponsesRequest() == nil { - t.Fatal("ResponsesRequest = nil") + req, ok := decoded.Request.(*ResponsesRequest) + if !ok || req == nil { + t.Fatalf("Request = %T, want *ResponsesRequest", decoded.Request) } - if decoded.ResponsesRequest().Model != "gpt-4o-mini" { - t.Fatalf("ResponsesRequest.Model = %q, want gpt-4o-mini", decoded.ResponsesRequest().Model) + if req.Model != "gpt-4o-mini" { + t.Fatalf("ResponsesRequest.Model = %q, want gpt-4o-mini", req.Model) } } diff --git a/internal/core/chat_content.go b/internal/core/chat_content.go index 4fc2c4740..048e662f6 100644 --- a/internal/core/chat_content.go +++ b/internal/core/chat_content.go @@ -326,20 +326,6 @@ func HasStructuredContent(content any) bool { } } -// HasNonTextContent reports whether the content contains image/audio parts. -func HasNonTextContent(content any) bool { - parts, ok := NormalizeContentParts(content) - if !ok { - return false - } - for _, part := range parts { - if part.Type != "text" { - return true - } - } - return false -} - // NormalizeContentParts converts dynamic JSON-decoded content into typed parts. func NormalizeContentParts(content any) ([]ContentPart, bool) { normalized, err := NormalizeMessageContent(content) diff --git a/internal/core/chat_content_test.go b/internal/core/chat_content_test.go index 03c7b9868..fccf49f83 100644 --- a/internal/core/chat_content_test.go +++ b/internal/core/chat_content_test.go @@ -410,13 +410,3 @@ func TestMessageUnmarshalJSON_MixedTextImageAudio(t *testing.T) { t.Fatalf("unexpected part 2: %+v", parts[2]) } } - -func TestHasNonTextContent_InputAudio(t *testing.T) { - result := HasNonTextContent([]ContentPart{{ - Type: "input_audio", - InputAudio: &InputAudioContent{Data: "abc", Format: "wav"}, - }}) - if !result { - t.Fatal("HasNonTextContent() = false, want true") - } -} diff --git a/internal/core/semantic.go b/internal/core/semantic.go index b6aad4862..22d368e0a 100644 --- a/internal/core/semantic.go +++ b/internal/core/semantic.go @@ -17,10 +17,11 @@ import ( // Lifecycle: // - DeriveWhiteBoxPrompt seeds these values directly from transport/body data. // - Canonical JSON decode may refine them from a cached request object. -// - NormalizeModelSelector canonicalizes model/provider values in place. +// - Selector normalization (ParseModelSelector / RequestedModelSelector.Normalize) +// canonicalizes model/provider values in place. // // Consumers that require canonical selector state should prefer a cached canonical -// request or call NormalizeModelSelector before relying on these fields. +// request or normalize the selector before relying on these fields. type RouteHints struct { Model string Provider string @@ -65,7 +66,7 @@ const ( // - transport seeds RouteType/OperationType plus sparse RouteHints // - route-specific metadata may be cached on demand // - canonical request decode may cache a parsed request and refine RouteHints -// - NormalizeModelSelector may rewrite selector hints into canonical form +// - selector normalization may rewrite selector hints into canonical form type WhiteBoxPrompt struct { RouteType string OperationType string diff --git a/internal/core/semantic_canonical.go b/internal/core/semantic_canonical.go index 7e1775b94..10a408308 100644 --- a/internal/core/semantic_canonical.go +++ b/internal/core/semantic_canonical.go @@ -172,32 +172,6 @@ func FileRouteMetadata(env *WhiteBoxPrompt, method, path string, routeParams map ) } -// NormalizeModelSelector canonicalizes model/provider selector inputs and keeps -// semantic selector hints aligned with the normalized request state. -// -// This is the point where RouteHints transition from raw ingress values -// (which may still contain a qualified model string like "openai/gpt-5-mini") -// to canonical model/provider fields. -func NormalizeModelSelector(env *WhiteBoxPrompt, model, provider *string) error { - if model == nil || provider == nil { - return NewInvalidRequestError("model selector targets are required", nil) - } - - selector, err := ParseModelSelector(*model, *provider) - if err != nil { - return NewInvalidRequestError(err.Error(), err) - } - - *model = selector.Model - *provider = selector.Provider - - if env != nil { - env.RouteHints.Model = selector.Model - env.RouteHints.Provider = selector.Provider - } - return nil -} - // DecodeCanonicalSelector decodes a canonical request body using the codec // resolved by canonicalOperationCodecFor for env, then extracts the model and // provider via semanticSelectorFromCanonicalRequest. It returns ok=false for a diff --git a/internal/core/semantic_canonical_test.go b/internal/core/semantic_canonical_test.go index 725057071..475bac79b 100644 --- a/internal/core/semantic_canonical_test.go +++ b/internal/core/semantic_canonical_test.go @@ -95,37 +95,6 @@ func TestFileRouteMetadata_CachesProviderHint(t *testing.T) { } } -func TestNormalizeModelSelector_UpdatesSemanticHints(t *testing.T) { - t.Parallel() - - env := &WhiteBoxPrompt{ - RouteHints: RouteHints{ - Model: "openai/gpt-4o-mini", - Provider: "", - }, - } - model := "openai/gpt-4o-mini" - provider := "" - - err := NormalizeModelSelector(env, &model, &provider) - if err != nil { - t.Fatalf("NormalizeModelSelector() error = %v", err) - } - - if model != "gpt-4o-mini" { - t.Fatalf("model = %q, want gpt-4o-mini", model) - } - if provider != "openai" { - t.Fatalf("provider = %q, want openai", provider) - } - if env.RouteHints.Model != "gpt-4o-mini" { - t.Fatalf("RouteHints.Model = %q, want gpt-4o-mini", env.RouteHints.Model) - } - if env.RouteHints.Provider != "openai" { - t.Fatalf("RouteHints.Provider = %q, want openai", env.RouteHints.Provider) - } -} - func TestDecodeCanonicalSelector_UsesOperationCodec(t *testing.T) { t.Parallel() diff --git a/internal/embedding/embedding_test.go b/internal/embedding/embedding_test.go index a4d388990..d75f7ae36 100644 --- a/internal/embedding/embedding_test.go +++ b/internal/embedding/embedding_test.go @@ -1,7 +1,6 @@ package embedding import ( - "context" "testing" "gomodel/config" @@ -145,17 +144,3 @@ func TestAPIEmbedder_UsesProviderCredentials(t *testing.T) { t.Errorf("endpointURL = %q, want %q", a.endpointURL, want) } } - -// MockEmbedder is an Embedder implementation for testing that returns a fixed vector. -type MockEmbedder struct { - Vector []float32 - Err error - Calls int -} - -func (m *MockEmbedder) Embed(_ context.Context, _ string) ([]float32, error) { - m.Calls++ - return m.Vector, m.Err -} - -func (m *MockEmbedder) Close() error { return nil } diff --git a/internal/failover/resolver.go b/internal/failover/resolver.go index 016add283..25e1b4ee1 100644 --- a/internal/failover/resolver.go +++ b/internal/failover/resolver.go @@ -42,14 +42,9 @@ type Resolver struct { registry Registry } -// NewResolver builds a failover resolver from config and the current model -// inventory. Returns nil when failover is effectively disabled. -func NewResolver(cfg config.FailoverConfig, registry Registry) *Resolver { - return NewResolverWithRuleProvider(cfg, registry, nil) -} - -// NewResolverWithRuleProvider builds a resolver backed by static config and an -// optional dynamic manual-rule provider. +// NewResolverWithRuleProvider builds a failover resolver from config and the +// current model inventory, backed by an optional dynamic manual-rule provider. +// Returns nil when failover is effectively disabled. func NewResolverWithRuleProvider(cfg config.FailoverConfig, registry Registry, ruleProvider RuleProvider) *Resolver { if registry == nil { return nil diff --git a/internal/failover/resolver_test.go b/internal/failover/resolver_test.go index b17eb664d..76c3989b9 100644 --- a/internal/failover/resolver_test.go +++ b/internal/failover/resolver_test.go @@ -8,6 +8,13 @@ import ( "gomodel/internal/providers" ) +// NewResolver builds a failover resolver from config and the current model +// inventory without a dynamic manual-rule provider. Test-only convenience +// over NewResolverWithRuleProvider. +func NewResolver(cfg config.FailoverConfig, registry Registry) *Resolver { + return NewResolverWithRuleProvider(cfg, registry, nil) +} + type fakeRegistry struct { byKey map[string]*providers.ModelInfo models []providers.ModelWithProvider diff --git a/internal/gateway/batch_selection.go b/internal/gateway/batch_selection.go index ec1095b16..eddb56c6b 100644 --- a/internal/gateway/batch_selection.go +++ b/internal/gateway/batch_selection.go @@ -20,26 +20,6 @@ type BatchInputFileProviderResolver interface { ResolveBatchInputFileProvider(ctx context.Context, fileID string) (providerType string, ok bool, err error) } -// DetermineBatchExecutionSelection resolves a native batch to one provider. -func DetermineBatchExecutionSelection( - provider core.RoutableProvider, - resolver ModelResolver, - req *core.BatchRequest, -) (BatchExecutionSelection, error) { - return DetermineBatchExecutionSelectionWithAuthorizer(context.Background(), provider, resolver, nil, req) -} - -// DetermineBatchExecutionSelectionWithAuthorizer resolves and authorizes native batch items. -func DetermineBatchExecutionSelectionWithAuthorizer( - ctx context.Context, - provider core.RoutableProvider, - resolver ModelResolver, - authorizer ModelAuthorizer, - req *core.BatchRequest, -) (BatchExecutionSelection, error) { - return DetermineBatchExecutionSelectionWithAuthorizerAndInputFileResolver(ctx, provider, resolver, authorizer, nil, req) -} - // DetermineBatchExecutionSelectionWithAuthorizerAndInputFileResolver resolves // and authorizes native batch items, using file ownership metadata for // file-backed batches when no explicit provider hint is supplied. diff --git a/internal/gateway/edge_cases_test.go b/internal/gateway/edge_cases_test.go index 8033b4642..b657bad5f 100644 --- a/internal/gateway/edge_cases_test.go +++ b/internal/gateway/edge_cases_test.go @@ -46,9 +46,9 @@ func TestMergeStoredBatchFromUpstreamPreservesGatewayOwnedMetadata(t *testing.T) } func TestDetermineBatchExecutionSelectionRejectsNilRequest(t *testing.T) { - _, err := DetermineBatchExecutionSelectionWithAuthorizer(context.Background(), nil, nil, nil, nil) + _, err := DetermineBatchExecutionSelectionWithAuthorizerAndInputFileResolver(context.Background(), nil, nil, nil, nil, nil) if err == nil { - t.Fatal("DetermineBatchExecutionSelectionWithAuthorizer() error = nil, want error") + t.Fatal("DetermineBatchExecutionSelectionWithAuthorizerAndInputFileResolver() error = nil, want error") } var gatewayErr *core.GatewayError diff --git a/internal/gateway/request_model_resolution.go b/internal/gateway/request_model_resolution.go index 715ec6fe4..449c0fd72 100644 --- a/internal/gateway/request_model_resolution.go +++ b/internal/gateway/request_model_resolution.go @@ -56,11 +56,6 @@ func WorkflowProviderNameForType(provider core.RoutableProvider, providerType st return "" } -// ResolveRequestModel resolves a requested selector into a concrete provider/model selector. -func ResolveRequestModel(provider core.RoutableProvider, resolver ModelResolver, requested core.RequestedModelSelector) (*core.RequestModelResolution, error) { - return ResolveRequestModelWithAuthorizer(context.Background(), provider, resolver, nil, requested) -} - // ResolveRequestModelWithAuthorizer resolves and validates a requested selector. func ResolveRequestModelWithAuthorizer( ctx context.Context, diff --git a/internal/guardrails/provider_harness_test.go b/internal/guardrails/provider_harness_test.go index 8e2dc39e8..ecf9f6f0e 100644 --- a/internal/guardrails/provider_harness_test.go +++ b/internal/guardrails/provider_harness_test.go @@ -107,13 +107,14 @@ func (g *GuardedProvider) CreateBatch(ctx context.Context, providerType string, if err != nil { return nil, err } - batchrewrite.RecordPreparation(ctx, req, result.Request) + batchrewrite.RecordResult(ctx, result) resp, err := bp.CreateBatch(ctx, providerType, result.Request) if err != nil { - batchrewrite.CleanupFileFromRouter(ctx, g.nativeFileRouter, providerType, result.RewrittenInputFileID, "") + if files, routeErr := g.nativeFileRouter(); routeErr == nil { + batchrewrite.CleanupFile(ctx, files, providerType, result.RewrittenInputFileID, "") + } return nil, err } - batchrewrite.CleanupSupersededFileFromRouter(ctx, g.nativeFileRouter, providerType, result.RewrittenInputFileID, "") return resp, nil } diff --git a/internal/observability/metrics.go b/internal/observability/metrics.go index df1477539..2d17a9679 100644 --- a/internal/observability/metrics.go +++ b/internal/observability/metrics.go @@ -3,7 +3,6 @@ package observability import ( "context" - "fmt" "strconv" "github.com/prometheus/client_golang/prometheus" @@ -175,26 +174,6 @@ func NewPrometheusHooks() llmclient.Hooks { // Panel 5: Requests by Model // Query: sum(rate(gomodel_requests_total[5m])) by (model) -// PrometheusMetrics provides access to all registered metrics for testing -type PrometheusMetrics struct { - RequestsTotal *prometheus.CounterVec - RequestDuration *prometheus.HistogramVec - InFlightRequests *prometheus.GaugeVec - ResponseSnapshotStoreFailures *prometheus.CounterVec - CircuitBreakerState *prometheus.GaugeVec -} - -// GetMetrics returns the prometheus metrics for testing and introspection -func GetMetrics() *PrometheusMetrics { - return &PrometheusMetrics{ - RequestsTotal: RequestsTotal, - RequestDuration: RequestDuration, - InFlightRequests: InFlightRequests, - ResponseSnapshotStoreFailures: ResponseSnapshotStoreFailures, - CircuitBreakerState: CircuitBreakerState, - } -} - // ResetMetrics resets all metrics to zero (useful for testing) func ResetMetrics() { RequestsTotal.Reset() @@ -203,19 +182,3 @@ func ResetMetrics() { ResponseSnapshotStoreFailures.Reset() CircuitBreakerState.Reset() } - -// HealthCheck verifies that metrics are being collected -func HealthCheck() error { - // Try to collect metrics - mfs, err := prometheus.DefaultGatherer.Gather() - if err != nil { - return fmt.Errorf("failed to gather metrics: %w", err) - } - - // Check that we have some metrics - if len(mfs) == 0 { - return fmt.Errorf("no metrics registered") - } - - return nil -} diff --git a/internal/observability/metrics_test.go b/internal/observability/metrics_test.go index ba02eab8e..2baad8dba 100644 --- a/internal/observability/metrics_test.go +++ b/internal/observability/metrics_test.go @@ -386,39 +386,3 @@ func TestRequestDuration(t *testing.T) { t.Fatal("Expected histogram, got nil") } } - -func TestHealthCheck(t *testing.T) { - // Reset metrics before test - ResetMetrics() - - // Health check should succeed - err := HealthCheck() - if err != nil { - t.Errorf("HealthCheck failed: %v", err) - } -} - -func TestGetMetrics(t *testing.T) { - metrics := GetMetrics() - - if metrics == nil { - t.Fatal("GetMetrics returned nil") - return - } - - if metrics.RequestsTotal == nil { - t.Error("RequestsTotal metric is nil") - } - - if metrics.RequestDuration == nil { - t.Error("RequestDuration metric is nil") - } - - if metrics.InFlightRequests == nil { - t.Error("InFlightRequests metric is nil") - } - - if metrics.ResponseSnapshotStoreFailures == nil { - t.Error("ResponseSnapshotStoreFailures metric is nil") - } -} diff --git a/internal/providers/config.go b/internal/providers/config.go index dd2a5bfc3..ac23d3597 100644 --- a/internal/providers/config.go +++ b/internal/providers/config.go @@ -468,23 +468,6 @@ func rawProviderMatchesType(cfg config.RawProviderConfig, providerType string) b return strings.TrimSpace(cfg.Type) == strings.TrimSpace(providerType) } -type providerEnvNames struct { - APIKey string - BaseURL string - APIVersion string - Models string -} - -func derivedEnvNames(providerType string) providerEnvNames { - prefix := envPrefix(providerType) - return providerEnvNames{ - APIKey: prefix + "_API_KEY", - BaseURL: prefix + "_BASE_URL", - APIVersion: prefix + "_API_VERSION", - Models: prefix + "_MODELS", - } -} - func envPrefix(providerType string) string { var b strings.Builder b.Grow(len(providerType)) diff --git a/internal/providers/config_test.go b/internal/providers/config_test.go index 74ecc01dd..c9388887a 100644 --- a/internal/providers/config_test.go +++ b/internal/providers/config_test.go @@ -1243,6 +1243,25 @@ func TestApplyProviderEnvVars_SuffixedEnvOverlaysMatchingYAMLProvider(t *testing } } +// providerEnvNames mirrors the env-var naming convention applied by +// applyProviderEnvVars, so tests can clear ambient variables per provider. +type providerEnvNames struct { + APIKey string + BaseURL string + APIVersion string + Models string +} + +func derivedEnvNames(providerType string) providerEnvNames { + prefix := envPrefix(providerType) + return providerEnvNames{ + APIKey: prefix + "_API_KEY", + BaseURL: prefix + "_BASE_URL", + APIVersion: prefix + "_API_VERSION", + Models: prefix + "_MODELS", + } +} + func TestApplyProviderEnvVars_SkipsWhenNoEnvVars(t *testing.T) { // Ensure no ambient env vars interfere for providerType, spec := range testDiscoveryConfigs { diff --git a/internal/providers/file_adapter_openai_compat.go b/internal/providers/file_adapter_openai_compat.go index 9662fcea1..96bbdfa5e 100644 --- a/internal/providers/file_adapter_openai_compat.go +++ b/internal/providers/file_adapter_openai_compat.go @@ -35,10 +35,6 @@ func prepareOpenAICompatibleRequest(prepare openAICompatibleRequestPreparer, req return prepare(req) } -func doOpenAICompatibleFileIDRequest[T any](ctx context.Context, client *llmclient.Client, method, id string, defaultObject string) (*T, error) { - return doOpenAICompatibleFileIDRequestWithPreparer[T](ctx, client, method, id, defaultObject, nil) -} - func doOpenAICompatibleFileIDRequestWithPreparer[T any](ctx context.Context, client *llmclient.Client, method, id string, defaultObject string, prepare openAICompatibleRequestPreparer) (*T, error) { trimmedID, err := validatedOpenAICompatibleFileID(client, id) if err != nil { diff --git a/internal/providers/file_adapter_openai_compat_test.go b/internal/providers/file_adapter_openai_compat_test.go index 8f736406f..b9b4a8b3b 100644 --- a/internal/providers/file_adapter_openai_compat_test.go +++ b/internal/providers/file_adapter_openai_compat_test.go @@ -168,10 +168,10 @@ func TestDoOpenAICompatibleFileIDRequest(t *testing.T) { switch tt.method { case http.MethodDelete: - resp, err := doOpenAICompatibleFileIDRequest[core.FileDeleteResponse](context.Background(), client, tt.method, tt.id, tt.defaultObject) + resp, err := doOpenAICompatibleFileIDRequestWithPreparer[core.FileDeleteResponse](context.Background(), client, tt.method, tt.id, tt.defaultObject, nil) tt.check(t, gotPath, nil, resp, err) default: - resp, err := doOpenAICompatibleFileIDRequest[core.FileObject](context.Background(), client, tt.method, tt.id, tt.defaultObject) + resp, err := doOpenAICompatibleFileIDRequestWithPreparer[core.FileObject](context.Background(), client, tt.method, tt.id, tt.defaultObject, nil) tt.check(t, gotPath, resp, nil, err) } if gotMethod != tt.method { diff --git a/internal/responsecache/handle_request_test.go b/internal/responsecache/handle_request_test.go index 0b44a9d24..5324f4b49 100644 --- a/internal/responsecache/handle_request_test.go +++ b/internal/responsecache/handle_request_test.go @@ -3,7 +3,6 @@ package responsecache import ( "bytes" "context" - "encoding/json" "errors" "net/http" "net/http/httptest" @@ -1083,479 +1082,3 @@ func TestHandleRequest_InvalidStreamingBodySkipsExactCacheWrite(t *testing.T) { t.Fatalf("expected invalid stream to bypass cache on follow-up, got %d calls", handlerCalls) } } - -func TestReconstructStreamingResponse_PreservesChatReasoningContent(t *testing.T) { - raw := []byte( - "data: {\"id\":\"chatcmpl-reasoning\",\"object\":\"chat.completion.chunk\",\"created\":1234567890,\"model\":\"claude-sonnet\",\"provider\":\"anthropic\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"reasoning_content\":\"think first\"},\"finish_reason\":null}]}\n\n" + - "data: {\"id\":\"chatcmpl-reasoning\",\"object\":\"chat.completion.chunk\",\"created\":1234567890,\"model\":\"claude-sonnet\",\"provider\":\"anthropic\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"final answer\"},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":10,\"completion_tokens\":4,\"total_tokens\":14}}\n\n" + - "data: [DONE]\n\n", - ) - - cached, ok := reconstructStreamingResponse("/v1/chat/completions", raw, streamResponseDefaults{ - Model: "claude-sonnet", - Provider: "anthropic", - }) - if !ok { - t.Fatal("expected streamed chat response to reconstruct successfully") - } - if !bytes.Contains(cached, []byte(`"reasoning_content":"think first"`)) { - t.Fatalf("reconstructed chat response = %q, want reasoning_content preserved", string(cached)) - } - - replay, err := renderCachedChatStream([]byte(`{"model":"claude-sonnet","stream":true}`), cached) - if err != nil { - t.Fatalf("renderCachedChatStream() error = %v", err) - } - if !bytes.Contains(replay, []byte(`"reasoning_content":"think first"`)) { - t.Fatalf("cached chat replay = %q, want reasoning_content delta", string(replay)) - } - if bytes.Contains(replay, []byte(`"usage"`)) { - t.Fatalf("cached chat replay without include_usage = %q, did not expect usage chunk", string(replay)) - } - - replayWithUsage, err := renderCachedChatStream([]byte(`{"model":"claude-sonnet","stream":true,"stream_options":{"include_usage":true}}`), cached) - if err != nil { - t.Fatalf("renderCachedChatStream(include_usage) error = %v", err) - } - if !bytes.Contains(replayWithUsage, []byte(`"usage"`)) { - t.Fatalf("cached chat replay with include_usage = %q, want usage chunk", string(replayWithUsage)) - } -} - -func TestRenderCachedChatStream_EmitsStandaloneUsageChunk(t *testing.T) { - raw := []byte( - "data: {\"id\":\"chatcmpl-usage\",\"object\":\"chat.completion.chunk\",\"created\":1234567890,\"model\":\"gpt-4o-mini\",\"provider\":\"openai\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"Hello\"},\"finish_reason\":null}]}\n\n" + - "data: {\"id\":\"chatcmpl-usage\",\"object\":\"chat.completion.chunk\",\"created\":1234567890,\"model\":\"gpt-4o-mini\",\"provider\":\"openai\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" world\"},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":11,\"completion_tokens\":2,\"total_tokens\":13}}\n\n" + - "data: [DONE]\n\n", - ) - - cached, ok := reconstructStreamingResponse("/v1/chat/completions", raw, streamResponseDefaults{ - Model: "gpt-4o-mini", - Provider: "openai", - }) - if !ok { - t.Fatal("expected streamed chat response to reconstruct successfully") - } - - replay, err := renderCachedChatStream([]byte(`{"model":"gpt-4o-mini","stream":true,"stream_options":{"include_usage":true}}`), cached) - if err != nil { - t.Fatalf("renderCachedChatStream() error = %v", err) - } - - var events []map[string]any - parseSSEJSONEvents(replay, func(event map[string]any) { - events = append(events, event) - }) - if len(events) != 2 { - t.Fatalf("len(events) = %d, want 2 chat events before [DONE]", len(events)) - } - - firstChoices, ok := events[0]["choices"].([]any) - if !ok || len(firstChoices) != 1 { - t.Fatalf("first event choices = %#v, want len=1", events[0]["choices"]) - } - if _, ok := events[0]["usage"]; ok { - t.Fatalf("first event should not carry usage, got %#v", events[0]["usage"]) - } - - secondChoices, ok := events[1]["choices"].([]any) - if !ok || len(secondChoices) != 0 { - t.Fatalf("usage event choices = %#v, want empty slice", events[1]["choices"]) - } - usage, ok := events[1]["usage"].(map[string]any) - if !ok { - t.Fatalf("usage event usage = %#v, want object", events[1]["usage"]) - } - if got, ok := jsonNumberToInt(usage["total_tokens"]); !ok || got != 13 { - t.Fatalf("usage.total_tokens = %#v, want 13", usage["total_tokens"]) - } -} - -func TestReconstructStreamingResponse_PreservesChatLogprobs(t *testing.T) { - raw := []byte( - "data: {\"id\":\"chatcmpl-logprobs\",\"object\":\"chat.completion.chunk\",\"created\":1234567890,\"model\":\"gpt-4o-mini\",\"provider\":\"openai\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"Hello\"},\"logprobs\":null,\"finish_reason\":null}]}\n\n" + - "data: {\"id\":\"chatcmpl-logprobs\",\"object\":\"chat.completion.chunk\",\"created\":1234567890,\"model\":\"gpt-4o-mini\",\"provider\":\"openai\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" world\"},\"logprobs\":null,\"finish_reason\":\"stop\"}]}\n\n" + - "data: [DONE]\n\n", - ) - - cached, ok := reconstructStreamingResponse("/v1/chat/completions", raw, streamResponseDefaults{ - Model: "gpt-4o-mini", - Provider: "openai", - }) - if !ok { - t.Fatal("expected streamed chat response to reconstruct successfully") - } - if !bytes.Contains(cached, []byte(`"logprobs":null`)) { - t.Fatalf("reconstructed chat response = %q, want choice.logprobs preserved", string(cached)) - } - - replay, err := renderCachedChatStream([]byte(`{"model":"gpt-4o-mini","stream":true}`), cached) - if err != nil { - t.Fatalf("renderCachedChatStream() error = %v", err) - } - if !bytes.Contains(replay, []byte(`"logprobs":null`)) { - t.Fatalf("cached chat replay = %q, want choice.logprobs preserved", string(replay)) - } -} - -func TestReconstructStreamingResponse_PreservesResponsesReasoningText(t *testing.T) { - raw := []byte( - "event: response.created\n" + - "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_reasoning_build\",\"object\":\"response\",\"created_at\":1234567890,\"model\":\"grok-4\",\"provider\":\"xai\",\"status\":\"in_progress\",\"output\":[]}}\n\n" + - "event: response.output_item.added\n" + - "data: {\"type\":\"response.output_item.added\",\"output_index\":0,\"item\":{\"id\":\"rs_1\",\"type\":\"reasoning\",\"status\":\"in_progress\",\"summary\":[]}}\n\n" + - "event: response.reasoning_text.delta\n" + - "data: {\"type\":\"response.reasoning_text.delta\",\"item_id\":\"rs_1\",\"output_index\":0,\"content_index\":0,\"delta\":\"step by\"}\n\n" + - "event: response.reasoning_text.delta\n" + - "data: {\"type\":\"response.reasoning_text.delta\",\"item_id\":\"rs_1\",\"output_index\":0,\"content_index\":1,\"delta\":\"step\"}\n\n" + - "event: response.output_item.done\n" + - "data: {\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{\"id\":\"rs_1\",\"type\":\"reasoning\",\"status\":\"completed\",\"summary\":[]}}\n\n" + - "event: response.completed\n" + - "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_reasoning_build\",\"object\":\"response\",\"created_at\":1234567890,\"model\":\"grok-4\",\"provider\":\"xai\",\"status\":\"completed\",\"output\":[{\"id\":\"rs_1\",\"type\":\"reasoning\",\"status\":\"completed\",\"summary\":[]}]}}\n\n" + - "data: [DONE]\n\n", - ) - - cached, ok := reconstructStreamingResponse("/v1/responses", raw, streamResponseDefaults{ - Model: "grok-4", - Provider: "xai", - }) - if !ok { - t.Fatal("expected streamed responses payload to reconstruct successfully") - } - - var response map[string]any - if err := json.Unmarshal(cached, &response); err != nil { - t.Fatalf("json.Unmarshal(cached) error = %v", err) - } - output, ok := response["output"].([]any) - if !ok || len(output) != 1 { - t.Fatalf("reconstructed output = %#v, want len=1", response["output"]) - } - item, ok := output[0].(map[string]any) - if !ok { - t.Fatalf("reconstructed output[0] = %#v, want object", output[0]) - } - if _, ok := item["content"]; ok { - t.Fatalf("reasoning item should not use content field, got %#v", item["content"]) - } - summary, ok := item["summary"].([]any) - if !ok || len(summary) != 2 { - t.Fatalf("reconstructed reasoning summary = %#v, want len=2", item["summary"]) - } - wantTexts := []string{"step by", "step"} - for i, wantText := range wantTexts { - part, ok := summary[i].(map[string]any) - if !ok { - t.Fatalf("reconstructed summary[%d] = %#v, want object", i, summary[i]) - } - if got, _ := part["type"].(string); got != "reasoning_text" { - t.Fatalf("reconstructed summary part type = %q, want reasoning_text", got) - } - if got, _ := part["text"].(string); got != wantText { - t.Fatalf("reconstructed summary part text = %q, want %q", got, wantText) - } - } - - replay, err := renderCachedResponsesStream([]byte(`{"model":"grok-4","stream":true}`), cached) - if err != nil { - t.Fatalf("renderCachedResponsesStream() error = %v", err) - } - var reasoningDeltas []map[string]any - parseSSEJSONEvents(replay, func(event map[string]any) { - if eventType, _ := event["type"].(string); eventType == "response.reasoning_text.delta" { - reasoningDeltas = append(reasoningDeltas, event) - } - }) - if len(reasoningDeltas) != 2 { - t.Fatalf("cached responses replay = %q, want 2 reasoning_text delta events", string(replay)) - } - for i, deltaEvent := range reasoningDeltas { - if got, _ := deltaEvent["delta"].(string); got != wantTexts[i] { - t.Fatalf("reasoning delta text = %q, want %q", got, wantTexts[i]) - } - if got, _ := deltaEvent["item_id"].(string); got != "rs_1" { - t.Fatalf("reasoning delta item_id = %q, want rs_1", got) - } - if got, ok := jsonNumberToInt(deltaEvent["output_index"]); !ok || got != 0 { - t.Fatalf("reasoning delta output_index = %#v, want 0", deltaEvent["output_index"]) - } - if got, ok := jsonNumberToInt(deltaEvent["content_index"]); !ok || got != i { - t.Fatalf("reasoning delta content_index = %#v, want %d", deltaEvent["content_index"], i) - } - } -} - -func TestReconstructStreamingResponse_HonorsResponsesTextDeltaLocators(t *testing.T) { - raw := []byte( - "event: response.created\n" + - "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_text_locator\",\"object\":\"response\",\"created_at\":1234567890,\"model\":\"gpt-4o-mini\",\"provider\":\"openai\",\"status\":\"in_progress\",\"output\":[{\"id\":\"rs_1\",\"type\":\"reasoning\",\"status\":\"in_progress\",\"summary\":[]}]}}\n\n" + - "event: response.output_text.delta\n" + - "data: {\"type\":\"response.output_text.delta\",\"item_id\":\"msg_1\",\"output_index\":1,\"content_index\":0,\"delta\":\"final\"}\n\n" + - "event: response.output_text.delta\n" + - "data: {\"type\":\"response.output_text.delta\",\"item_id\":\"msg_1\",\"output_index\":1,\"content_index\":1,\"delta\":\"answer\"}\n\n" + - "event: response.completed\n" + - "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_text_locator\",\"object\":\"response\",\"created_at\":1234567890,\"model\":\"gpt-4o-mini\",\"provider\":\"openai\",\"status\":\"completed\",\"output\":[{\"id\":\"rs_1\",\"type\":\"reasoning\",\"status\":\"completed\",\"summary\":[{\"type\":\"reasoning_text\",\"text\":\"step by step\"}]},{\"id\":\"msg_1\",\"type\":\"message\",\"role\":\"assistant\",\"status\":\"completed\",\"content\":[]}]}}\n\n" + - "data: [DONE]\n\n", - ) - - cached, ok := reconstructStreamingResponse("/v1/responses", raw, streamResponseDefaults{ - Model: "gpt-4o-mini", - Provider: "openai", - }) - if !ok { - t.Fatal("expected streamed responses payload to reconstruct successfully") - } - - var response map[string]any - if err := json.Unmarshal(cached, &response); err != nil { - t.Fatalf("json.Unmarshal(cached) error = %v", err) - } - output, ok := response["output"].([]any) - if !ok || len(output) != 2 { - t.Fatalf("reconstructed output = %#v, want len=2", response["output"]) - } - reasoningItem, ok := output[0].(map[string]any) - if !ok { - t.Fatalf("reconstructed output[0] = %#v, want object", output[0]) - } - if _, ok := reasoningItem["content"]; ok { - t.Fatalf("reasoning item should not use content field, got %#v", reasoningItem["content"]) - } - messageItem, ok := output[1].(map[string]any) - if !ok { - t.Fatalf("reconstructed output[1] = %#v, want object", output[1]) - } - content, ok := messageItem["content"].([]any) - if !ok || len(content) != 2 { - t.Fatalf("reconstructed message content = %#v, want len=2", messageItem["content"]) - } - wantTexts := []string{"final", "answer"} - for i, wantText := range wantTexts { - messagePart, ok := content[i].(map[string]any) - if !ok { - t.Fatalf("reconstructed message content[%d] = %#v, want object", i, content[i]) - } - if got, _ := messagePart["text"].(string); got != wantText { - t.Fatalf("reconstructed message text = %q, want %q", got, wantText) - } - } - - replay, err := renderCachedResponsesStream([]byte(`{"model":"gpt-4o-mini","stream":true}`), cached) - if err != nil { - t.Fatalf("renderCachedResponsesStream() error = %v", err) - } - var textDeltas []map[string]any - parseSSEJSONEvents(replay, func(event map[string]any) { - if eventType, _ := event["type"].(string); eventType == "response.output_text.delta" { - textDeltas = append(textDeltas, event) - } - }) - if len(textDeltas) != 2 { - t.Fatalf("cached responses replay = %q, want 2 output_text delta events", string(replay)) - } - for i, deltaEvent := range textDeltas { - if got, _ := deltaEvent["delta"].(string); got != wantTexts[i] { - t.Fatalf("text delta text = %q, want %q", got, wantTexts[i]) - } - if got, _ := deltaEvent["item_id"].(string); got != "msg_1" { - t.Fatalf("text delta item_id = %q, want msg_1", got) - } - if got, ok := jsonNumberToInt(deltaEvent["output_index"]); !ok || got != 1 { - t.Fatalf("text delta output_index = %#v, want 1", deltaEvent["output_index"]) - } - if got, ok := jsonNumberToInt(deltaEvent["content_index"]); !ok || got != i { - t.Fatalf("text delta content_index = %#v, want %d", deltaEvent["content_index"], i) - } - } -} - -func TestRenderCachedResponsesStream_PreservesReasoningTextDeltas(t *testing.T) { - cached := []byte(`{ - "id":"resp_reasoning", - "object":"response", - "created_at":1234567890, - "model":"grok-4", - "provider":"xai", - "status":"completed", - "output":[ - { - "id":"rs_1", - "type":"reasoning", - "status":"completed", - "summary":[{"type":"reasoning_text","text":"step by step"}] - }, - { - "id":"msg_1", - "type":"message", - "role":"assistant", - "status":"completed", - "content":[{"type":"output_text","text":"final answer"}] - } - ] - }`) - - replay, err := renderCachedResponsesStream([]byte(`{"model":"grok-4","stream":true}`), cached) - if err != nil { - t.Fatalf("renderCachedResponsesStream() error = %v", err) - } - if !bytes.Contains(replay, []byte("event: response.reasoning_text.delta")) { - t.Fatalf("cached responses replay = %q, want reasoning_text delta event", string(replay)) - } - if !bytes.Contains(replay, []byte("step by step")) { - t.Fatalf("cached responses replay = %q, want reasoning delta text", string(replay)) - } - if !bytes.Contains(replay, []byte("event: response.output_text.delta")) { - t.Fatalf("cached responses replay = %q, want output_text delta event", string(replay)) - } -} - -func TestRenderCachedResponsesStream_FunctionCallAddedItemOmitsArguments(t *testing.T) { - cached := []byte(`{ - "id":"resp_function_call", - "object":"response", - "created_at":1234567890, - "model":"gpt-4o-mini", - "provider":"openai", - "status":"completed", - "output":[ - { - "id":"fc_1", - "type":"function_call", - "status":"completed", - "call_id":"call_1", - "name":"lookup_weather", - "arguments":"{\"city\":\"Warsaw\"}" - } - ] - }`) - - replay, err := renderCachedResponsesStream([]byte(`{"model":"gpt-4o-mini","stream":true}`), cached) - if err != nil { - t.Fatalf("renderCachedResponsesStream() error = %v", err) - } - - var addedItem map[string]any - var argDelta map[string]any - var argDone map[string]any - parseSSEJSONEvents(replay, func(event map[string]any) { - switch eventType, _ := event["type"].(string); eventType { - case "response.output_item.added": - addedItem, _ = event["item"].(map[string]any) - case "response.function_call_arguments.delta": - argDelta = event - case "response.function_call_arguments.done": - argDone = event - } - }) - - if addedItem == nil { - t.Fatalf("cached responses replay = %q, want output_item.added event", string(replay)) - } - if _, ok := addedItem["arguments"]; ok { - t.Fatalf("added item arguments = %#v, want omitted", addedItem["arguments"]) - } - if argDelta == nil || argDone == nil { - t.Fatalf("cached responses replay = %q, want function_call_arguments delta and done events", string(replay)) - } - if got, _ := argDelta["delta"].(string); got != `{"city":"Warsaw"}` { - t.Fatalf("arguments delta = %q, want full arguments", got) - } - if got, _ := argDone["arguments"].(string); got != `{"city":"Warsaw"}` { - t.Fatalf("arguments done = %q, want full arguments", got) - } -} - -func TestReconstructStreamingResponse_PreservesResponsesTerminalEvents(t *testing.T) { - tests := []struct { - name string - eventName string - status string - terminalResponse string - assertTerminal func(*testing.T, map[string]any) - }{ - { - name: "failed", - eventName: "response.failed", - status: "failed", - terminalResponse: `{"id":"resp_failed","object":"response","created_at":1234567890,"model":"gpt-4o-mini","provider":"openai","status":"failed","error":{"code":"boom","message":"upstream failed"},"metadata":{"trace":"abc"},"output":[]}`, - assertTerminal: func(t *testing.T, response map[string]any) { - t.Helper() - errMap, ok := response["error"].(map[string]any) - if !ok { - t.Fatalf("terminal response error = %#v, want object", response["error"]) - } - if got, _ := errMap["code"].(string); got != "boom" { - t.Fatalf("terminal response error.code = %q, want boom", got) - } - }, - }, - { - name: "incomplete", - eventName: "response.incomplete", - status: "incomplete", - terminalResponse: `{"id":"resp_incomplete","object":"response","created_at":1234567890,"model":"gpt-4o-mini","provider":"openai","status":"incomplete","incomplete_details":{"reason":"max_output_tokens"},"metadata":{"trace":"def"},"output":[]}`, - assertTerminal: func(t *testing.T, response map[string]any) { - t.Helper() - details, ok := response["incomplete_details"].(map[string]any) - if !ok { - t.Fatalf("terminal response incomplete_details = %#v, want object", response["incomplete_details"]) - } - if got, _ := details["reason"].(string); got != "max_output_tokens" { - t.Fatalf("terminal response incomplete_details.reason = %q, want max_output_tokens", got) - } - }, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - raw := []byte( - "event: response.created\n" + - "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_terminal\",\"object\":\"response\",\"created_at\":1234567890,\"model\":\"gpt-4o-mini\",\"provider\":\"openai\",\"status\":\"in_progress\",\"output\":[]}}\n\n" + - "event: " + tt.eventName + "\n" + - "data: {\"type\":\"" + tt.eventName + "\",\"response\":" + tt.terminalResponse + "}\n\n" + - "data: [DONE]\n\n", - ) - - cached, ok := reconstructStreamingResponse("/v1/responses", raw, streamResponseDefaults{ - Model: "gpt-4o-mini", - Provider: "openai", - }) - if !ok { - t.Fatal("expected streamed responses payload to reconstruct successfully") - } - - var cachedResponse map[string]any - if err := json.Unmarshal(cached, &cachedResponse); err != nil { - t.Fatalf("json.Unmarshal(cached) error = %v", err) - } - if got, _ := cachedResponse["status"].(string); got != tt.status { - t.Fatalf("cached response status = %q, want %q", got, tt.status) - } - tt.assertTerminal(t, cachedResponse) - - replay, err := renderCachedResponsesStream([]byte(`{"model":"gpt-4o-mini","stream":true}`), cached) - if err != nil { - t.Fatalf("renderCachedResponsesStream() error = %v", err) - } - - var terminalEvent map[string]any - parseSSEJSONEvents(replay, func(event map[string]any) { - if eventType, _ := event["type"].(string); eventType == tt.eventName { - terminalEvent = event - } - }) - if terminalEvent == nil { - t.Fatalf("cached responses replay = %q, want terminal event %s", string(replay), tt.eventName) - } - response, ok := terminalEvent["response"].(map[string]any) - if !ok { - t.Fatalf("terminal event response = %#v, want object", terminalEvent["response"]) - } - if got, _ := response["status"].(string); got != tt.status { - t.Fatalf("terminal event status = %q, want %q", got, tt.status) - } - tt.assertTerminal(t, response) - }) - } -} diff --git a/internal/responsecache/semantic.go b/internal/responsecache/semantic.go index ef73f6e92..c88bb9034 100644 --- a/internal/responsecache/semantic.go +++ b/internal/responsecache/semantic.go @@ -581,17 +581,6 @@ func sha256HexOf(s string) string { return hex.EncodeToString(h[:]) } -// GuardrailsHashFromContext retrieves the guardrails hash from the context, -// using the core package's storage key. -func GuardrailsHashFromContext(ctx context.Context) string { - return core.GetGuardrailsHash(ctx) -} - -// WithGuardrailsHash stores the guardrails hash into the context. -func WithGuardrailsHash(ctx context.Context, hash string) context.Context { - return core.WithGuardrailsHash(ctx, hash) -} - // CacheTypeHeader values for X-Cache-Type. const ( CacheTypeExact = "exact" @@ -600,17 +589,8 @@ const ( CacheHeaderSemantic = "HIT (semantic)" ) -// ShouldSkipExactCache reports whether the X-Cache-Type header requests semantic-only mode. -func ShouldSkipExactCache(req *http.Request) bool { - return strings.EqualFold(req.Header.Get("X-Cache-Type"), CacheTypeSemantic) -} - -// ShouldSkipAllCache reports whether caching must be bypassed for this request, -// matching the exact-cache middleware semantics for no-cache and no-store. -func ShouldSkipAllCache(req *http.Request) bool { - return shouldSkipAllCacheHeaders(req.Header.Get) -} - +// shouldSkipAllCacheHeaders reports whether caching must be bypassed for this +// request, matching the exact-cache middleware semantics for no-cache and no-store. func shouldSkipAllCacheHeaders(header func(string) string) bool { if strings.EqualFold(header("X-Cache-Control"), "no-store") { return true diff --git a/internal/responsecache/semantic_test.go b/internal/responsecache/semantic_test.go index 4e5c369ec..96164b22e 100644 --- a/internal/responsecache/semantic_test.go +++ b/internal/responsecache/semantic_test.go @@ -576,19 +576,19 @@ func TestMapVecStore_DeleteExpiredOnlyRemovesExpired(t *testing.T) { } } -func TestShouldSkipAllCache_CacheControlNoStore(t *testing.T) { +func TestShouldSkipAllCacheHeaders_CacheControlNoStore(t *testing.T) { req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil) req.Header.Set("Cache-Control", "private, no-store, max-age=0") - if !ShouldSkipAllCache(req) { - t.Fatal("expected ShouldSkipAllCache for Cache-Control: no-store") + if !shouldSkipAllCacheHeaders(req.Header.Get) { + t.Fatal("expected cache skip for Cache-Control: no-store") } } -func TestShouldSkipAllCache_CacheControlNoCache(t *testing.T) { +func TestShouldSkipAllCacheHeaders_CacheControlNoCache(t *testing.T) { req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil) req.Header.Set("Cache-Control", "private, no-cache, max-age=0") - if !ShouldSkipAllCache(req) { - t.Fatal("expected ShouldSkipAllCache for Cache-Control: no-cache") + if !shouldSkipAllCacheHeaders(req.Header.Get) { + t.Fatal("expected cache skip for Cache-Control: no-cache") } } diff --git a/internal/responsecache/stream_cache.go b/internal/responsecache/stream_cache.go index 74cac062a..f21232f57 100644 --- a/internal/responsecache/stream_cache.go +++ b/internal/responsecache/stream_cache.go @@ -2,7 +2,6 @@ package responsecache import ( "bytes" - "maps" "net/http" "strings" @@ -21,11 +20,6 @@ var ( cacheDonePayload = []byte("[DONE]") ) -type streamResponseDefaults struct { - Model string - Provider string -} - func cacheKeyRequestBody(path string, body []byte) []byte { switch path { case "/v1/chat/completions": @@ -103,79 +97,6 @@ func cacheHeaderValue(cacheType string) string { } } -func reconstructStreamingResponse(path string, raw []byte, defaults streamResponseDefaults) ([]byte, bool) { - switch path { - case "/v1/chat/completions": - builder := &chatStreamCacheBuilder{ - defaults: defaults, - Choices: make(map[int]*chatChoiceState), - } - parseSSEJSONEvents(raw, builder.OnJSONEvent) - return builder.Build() - case "/v1/responses": - builder := &responsesStreamCacheBuilder{ - defaults: defaults, - Output: make(map[int]*responsesOutputState), - ItemIDs: make(map[string]int), - } - parseSSEJSONEvents(raw, builder.OnJSONEvent) - return builder.Build() - default: - return nil, false - } -} - -func parseSSEJSONEvents(raw []byte, onJSON func(map[string]any)) { - for len(raw) > 0 { - idx, sepLen := nextCacheEventBoundary(raw) - event := raw - if idx != -1 { - event = raw[:idx] - raw = raw[idx+sepLen:] - } else { - raw = nil - } - - payload, ok := parseCacheEventJSON(event) - if !ok { - if idx == -1 { - break - } - continue - } - onJSON(payload) - if idx == -1 { - break - } - } -} - -func parseCacheEventJSON(event []byte) (map[string]any, bool) { - lines := bytes.Split(event, []byte("\n")) - payloadLines := make([][]byte, 0, len(lines)) - for _, line := range lines { - data, ok := parseCacheDataLine(line) - if !ok { - continue - } - payloadLines = append(payloadLines, data) - } - if len(payloadLines) == 0 { - return nil, false - } - - jsonData := bytes.Join(payloadLines, []byte("\n")) - if bytes.Equal(jsonData, cacheDonePayload) { - return nil, false - } - - var payload map[string]any - if err := json.Unmarshal(jsonData, &payload); err != nil { - return nil, false - } - return payload, true -} - func nextCacheEventBoundary(data []byte) (idx int, sepLen int) { lfIdx := bytes.Index(data, cacheLFEventBoundary) crlfIdx := bytes.Index(data, cacheCRLFEventBoundary) @@ -212,88 +133,3 @@ func normalizeStreamOptionsForCache(src *core.StreamOptions) *core.StreamOptions cloned := *src return &cloned } - -func streamIncludeUsageRequested(path string, requestBody []byte) bool { - switch path { - case "/v1/chat/completions": - req, err := core.DecodeChatRequest(requestBody, nil) - if err != nil || req == nil { - return false - } - return normalizeStreamOptionsForCache(req.StreamOptions) != nil - case "/v1/responses": - req, err := core.DecodeResponsesRequest(requestBody, nil) - if err != nil || req == nil { - return false - } - return normalizeStreamOptionsForCache(req.StreamOptions) != nil - default: - return false - } -} - -func appendSSEJSONEvent(out *bytes.Buffer, eventName string, payload any) error { - data, err := json.Marshal(payload) - if err != nil { - return err - } - if eventName != "" { - out.WriteString("event: ") - out.WriteString(eventName) - out.WriteByte('\n') - } - out.WriteString("data: ") - out.Write(data) - out.WriteString("\n\n") - return nil -} - -func toJSONMap(value any) (map[string]any, error) { - data, err := json.Marshal(value) - if err != nil { - return nil, err - } - var result map[string]any - if err := json.Unmarshal(data, &result); err != nil { - return nil, err - } - return result, nil -} - -func cloneJSONMap(src map[string]any) map[string]any { - return maps.Clone(src) -} - -func jsonNumberToInt(value any) (int, bool) { - switch v := value.(type) { - case float64: - return int(v), true - case int: - return v, true - case int64: - return int(v), true - default: - return 0, false - } -} - -func jsonNumberToInt64(value any) (int64, bool) { - switch v := value.(type) { - case float64: - return int64(v), true - case int: - return int64(v), true - case int64: - return v, true - default: - return 0, false - } -} - -func nonEmpty(value, fallback string) string { - value = strings.TrimSpace(value) - if value != "" { - return value - } - return strings.TrimSpace(fallback) -} diff --git a/internal/responsecache/stream_cache_chat.go b/internal/responsecache/stream_cache_chat.go deleted file mode 100644 index 51da756e1..000000000 --- a/internal/responsecache/stream_cache_chat.go +++ /dev/null @@ -1,396 +0,0 @@ -package responsecache - -import ( - "bytes" - "sort" - "strings" - - "github.com/goccy/go-json" - - "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 deleted file mode 100644 index 0531f297d..000000000 --- a/internal/responsecache/stream_cache_responses.go +++ /dev/null @@ -1,628 +0,0 @@ -package responsecache - -import ( - "bytes" - "sort" - "strings" - - "github.com/goccy/go-json" -) - -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 -} diff --git a/internal/responsecache/vecstore_map.go b/internal/responsecache/vecstore_map_test.go similarity index 100% rename from internal/responsecache/vecstore_map.go rename to internal/responsecache/vecstore_map_test.go diff --git a/internal/server/auth.go b/internal/server/auth.go index 1ea9fa663..85390f965 100644 --- a/internal/server/auth.go +++ b/internal/server/auth.go @@ -20,15 +20,10 @@ type BearerTokenAuthenticator interface { Authenticate(ctx context.Context, token string) (authkeys.AuthenticationResult, error) } -// AuthMiddleware creates an Echo middleware that validates the master key -// if it's configured. If masterKey is empty, no authentication is required. -// skipPaths is a list of paths that should bypass authentication. -func AuthMiddleware(masterKey string, skipPaths []string) echo.MiddlewareFunc { - return AuthMiddlewareWithAuthenticator(masterKey, nil, skipPaths) -} - -// AuthMiddlewareWithAuthenticator validates the legacy master key and, when -// configured, managed auth keys from the auth key service. +// AuthMiddlewareWithAuthenticator creates an Echo middleware that validates +// the legacy master key and, when configured, managed auth keys from the auth +// key service. If no auth mechanism is configured, no authentication is +// required. skipPaths is a list of paths that should bypass authentication. func AuthMiddlewareWithAuthenticator(masterKey string, authenticator BearerTokenAuthenticator, skipPaths []string, userPathHeader ...string) echo.MiddlewareFunc { userPathHeaderName := configuredUserPathHeaderName(userPathHeader...) return func(next echo.HandlerFunc) echo.HandlerFunc { diff --git a/internal/server/handlers.go b/internal/server/handlers.go index d0f0c6b24..f72d4f37e 100644 --- a/internal/server/handlers.go +++ b/internal/server/handlers.go @@ -51,34 +51,8 @@ type Handler struct { translatedSvcOnce sync.Once } -// NewHandler creates a new handler with the given routable provider (typically the Router) -func NewHandler(provider core.RoutableProvider, logger auditlog.LoggerInterface, usageLogger usage.LoggerInterface, pricingResolver usage.PricingResolver) *Handler { - return newHandler(provider, logger, usageLogger, pricingResolver, nil, nil, nil, nil) -} - -func newHandler( - provider core.RoutableProvider, - logger auditlog.LoggerInterface, - usageLogger usage.LoggerInterface, - pricingResolver usage.PricingResolver, - modelResolver RequestModelResolver, - workflowPolicyResolver RequestWorkflowPolicyResolver, - failoverResolver RequestFailoverResolver, - translatedRequestPatcher TranslatedRequestPatcher, -) *Handler { - return newHandlerWithAuthorizer( - provider, - logger, - usageLogger, - pricingResolver, - modelResolver, - nil, - workflowPolicyResolver, - failoverResolver, - translatedRequestPatcher, - ) -} - +// newHandlerWithAuthorizer creates a new handler with the given routable +// provider (typically the Router) and optional resolvers. func newHandlerWithAuthorizer( provider core.RoutableProvider, logger auditlog.LoggerInterface, diff --git a/internal/server/http_test.go b/internal/server/http_test.go index a29171011..8f448bce4 100644 --- a/internal/server/http_test.go +++ b/internal/server/http_test.go @@ -580,7 +580,7 @@ func TestServer_ManagedAuthKeyUserPathOverridesHeaderBeforeWorkflowResolution(t func newDashboardHandler(t *testing.T) *dashboard.Handler { t.Helper() - h, err := dashboard.New() + h, err := dashboard.NewWithBasePath("/") if err != nil { t.Fatalf("failed to create dashboard handler: %v", err) } diff --git a/internal/server/model_validation.go b/internal/server/model_validation.go index 728e43af0..37d924270 100644 --- a/internal/server/model_validation.go +++ b/internal/server/model_validation.go @@ -10,23 +10,13 @@ import ( "gomodel/internal/core" ) -// WorkflowResolution resolves the request-scoped workflow for model-facing -// routes. The workflow centralizes endpoint capabilities, execution mode, resolved -// provider type, and any early model routing decision that downstream handlers -// or middleware need to consume. -func WorkflowResolution(provider core.RoutableProvider) echo.MiddlewareFunc { - return WorkflowResolutionWithResolverAndPolicy(provider, nil, nil) -} - -// WorkflowResolutionWithResolver resolves request-scoped workflows using -// an explicit selector resolver when provided. This lets workflow resolution own -// alias policy instead of depending on provider decorators. -func WorkflowResolutionWithResolver(provider core.RoutableProvider, resolver RequestModelResolver) echo.MiddlewareFunc { - return WorkflowResolutionWithResolverAndPolicy(provider, resolver, nil) -} - -// WorkflowResolutionWithResolverAndPolicy resolves request-scoped workflows -// and matches one persisted workflow policy when configured. +// WorkflowResolutionWithResolverAndPolicy resolves the request-scoped workflow +// for model-facing routes and matches one persisted workflow policy when +// configured. The workflow centralizes endpoint capabilities, execution mode, +// resolved provider type, and any early model routing decision that downstream +// handlers or middleware need to consume. An explicit selector resolver lets +// workflow resolution own alias policy instead of depending on provider +// decorators. func WorkflowResolutionWithResolverAndPolicy( provider core.RoutableProvider, resolver RequestModelResolver, @@ -277,13 +267,3 @@ func passthroughRouteInfo(c *echo.Context) *core.PassthroughRouteInfo { AuditPath: c.Request().URL.Path, } } - -// GetProviderType returns the provider type captured in the workflow for this request. -func GetProviderType(c *echo.Context) string { - if workflow := core.GetWorkflow(c.Request().Context()); workflow != nil { - if providerType := strings.TrimSpace(workflow.ProviderType); providerType != "" { - return providerType - } - } - return "" -} diff --git a/internal/server/model_validation_test.go b/internal/server/model_validation_test.go index e66f12951..068da175d 100644 --- a/internal/server/model_validation_test.go +++ b/internal/server/model_validation_test.go @@ -916,27 +916,6 @@ func TestModelValidation_DoesNotCacheCanonicalResponsesRequestWhenRouteHintsAlre assert.Equal(t, "gpt-4o-mini", capturedEnv.RouteHints.Model) } -func TestGetProviderType_EmptyWhenNotSet(t *testing.T) { - e := echo.New() - req := httptest.NewRequest(http.MethodGet, "/health", nil) - rec := httptest.NewRecorder() - c := e.NewContext(req, rec) - - assert.Equal(t, "", GetProviderType(c)) -} - -func TestGetProviderType_UsesWorkflow(t *testing.T) { - e := echo.New() - req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil) - req = req.WithContext(core.WithWorkflow(req.Context(), &core.Workflow{ - ProviderType: "openai", - })) - rec := httptest.NewRecorder() - c := e.NewContext(req, rec) - - assert.Equal(t, "openai", GetProviderType(c)) -} - func TestSelectorHintsFromJSONGJSON_MatchesStdlibSemantics(t *testing.T) { tests := []struct { name string diff --git a/internal/server/request_model_resolution.go b/internal/server/request_model_resolution.go index 624d15873..4d19cfb5d 100644 --- a/internal/server/request_model_resolution.go +++ b/internal/server/request_model_resolution.go @@ -22,10 +22,6 @@ func workflowProviderNameForType(provider core.RoutableProvider, providerType st return gateway.WorkflowProviderNameForType(provider, providerType) } -func resolveRequestModel(provider core.RoutableProvider, resolver RequestModelResolver, requested core.RequestedModelSelector) (*core.RequestModelResolution, error) { - return gateway.ResolveRequestModel(provider, resolver, requested) -} - func resolveRequestModelWithAuthorizer( ctx context.Context, provider core.RoutableProvider, diff --git a/internal/server/request_model_resolution_test.go b/internal/server/request_model_resolution_test.go index deeb72cff..459906646 100644 --- a/internal/server/request_model_resolution_test.go +++ b/internal/server/request_model_resolution_test.go @@ -91,9 +91,9 @@ func TestResolveRequestModel_UsesResolvedProviderNameInsteadOfSelectorPrefix(t * }, } - resolution, err := resolveRequestModel(provider, nil, core.NewRequestedModelSelector("openai/gpt-5-nano", "")) + resolution, err := resolveRequestModelWithAuthorizer(context.Background(), provider, nil, nil, core.NewRequestedModelSelector("openai/gpt-5-nano", "")) if err != nil { - t.Fatalf("resolveRequestModel() error = %v", err) + t.Fatalf("resolveRequestModelWithAuthorizer() error = %v", err) } if got := resolution.ResolvedSelector.Provider; got != "openai" { @@ -117,9 +117,9 @@ func TestResolveRequestModel_CanonicalizesProviderTypeSelectorToConcreteProvider }, } - resolution, err := resolveRequestModel(provider, nil, core.NewRequestedModelSelector("openai/gpt-5-nano", "")) + resolution, err := resolveRequestModelWithAuthorizer(context.Background(), provider, nil, nil, core.NewRequestedModelSelector("openai/gpt-5-nano", "")) if err != nil { - t.Fatalf("resolveRequestModel() error = %v", err) + t.Fatalf("resolveRequestModelWithAuthorizer() error = %v", err) } if got := resolution.ResolvedQualifiedModel(); got != "openai_test/gpt-5-nano" { @@ -156,9 +156,9 @@ func TestResolveRequestModel_CanonicalizesAliasOutputThroughProviderResolver(t * }, } - resolution, err := resolveRequestModel(provider, aliasResolverStub{}, core.NewRequestedModelSelector("anthropic/claude-opus-4-6", "")) + resolution, err := resolveRequestModelWithAuthorizer(context.Background(), provider, aliasResolverStub{}, nil, core.NewRequestedModelSelector("anthropic/claude-opus-4-6", "")) if err != nil { - t.Fatalf("resolveRequestModel() error = %v", err) + t.Fatalf("resolveRequestModelWithAuthorizer() error = %v", err) } if !resolution.AliasApplied { diff --git a/internal/server/seams_test.go b/internal/server/seams_test.go new file mode 100644 index 000000000..86afeb331 --- /dev/null +++ b/internal/server/seams_test.go @@ -0,0 +1,90 @@ +package server + +// Test-only convenience wrappers over the production constructors and +// middleware. Production wires the fuller variants directly +// (newHandlerWithAuthorizer at http.go, AuthMiddlewareWithAuthenticator, +// WorkflowResolutionWithResolverAndPolicy); tests use these to avoid +// repeating nil arguments. + +import ( + "io" + "strings" + + "github.com/labstack/echo/v5" + + "gomodel/internal/auditlog" + "gomodel/internal/core" + "gomodel/internal/usage" +) + +// NewHandler creates a handler with the given routable provider and no +// optional resolvers. +func NewHandler(provider core.RoutableProvider, logger auditlog.LoggerInterface, usageLogger usage.LoggerInterface, pricingResolver usage.PricingResolver) *Handler { + return newHandler(provider, logger, usageLogger, pricingResolver, nil, nil, nil, nil) +} + +func newHandler( + provider core.RoutableProvider, + logger auditlog.LoggerInterface, + usageLogger usage.LoggerInterface, + pricingResolver usage.PricingResolver, + modelResolver RequestModelResolver, + workflowPolicyResolver RequestWorkflowPolicyResolver, + failoverResolver RequestFailoverResolver, + translatedRequestPatcher TranslatedRequestPatcher, +) *Handler { + return newHandlerWithAuthorizer( + provider, + logger, + usageLogger, + pricingResolver, + modelResolver, + nil, + workflowPolicyResolver, + failoverResolver, + translatedRequestPatcher, + ) +} + +// AuthMiddleware validates the master key without a managed-key authenticator. +func AuthMiddleware(masterKey string, skipPaths []string) echo.MiddlewareFunc { + return AuthMiddlewareWithAuthenticator(masterKey, nil, skipPaths) +} + +// WorkflowResolution resolves request-scoped workflows without an explicit +// selector resolver or policy resolver. +func WorkflowResolution(provider core.RoutableProvider) echo.MiddlewareFunc { + return WorkflowResolutionWithResolverAndPolicy(provider, nil, nil) +} + +// WorkflowResolutionWithResolver resolves request-scoped workflows using an +// explicit selector resolver when provided. +func WorkflowResolutionWithResolver(provider core.RoutableProvider, resolver RequestModelResolver) echo.MiddlewareFunc { + return WorkflowResolutionWithResolverAndPolicy(provider, resolver, nil) +} + +// GetProviderType returns the provider type captured in the workflow for this request. +func GetProviderType(c *echo.Context) string { + if workflow := core.GetWorkflow(c.Request().Context()); workflow != nil { + if providerType := strings.TrimSpace(workflow.ProviderType); providerType != "" { + return providerType + } + } + return "" +} + +// handleStreamingResponse obtains the stream from streamFn and relays it, +// recording dispatch errors the way the live inference paths do before they +// call handleStreamingReadCloser. +func (s *translatedInferenceService) handleStreamingResponse( + c *echo.Context, + workflow *core.Workflow, + model, provider, providerName string, + streamFn func() (io.ReadCloser, error), +) error { + stream, err := streamFn() + if err != nil { + return handleStreamingDispatchError(c, err) + } + return s.handleStreamingReadCloser(c, workflow, model, provider, providerName, "", stream, nil) +} diff --git a/internal/server/translated_inference_service.go b/internal/server/translated_inference_service.go index 2ab81d0f2..e6884e05f 100644 --- a/internal/server/translated_inference_service.go +++ b/internal/server/translated_inference_service.go @@ -565,19 +565,6 @@ func (s *translatedInferenceService) handleStreamingReadCloser( return nil } -func (s *translatedInferenceService) handleStreamingResponse( - c *echo.Context, - workflow *core.Workflow, - model, provider, providerName string, - streamFn func() (io.ReadCloser, error), -) error { - stream, err := streamFn() - if err != nil { - return handleStreamingDispatchError(c, err) - } - return s.handleStreamingReadCloser(c, workflow, model, provider, providerName, "", stream, nil) -} - // handleStreamingDispatchError records audit context for a streaming request // that failed before any chunks could be flushed. It marks the entry as // streaming and distinguishes client cancellations from upstream failures so diff --git a/internal/server/workflow_helpers.go b/internal/server/workflow_helpers.go index 6b5f758c8..eacdfd139 100644 --- a/internal/server/workflow_helpers.go +++ b/internal/server/workflow_helpers.go @@ -10,17 +10,6 @@ import ( "gomodel/internal/gateway" ) -func ensureTranslatedRequestWorkflow( - c *echo.Context, - provider core.RoutableProvider, - resolver RequestModelResolver, - policyResolver RequestWorkflowPolicyResolver, - model, - providerHint *string, -) (*core.Workflow, error) { - return ensureTranslatedRequestWorkflowWithAuthorizer(c, provider, resolver, nil, policyResolver, model, providerHint) -} - func ensureTranslatedRequestWorkflowWithAuthorizer( c *echo.Context, provider core.RoutableProvider, diff --git a/internal/server/workflow_helpers_test.go b/internal/server/workflow_helpers_test.go index 9a0ea1299..1f4685daf 100644 --- a/internal/server/workflow_helpers_test.go +++ b/internal/server/workflow_helpers_test.go @@ -35,7 +35,7 @@ func TestEnsureTranslatedRequestWorkflow_CompletesPartialWorkflowFromDecodedSele model := "gpt-4o-mini" providerHint := "" - workflow, err := ensureTranslatedRequestWorkflow(c, provider, nil, nil, &model, &providerHint) + workflow, err := ensureTranslatedRequestWorkflowWithAuthorizer(c, provider, nil, nil, nil, &model, &providerHint) require.NoError(t, err) require.NotNil(t, workflow) diff --git a/internal/server/workflow_policy_test.go b/internal/server/workflow_policy_test.go index 131e8f4bf..e7f362de3 100644 --- a/internal/server/workflow_policy_test.go +++ b/internal/server/workflow_policy_test.go @@ -74,9 +74,9 @@ func TestDetermineBatchExecutionSelection_UsesSingleResolutionPass(t *testing.T) }, } - selection, err := gateway.DetermineBatchExecutionSelectionWithAuthorizer(context.Background(), provider, resolver, nil, req) + selection, err := gateway.DetermineBatchExecutionSelectionWithAuthorizerAndInputFileResolver(context.Background(), provider, resolver, nil, nil, req) if err != nil { - t.Fatalf("DetermineBatchExecutionSelectionWithAuthorizer() error = %v", err) + t.Fatalf("DetermineBatchExecutionSelectionWithAuthorizerAndInputFileResolver() error = %v", err) } if selection.ProviderType != "openai" { t.Fatalf("providerType = %q, want openai", selection.ProviderType) diff --git a/internal/workflows/compiler.go b/internal/workflows/compiler.go index 66184dd1f..ddcd520ac 100644 --- a/internal/workflows/compiler.go +++ b/internal/workflows/compiler.go @@ -13,11 +13,6 @@ type compiler struct { featureCaps core.WorkflowFeatures } -// NewCompiler creates the default workflow compiler for the v1 payload. -func NewCompiler(registry guardrails.Catalog) Compiler { - return NewCompilerWithFeatureCaps(registry, core.DefaultWorkflowFeatures()) -} - // NewCompilerWithFeatureCaps creates the default workflow compiler for the // v1 payload with process-level feature caps applied at compile time. func NewCompilerWithFeatureCaps(registry guardrails.Catalog, featureCaps core.WorkflowFeatures) Compiler { diff --git a/internal/workflows/compiler_test.go b/internal/workflows/compiler_test.go index e9a3f7cfa..72ef14767 100644 --- a/internal/workflows/compiler_test.go +++ b/internal/workflows/compiler_test.go @@ -22,7 +22,7 @@ func TestCompilerCompile_Guardrails(t *testing.T) { t.Fatalf("Register() error = %v", err) } - compiled, err := NewCompiler(registry).Compile(Version{ + compiled, err := NewCompilerWithFeatureCaps(registry, core.DefaultWorkflowFeatures()).Compile(Version{ ID: "workflow-1", Scope: Scope{}, Version: 3, @@ -133,7 +133,7 @@ func TestCompilerCompile_DefaultsFailoverEnabledWhenUnset(t *testing.T) { } func TestCompilerCompile_ReturnsGatewayErrorWhenGuardrailsCatalogIsEmpty(t *testing.T) { - _, err := NewCompiler(guardrails.NewRegistry()).Compile(Version{ + _, err := NewCompilerWithFeatureCaps(guardrails.NewRegistry(), core.DefaultWorkflowFeatures()).Compile(Version{ ID: "workflow-1", Scope: Scope{}, Version: 1, @@ -169,7 +169,7 @@ func TestCompilerCompile_WrapsBuildPipelineErrorsAsGatewayErrors(t *testing.T) { }); err != nil { t.Fatalf("Register() error = %v", err) } - _, err = NewCompiler(registry).Compile(Version{ + _, err = NewCompilerWithFeatureCaps(registry, core.DefaultWorkflowFeatures()).Compile(Version{ ID: "workflow-1", Scope: Scope{}, Version: 1, diff --git a/internal/workflows/service_test.go b/internal/workflows/service_test.go index 49674c053..f45eb7dfe 100644 --- a/internal/workflows/service_test.go +++ b/internal/workflows/service_test.go @@ -361,7 +361,7 @@ func TestServiceMatch_MostSpecificWins(t *testing.T) { }, } - service, err := NewService(store, NewCompiler(nil)) + service, err := NewService(store, NewCompilerWithFeatureCaps(nil, core.DefaultWorkflowFeatures())) if err != nil { t.Fatalf("NewService() error = %v", err) } @@ -393,7 +393,7 @@ func TestServiceMatch_MostSpecificWins(t *testing.T) { func TestServiceEnsureDefaultGlobal_CreatesWhenMissing(t *testing.T) { store := &staticStore{} - service, err := NewService(store, NewCompiler(nil)) + service, err := NewService(store, NewCompilerWithFeatureCaps(nil, core.DefaultWorkflowFeatures())) if err != nil { t.Fatalf("NewService() error = %v", err) } @@ -450,7 +450,7 @@ func TestServiceEnsureDefaultGlobal_ReconcilesManagedDefault(t *testing.T) { }, }, } - service, err := NewService(store, NewCompiler(nil)) + service, err := NewService(store, NewCompilerWithFeatureCaps(nil, core.DefaultWorkflowFeatures())) if err != nil { t.Fatalf("NewService() error = %v", err) } @@ -506,7 +506,7 @@ func TestServiceEnsureDefaultGlobal_PreservesCustomGlobal(t *testing.T) { }, }, } - service, err := NewService(store, NewCompiler(nil)) + service, err := NewService(store, NewCompilerWithFeatureCaps(nil, core.DefaultWorkflowFeatures())) if err != nil { t.Fatalf("NewService() error = %v", err) } @@ -553,7 +553,7 @@ func TestServiceEnsureDefaultGlobal_LoadsPreservedCustomGlobalIntoSnapshot(t *te }, }, } - service, err := NewService(store, NewCompiler(nil)) + service, err := NewService(store, NewCompilerWithFeatureCaps(nil, core.DefaultWorkflowFeatures())) if err != nil { t.Fatalf("NewService() error = %v", err) } @@ -587,7 +587,7 @@ func TestServiceEnsureDefaultGlobal_ValidatesBeforeStoreMutation(t *testing.T) { store := &concurrentStore{ createCalled: make(chan struct{}, 1), } - service, err := NewService(store, &previewEmptyCompiler{delegate: NewCompiler(nil)}) + service, err := NewService(store, &previewEmptyCompiler{delegate: NewCompilerWithFeatureCaps(nil, core.DefaultWorkflowFeatures())}) if err != nil { t.Fatalf("NewService() error = %v", err) } @@ -805,7 +805,7 @@ func TestServiceCreate_RefreshesSnapshot(t *testing.T) { }, }, } - service, err := NewService(store, NewCompiler(nil)) + service, err := NewService(store, NewCompilerWithFeatureCaps(nil, core.DefaultWorkflowFeatures())) if err != nil { t.Fatalf("NewService() error = %v", err) } @@ -930,7 +930,7 @@ func TestServiceListViews_AnnotatesCompileFailuresPerRow(t *testing.T) { }, } service, err := NewService(store, &versionFailingCompiler{ - delegate: NewCompiler(nil), + delegate: NewCompilerWithFeatureCaps(nil, core.DefaultWorkflowFeatures()), version: "provider-v1", err: errors.New("compile failed for provider-v1"), }) @@ -1002,7 +1002,7 @@ func TestServiceDeactivate_RefreshesSnapshot(t *testing.T) { }, }, } - service, err := NewService(store, NewCompiler(nil)) + service, err := NewService(store, NewCompilerWithFeatureCaps(nil, core.DefaultWorkflowFeatures())) if err != nil { t.Fatalf("NewService() error = %v", err) } @@ -1043,7 +1043,7 @@ func TestServiceDeactivate_RejectsGlobalWorkflow(t *testing.T) { }, }, } - service, err := NewService(store, NewCompiler(nil)) + service, err := NewService(store, NewCompilerWithFeatureCaps(nil, core.DefaultWorkflowFeatures())) if err != nil { t.Fatalf("NewService() error = %v", err) } @@ -1089,7 +1089,7 @@ func TestServiceDeactivate_AllowsPathScopedWorkflow(t *testing.T) { }, }, } - service, err := NewService(store, NewCompiler(nil)) + service, err := NewService(store, NewCompilerWithFeatureCaps(nil, core.DefaultWorkflowFeatures())) if err != nil { t.Fatalf("NewService() error = %v", err) } @@ -1124,7 +1124,7 @@ func TestServiceCreateWaitsForInFlightRefreshBeforePersisting(t *testing.T) { createCalled: make(chan struct{}, 1), } compiler := &blockingCompiler{ - delegate: NewCompiler(nil), + delegate: NewCompilerWithFeatureCaps(nil, core.DefaultWorkflowFeatures()), blockCall: 2, blocked: make(chan struct{}), release: make(chan struct{}), @@ -1211,7 +1211,7 @@ func TestServiceCreateRejectsEmptyCompiledPreviewBeforePersisting(t *testing.T) }, createCalled: make(chan struct{}, 1), } - service, err := NewService(store, &previewEmptyCompiler{delegate: NewCompiler(nil)}) + service, err := NewService(store, &previewEmptyCompiler{delegate: NewCompilerWithFeatureCaps(nil, core.DefaultWorkflowFeatures())}) if err != nil { t.Fatalf("NewService() error = %v", err) } @@ -1266,7 +1266,7 @@ func TestServiceCreateRefreshIgnoresRequestContextCancellationAfterPersist(t *te }, }, } - service, err := NewService(store, NewCompiler(nil)) + service, err := NewService(store, NewCompilerWithFeatureCaps(nil, core.DefaultWorkflowFeatures())) if err != nil { t.Fatalf("NewService() error = %v", err) } @@ -1324,7 +1324,7 @@ func TestServiceCreateReturnsSuccessWhenReloadRefreshFailsAfterPersist(t *testin }, }, } - service, err := NewService(store, NewCompiler(nil)) + service, err := NewService(store, NewCompilerWithFeatureCaps(nil, core.DefaultWorkflowFeatures())) if err != nil { t.Fatalf("NewService() error = %v", err) } @@ -1391,7 +1391,7 @@ func TestServiceDeactivateRefreshIgnoresRequestContextCancellationAfterPersist(t }, }, } - service, err := NewService(store, NewCompiler(nil)) + service, err := NewService(store, NewCompilerWithFeatureCaps(nil, core.DefaultWorkflowFeatures())) if err != nil { t.Fatalf("NewService() error = %v", err) } @@ -1449,7 +1449,7 @@ func TestServiceDeactivateReturnsSuccessWhenReloadRefreshFailsAfterPersist(t *te }, }, } - service, err := NewService(store, NewCompiler(nil)) + service, err := NewService(store, NewCompilerWithFeatureCaps(nil, core.DefaultWorkflowFeatures())) if err != nil { t.Fatalf("NewService() error = %v", err) } diff --git a/tests/e2e/setup_test.go b/tests/e2e/setup_test.go index a8e234600..3689889bb 100644 --- a/tests/e2e/setup_test.go +++ b/tests/e2e/setup_test.go @@ -91,7 +91,7 @@ func setupE2EServer(t *testing.T, opts e2eServerOptions) *server.Server { if opts.adminUIEnabled { cfg.AdminUIEnabled = true - dashHandler, err := dashboard.New() + dashHandler, err := dashboard.NewWithBasePath("/") require.NoError(t, err, "failed to create dashboard handler") cfg.DashboardHandler = dashHandler }