From fec171a1fce1fbb04e5e210c5bb2289a0aa713a6 Mon Sep 17 00:00:00 2001 From: Jarvis Date: Sat, 1 Aug 2026 18:03:41 +1000 Subject: [PATCH 1/3] feat: add centralized inference gateway Centralize provider credentials, OAuth, search, fetch, policy, and attributed usage behind a versioned gateway while preserving satellite provider:model selection and local tool execution. Expose POST /v1/responses and GET /v1/models with provider/model namespaces for Discourse and other OpenAI clients, including streaming, reasoning, images, function calls, stateless continuation, usage, and cancellation. --- .github/workflows/ci.yml | 3 + cmd/flags.go | 2 +- cmd/gateway.go | 459 ++++++ cmd/gateway_enroll_test.go | 119 ++ cmd/gateway_tools_test.go | 28 + cmd/gateway_usage_test.go | 38 + cmd/models.go | 31 +- cmd/runner.go | 14 +- cmd/tools.go | 31 +- cmd/usage.go | 2 +- docs-site/content/guides/inference-gateway.md | 240 +++ docs-site/content/reference/configuration.md | 18 +- internal/config/config.go | 225 ++- internal/config/gateway_test.go | 112 ++ internal/config/schema.go | 15 + internal/gateway/catalog_test.go | 324 ++++ internal/gateway/error_test.go | 72 + internal/gateway/execution.go | 128 ++ internal/gateway/integration_test.go | 661 ++++++++ internal/gateway/limits_test.go | 56 + internal/gateway/protocol/protocol.go | 133 ++ internal/gateway/responses.go | 695 ++++++++ .../gateway/responses_integration_test.go | 816 ++++++++++ internal/gateway/responses_wire.go | 517 ++++++ internal/gateway/seal.go | 85 + internal/gateway/security_test.go | 414 +++++ internal/gateway/server.go | 1415 +++++++++++++++++ internal/gateway/store.go | 433 +++++ internal/gateway/store_lock_unix.go | 35 + internal/gateway/store_lock_windows.go | 37 + internal/gateway/testproviders_test.go | 89 ++ internal/gateway/usage.go | 96 ++ internal/llm/engine.go | 28 +- internal/llm/factory.go | 61 +- internal/llm/gateway_catalog_test.go | 165 ++ internal/llm/gateway_provider.go | 701 ++++++++ internal/llm/gateway_retry_budget_test.go | 105 ++ internal/llm/gateway_routing_test.go | 131 ++ internal/llm/gateway_stream_test.go | 65 + internal/llm/gateway_wire.go | 263 +++ internal/llm/gateway_wire_test.go | 100 ++ internal/llm/models.go | 45 +- internal/llm/retry.go | 18 +- internal/search/factory.go | 3 + internal/search/gateway.go | 109 ++ internal/search/gateway_test.go | 22 + internal/tools/config.go | 8 +- internal/tools/gateway_tools_test.go | 12 + internal/usage/gateway_test.go | 49 + internal/usage/logger.go | 1 + internal/usage/types.go | 17 +- ops/gateway-compose.yaml | 66 + ops/gateway-config.yaml | 10 + ops/gateway_compose_test.go | 46 + ops/satellite-config.yaml | 11 + 55 files changed, 9304 insertions(+), 75 deletions(-) create mode 100644 cmd/gateway.go create mode 100644 cmd/gateway_enroll_test.go create mode 100644 cmd/gateway_tools_test.go create mode 100644 cmd/gateway_usage_test.go create mode 100644 docs-site/content/guides/inference-gateway.md create mode 100644 internal/config/gateway_test.go create mode 100644 internal/gateway/catalog_test.go create mode 100644 internal/gateway/error_test.go create mode 100644 internal/gateway/execution.go create mode 100644 internal/gateway/integration_test.go create mode 100644 internal/gateway/limits_test.go create mode 100644 internal/gateway/protocol/protocol.go create mode 100644 internal/gateway/responses.go create mode 100644 internal/gateway/responses_integration_test.go create mode 100644 internal/gateway/responses_wire.go create mode 100644 internal/gateway/seal.go create mode 100644 internal/gateway/security_test.go create mode 100644 internal/gateway/server.go create mode 100644 internal/gateway/store.go create mode 100644 internal/gateway/store_lock_unix.go create mode 100644 internal/gateway/store_lock_windows.go create mode 100644 internal/gateway/testproviders_test.go create mode 100644 internal/gateway/usage.go create mode 100644 internal/llm/gateway_catalog_test.go create mode 100644 internal/llm/gateway_provider.go create mode 100644 internal/llm/gateway_retry_budget_test.go create mode 100644 internal/llm/gateway_routing_test.go create mode 100644 internal/llm/gateway_stream_test.go create mode 100644 internal/llm/gateway_wire.go create mode 100644 internal/llm/gateway_wire_test.go create mode 100644 internal/search/gateway.go create mode 100644 internal/search/gateway_test.go create mode 100644 internal/tools/gateway_tools_test.go create mode 100644 internal/usage/gateway_test.go create mode 100644 ops/gateway-compose.yaml create mode 100644 ops/gateway-config.yaml create mode 100644 ops/gateway_compose_test.go create mode 100644 ops/satellite-config.yaml diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 731555967..d8e8150ae 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -37,6 +37,9 @@ jobs: exit 1 fi + - name: Validate gateway Compose example + run: go test ./ops -run '^TestGatewayComposeExampleParsesAndHasHealthGating$' + - name: Build run: go build ./... diff --git a/cmd/flags.go b/cmd/flags.go index bd6f9aa84..881aa5d2d 100644 --- a/cmd/flags.go +++ b/cmd/flags.go @@ -260,7 +260,7 @@ func AddMaxOutputTokensFlag(cmd *cobra.Command, dest *int) { // AddToolFlags adds tool-related flags (--tools, --read-dir, --write-dir, --shell-allow) func AddToolFlags(cmd *cobra.Command, tools *string, readDirs, writeDirs, shellAllow *[]string) { - cmd.Flags().StringVar(tools, "tools", "", "Enable local tools (comma-separated, or 'all'): read_file,write_file,edit_file,shell,grep,glob,view_image,show_image,image_generate,ask_user,spawn_agent,queue_agent,wait_for_jobs") + cmd.Flags().StringVar(tools, "tools", "", "Enable local tools (comma-separated, 'all', or 'none'): read_file,write_file,edit_file,shell,grep,glob,view_image,show_image,image_generate,ask_user,spawn_agent,queue_agent,wait_for_jobs") cmd.Flags().StringArrayVar(readDirs, "read-dir", nil, "Directories for read_file/grep/glob/view_image tools (repeatable)") cmd.Flags().StringArrayVar(writeDirs, "write-dir", nil, "Directories for write_file/edit_file tools (repeatable)") cmd.Flags().StringArrayVar(shellAllow, "shell-allow", nil, "Shell command patterns to allow (repeatable, glob syntax)") diff --git a/cmd/gateway.go b/cmd/gateway.go new file mode 100644 index 000000000..256f34f9a --- /dev/null +++ b/cmd/gateway.go @@ -0,0 +1,459 @@ +package cmd + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "os" + "path/filepath" + "strings" + "time" + + "github.com/samsaffron/term-llm/internal/config" + "github.com/samsaffron/term-llm/internal/gateway" + "github.com/samsaffron/term-llm/internal/gateway/protocol" + "github.com/samsaffron/term-llm/internal/search" + "github.com/spf13/cobra" + "gopkg.in/yaml.v3" +) + +var ( + gatewayStateDir string + gatewayListen string + gatewayTLSCert string + gatewayTLSKey string + gatewayAllowProviders []string + gatewayDenyProviders []string + gatewayAllowModels []string + gatewayDenyModels []string + gatewayAllowCLI bool + gatewayNoSearch bool + gatewayNoFetch bool + gatewayIdleTimeout time.Duration + gatewayToolTimeout time.Duration + gatewayCatalogTTL time.Duration + gatewayRetryAttempts int + gatewayRetryElapsed time.Duration + gatewayClientAllowCLI bool + gatewayClientAllow []string + gatewayClientDeny []string + gatewayClientModels []string + gatewayClientDenyModel []string + gatewayClientSearch bool + gatewayClientFetch bool + gatewayClientInference int + gatewayClientSearchRPM int + gatewayClientSearchMax int + gatewayClientFetchRPM int + gatewayClientFetchMax int + gatewayClientEnroll bool + gatewayEnrollmentTTL time.Duration + gatewayEnrollName string + gatewayEnrollWrite bool + gatewayEnrollTokenFile string + gatewayEnrollPrintOnly bool + gatewayUsageClient string + gatewayUsageJSON bool +) + +var gatewayCmd = &cobra.Command{Use: "gateway", Short: "Serve and manage the private inference gateway"} + +var gatewayServeCmd = &cobra.Command{ + Use: "serve", + Short: "Serve the private gateway and OpenAI Responses edge", + RunE: runGatewayServe, +} + +var gatewayClientCmd = &cobra.Command{Use: "client", Short: "Manage gateway satellite credentials"} +var gatewayClientAddCmd = &cobra.Command{Use: "add NAME", Args: cobra.ExactArgs(1), Short: "Create a satellite credential or enrollment token", RunE: runGatewayClientAdd} +var gatewayClientListCmd = &cobra.Command{Use: "list", Args: cobra.NoArgs, Short: "List gateway clients", RunE: runGatewayClientList} +var gatewayClientRevokeCmd = &cobra.Command{Use: "revoke ID_OR_NAME", Args: cobra.ExactArgs(1), Short: "Revoke a satellite credential", RunE: runGatewayClientRevoke} +var gatewayEnrollCmd = &cobra.Command{Use: "enroll URL ENROLLMENT_TOKEN", Args: cobra.ExactArgs(2), Short: "Consume a one-use enrollment token and configure this satellite", RunE: runGatewayEnroll} +var gatewayHealthCmd = &cobra.Command{Use: "health URL", Args: cobra.ExactArgs(1), Short: "Check an HTTP health endpoint", RunE: runGatewayHealth} +var gatewayUsageCmd = &cobra.Command{Use: "usage", Args: cobra.NoArgs, Short: "Inspect attributed gateway inference usage", RunE: runGatewayUsage} + +func init() { + rootCmd.AddCommand(gatewayCmd) + gatewayCmd.AddCommand(gatewayServeCmd, gatewayClientCmd, gatewayEnrollCmd, gatewayHealthCmd, gatewayUsageCmd) + gatewayClientCmd.AddCommand(gatewayClientAddCmd, gatewayClientListCmd, gatewayClientRevokeCmd) + + gatewayCmd.PersistentFlags().StringVar(&gatewayStateDir, "state-dir", "", "Gateway state directory (default: config directory/gateway)") + gatewayServeCmd.Flags().StringVar(&gatewayListen, "listen", "127.0.0.1:8787", "Gateway listen address") + gatewayServeCmd.Flags().StringVar(&gatewayTLSCert, "tls-cert", "", "TLS certificate PEM (requires --tls-key)") + gatewayServeCmd.Flags().StringVar(&gatewayTLSKey, "tls-key", "", "TLS private key PEM (requires --tls-cert)") + gatewayServeCmd.Flags().StringSliceVar(&gatewayAllowProviders, "allow-provider", nil, "Allowed provider key/prefix (repeatable)") + gatewayServeCmd.Flags().StringSliceVar(&gatewayDenyProviders, "deny-provider", nil, "Denied provider key/prefix (repeatable)") + gatewayServeCmd.Flags().StringSliceVar(&gatewayAllowModels, "allow-model", nil, "Allowed provider:model/model pattern (repeatable)") + gatewayServeCmd.Flags().StringSliceVar(&gatewayDenyModels, "deny-model", nil, "Denied provider:model/model pattern (repeatable)") + gatewayServeCmd.Flags().BoolVar(&gatewayAllowCLI, "allow-cli", false, "Allow CLI providers globally (clients must also opt in)") + gatewayServeCmd.Flags().BoolVar(&gatewayNoSearch, "no-search", false, "Disable centralized gateway search") + gatewayServeCmd.Flags().BoolVar(&gatewayNoFetch, "no-fetch", false, "Disable centralized gateway fetch") + gatewayServeCmd.Flags().DurationVar(&gatewayIdleTimeout, "idle-timeout", 5*time.Minute, "Cancel provider streams with no events for this duration") + gatewayServeCmd.Flags().DurationVar(&gatewayToolTimeout, "tool-timeout", 10*time.Minute, "Maximum wait for a satellite tool callback") + gatewayServeCmd.Flags().DurationVar(&gatewayCatalogTTL, "catalog-ttl", 5*time.Minute, "Refresh provider config and live model catalogs after this interval") + gatewayServeCmd.Flags().IntVar(&gatewayRetryAttempts, "upstream-retry-attempts", gateway.DefaultUpstreamRetryAttempts, "Maximum upstream attempts per gateway inference request") + gatewayServeCmd.Flags().DurationVar(&gatewayRetryElapsed, "upstream-retry-elapsed", gateway.DefaultUpstreamRetryElapsed, "Maximum elapsed time across upstream attempts") + + gatewayClientAddCmd.Flags().BoolVar(&gatewayClientAllowCLI, "allow-cli", false, "Allow this client to use CLI providers") + gatewayClientAddCmd.Flags().StringSliceVar(&gatewayClientAllow, "allow-provider", nil, "Allowed provider key/prefix") + gatewayClientAddCmd.Flags().StringSliceVar(&gatewayClientDeny, "deny-provider", nil, "Denied provider key/prefix") + gatewayClientAddCmd.Flags().StringSliceVar(&gatewayClientModels, "allow-model", nil, "Allowed model/provider:model pattern") + gatewayClientAddCmd.Flags().StringSliceVar(&gatewayClientDenyModel, "deny-model", nil, "Denied model/provider:model pattern") + gatewayClientAddCmd.Flags().BoolVar(&gatewayClientSearch, "allow-search", false, "Allow centralized search for this client") + gatewayClientAddCmd.Flags().BoolVar(&gatewayClientFetch, "allow-fetch", false, "Allow centralized fetch for this client") + gatewayClientAddCmd.Flags().IntVar(&gatewayClientInference, "max-concurrent-inference", gateway.DefaultMaxConcurrentInference, "Maximum concurrent inference requests for this client") + gatewayClientAddCmd.Flags().IntVar(&gatewayClientSearchRPM, "search-rate", gateway.DefaultSearchRatePerMinute, "Maximum search requests per minute") + gatewayClientAddCmd.Flags().IntVar(&gatewayClientSearchMax, "max-concurrent-search", gateway.DefaultMaxConcurrentSearch, "Maximum concurrent search requests") + gatewayClientAddCmd.Flags().IntVar(&gatewayClientFetchRPM, "fetch-rate", gateway.DefaultFetchRatePerMinute, "Maximum fetch requests per minute") + gatewayClientAddCmd.Flags().IntVar(&gatewayClientFetchMax, "max-concurrent-fetch", gateway.DefaultMaxConcurrentFetch, "Maximum concurrent fetch requests") + gatewayClientAddCmd.Flags().BoolVar(&gatewayClientEnroll, "enroll", false, "Generate a persisted one-use enrollment token instead of a client token") + gatewayClientAddCmd.Flags().DurationVar(&gatewayEnrollmentTTL, "enroll-ttl", gateway.DefaultEnrollmentTTL, "Enrollment token lifetime (maximum 24h)") + gatewayEnrollCmd.Flags().StringVar(&gatewayEnrollName, "name", "", "Satellite name (must match the enrollment token; default hostname)") + gatewayEnrollCmd.Flags().BoolVar(&gatewayEnrollWrite, "write-config", true, "Atomically update satellite config and write a separate 0600 token file") + gatewayEnrollCmd.Flags().StringVar(&gatewayEnrollTokenFile, "token-file", "", "Token file path (default: config directory/gateway-token)") + gatewayEnrollCmd.Flags().BoolVar(&gatewayEnrollPrintOnly, "print-only", false, "Print config including the client token instead of writing files") + gatewayUsageCmd.Flags().StringVar(&gatewayUsageClient, "client", "", "Filter by client ID or name") + gatewayUsageCmd.Flags().BoolVar(&gatewayUsageJSON, "json", false, "Output usage records as JSON") +} + +func runGatewayServe(cmd *cobra.Command, _ []string) error { + cfg, err := loadConfigWithSetup() + if err != nil { + return err + } + if (strings.TrimSpace(gatewayTLSCert) == "") != (strings.TrimSpace(gatewayTLSKey) == "") { + return fmt.Errorf("--tls-cert and --tls-key must be provided together") + } + if gatewayRetryAttempts <= 0 { + return fmt.Errorf("--upstream-retry-attempts must be positive") + } + if gatewayRetryElapsed <= 0 { + return fmt.Errorf("--upstream-retry-elapsed must be positive") + } + stateDir, err := resolveGatewayStateDir() + if err != nil { + return err + } + clients, err := gateway.OpenClientStore(filepath.Join(stateDir, "clients.json")) + if err != nil { + return err + } + sealer, err := gateway.OpenStateSealer(filepath.Join(stateDir, "state.key")) + if err != nil { + return err + } + central := *cfg + central.Gateway = config.GatewayConfig{} + var searcher search.Searcher + if !gatewayNoSearch { + searcher, err = search.NewSearcher(¢ral) + if err != nil { + return fmt.Errorf("configure gateway search: %w", err) + } + } + fetchTool := newReadURLToolForConfig(¢ral) + if gatewayNoFetch { + fetchTool = nil + } + server, err := gateway.NewServer(gateway.ServerConfig{ + Config: ¢ral, + ConfigLoader: func() (*config.Config, error) { + loaded, loadErr := config.Load() + if loadErr != nil { + return nil, loadErr + } + loaded.Gateway = config.GatewayConfig{} + return loaded, nil + }, + Clients: clients, Sealer: sealer, + Usage: &gateway.JSONLUsageRecorder{Path: filepath.Join(stateDir, "usage.jsonl")}, + Searcher: searcher, FetchTool: fetchTool, + IdleTimeout: gatewayIdleTimeout, ToolTimeout: gatewayToolTimeout, CatalogTTL: gatewayCatalogTTL, + UpstreamRetryAttempts: gatewayRetryAttempts, UpstreamRetryMaxElapsed: gatewayRetryElapsed, + RunTempRoot: filepath.Join(stateDir, "runs"), + Policy: gateway.Policy{ + AllowProviders: gatewayAllowProviders, DenyProviders: gatewayDenyProviders, + AllowModels: gatewayAllowModels, DenyModels: gatewayDenyModels, + AllowCLI: gatewayAllowCLI, AllowSearch: !gatewayNoSearch, AllowFetch: !gatewayNoFetch, + }, + }) + if err != nil { + return err + } + httpServer := &http.Server{Addr: gatewayListen, Handler: server.Handler(), ReadHeaderTimeout: 10 * time.Second, IdleTimeout: 2 * time.Minute} + go func() { + <-cmd.Context().Done() + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + _ = httpServer.Shutdown(ctx) + }() + scheme := "http" + if gatewayTLSCert != "" { + scheme = "https" + } + fmt.Fprintf(cmd.OutOrStdout(), "Inference gateway listening on %s://%s (/g1 and /v1 Responses)\n", scheme, gatewayListen) + fmt.Fprintln(cmd.OutOrStdout(), "Create one-use enrollment tokens with `term-llm gateway client add NAME --enroll --allow-provider PROVIDER`.") + if gatewayTLSCert != "" { + err = httpServer.ListenAndServeTLS(gatewayTLSCert, gatewayTLSKey) + } else { + err = httpServer.ListenAndServe() + } + if errors.Is(err, http.ErrServerClosed) { + return nil + } + return err +} + +func gatewayClientPolicy() gateway.Policy { + return gateway.Policy{ + AllowProviders: gatewayClientAllow, DenyProviders: gatewayClientDeny, + AllowModels: gatewayClientModels, DenyModels: gatewayClientDenyModel, + AllowCLI: gatewayClientAllowCLI, AllowSearch: gatewayClientSearch, AllowFetch: gatewayClientFetch, + MaxConcurrentInference: gatewayClientInference, + SearchRatePerMinute: gatewayClientSearchRPM, SearchBurst: min(gatewayClientSearchRPM, gateway.DefaultSearchBurst), MaxConcurrentSearch: gatewayClientSearchMax, + FetchRatePerMinute: gatewayClientFetchRPM, FetchBurst: min(gatewayClientFetchRPM, gateway.DefaultFetchBurst), MaxConcurrentFetch: gatewayClientFetchMax, + } +} + +func runGatewayClientAdd(cmd *cobra.Command, args []string) error { + store, err := gatewayClientStore() + if err != nil { + return err + } + policy := gatewayClientPolicy() + if gatewayClientEnroll { + enrollment, token, err := store.CreateEnrollment(args[0], policy, gatewayEnrollmentTTL) + if err != nil { + return err + } + fmt.Fprintf(cmd.OutOrStdout(), "Enrollment token (shown once, expires %s): %s\n", enrollment.ExpiresAt.Format(time.RFC3339), token) + fmt.Fprintf(cmd.OutOrStdout(), "Satellite: term-llm gateway enroll https://gateway.example %s --name %s\n", token, enrollment.Name) + return nil + } + client, token, err := store.Add(args[0], policy) + if err != nil { + return err + } + fmt.Fprintf(cmd.OutOrStdout(), "Client: %s (%s)\nToken (shown once): %s\n", client.Name, client.ID, token) + fmt.Fprintln(cmd.OutOrStdout(), "\nSatellite config:\ngateway:\n url: https://gateway.example\n token: "+token) + return nil +} + +func runGatewayClientList(cmd *cobra.Command, _ []string) error { + store, err := gatewayClientStore() + if err != nil { + return err + } + for _, client := range store.List() { + status := "active" + if !client.RevokedAt.IsZero() { + status = "revoked" + } + fmt.Fprintf(cmd.OutOrStdout(), "%s\t%s\t%s\tcli=%t\tinference=%d\tsearch=%t\tfetch=%t\n", client.ID, client.Name, status, client.Policy.AllowCLI, client.Policy.InferenceConcurrency(), client.Policy.AllowSearch, client.Policy.AllowFetch) + } + return nil +} + +func runGatewayClientRevoke(_ *cobra.Command, args []string) error { + store, err := gatewayClientStore() + if err != nil { + return err + } + return store.Revoke(args[0]) +} + +func runGatewayEnroll(cmd *cobra.Command, args []string) error { + if !gatewayEnrollPrintOnly && !gatewayEnrollWrite && strings.TrimSpace(gatewayEnrollTokenFile) == "" { + return fmt.Errorf("select a credential destination with --write-config, --token-file, or --print-only") + } + name := strings.TrimSpace(gatewayEnrollName) + if name == "" { + name, _ = os.Hostname() + } + payload, _ := json.Marshal(protocol.EnrollmentRequest{Version: protocol.Version, Name: name}) + url := strings.TrimRight(args[0], "/") + "/g1/enroll" + req, err := http.NewRequestWithContext(cmd.Context(), http.MethodPost, url, bytes.NewReader(payload)) + if err != nil { + return err + } + req.Header.Set("Authorization", "Bearer "+args[1]) + req.Header.Set("Content-Type", "application/json") + resp, err := (&http.Client{Timeout: 10 * time.Second}).Do(req) + if err != nil { + return fmt.Errorf("enroll with gateway: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusCreated { + var wire protocol.Error + _ = json.NewDecoder(resp.Body).Decode(&wire) + if wire.Message != "" { + return fmt.Errorf("gateway enrollment failed: %s", wire.Message) + } + return fmt.Errorf("gateway enrollment failed: HTTP %d", resp.StatusCode) + } + var enrolled protocol.EnrollmentResponse + if err := json.NewDecoder(resp.Body).Decode(&enrolled); err != nil { + return err + } + if gatewayEnrollPrintOnly { + fmt.Fprintf(cmd.OutOrStdout(), "gateway:\n url: %s\n token: %s\n", strings.TrimRight(args[0], "/"), enrolled.Token) + return nil + } + tokenPath := strings.TrimSpace(gatewayEnrollTokenFile) + if tokenPath == "" { + configDir, pathErr := config.GetConfigDir() + if pathErr != nil { + return pathErr + } + tokenPath = filepath.Join(configDir, "gateway-token") + } + tokenPath, err = expandGatewayEnrollPath(tokenPath) + if err != nil { + return err + } + if err := os.MkdirAll(filepath.Dir(tokenPath), 0o700); err != nil { + return fmt.Errorf("create gateway token directory: %w", err) + } + if err := config.WriteFileAtomically(tokenPath, []byte(enrolled.Token+"\n"), 0o600); err != nil { + return fmt.Errorf("write gateway token: %w", err) + } + if err := os.Chmod(tokenPath, 0o600); err != nil { + return fmt.Errorf("secure gateway token: %w", err) + } + if !gatewayEnrollWrite { + fmt.Fprintf(cmd.OutOrStdout(), "Gateway token written to %s (mode 0600).\n", tokenPath) + return nil + } + configPath, err := config.GetConfigPath() + if err != nil { + return err + } + if err := writeGatewaySatelliteConfig(configPath, strings.TrimRight(args[0], "/"), tokenPath); err != nil { + return err + } + fmt.Fprintf(cmd.OutOrStdout(), "Gateway enrollment complete.\nConfig: %s\nToken: %s (mode 0600)\n", configPath, tokenPath) + return nil +} + +func expandGatewayEnrollPath(path string) (string, error) { + path = strings.TrimSpace(path) + if strings.HasPrefix(path, "~/") { + home, err := os.UserHomeDir() + if err != nil { + return "", fmt.Errorf("resolve token file: %w", err) + } + path = filepath.Join(home, strings.TrimPrefix(path, "~/")) + } + abs, err := filepath.Abs(path) + if err != nil { + return "", fmt.Errorf("resolve token file: %w", err) + } + return filepath.Clean(abs), nil +} + +func writeGatewaySatelliteConfig(path, gatewayURL, tokenPath string) error { + if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil { + return fmt.Errorf("create config directory: %w", err) + } + root := make(map[string]any) + if data, err := os.ReadFile(path); err == nil { + if err := yaml.Unmarshal(data, &root); err != nil { + return fmt.Errorf("parse existing config: %w", err) + } + } else if !os.IsNotExist(err) { + return fmt.Errorf("read existing config: %w", err) + } + gatewayMap := make(map[string]any) + if existing, ok := root["gateway"].(map[string]any); ok { + for key, value := range existing { + gatewayMap[key] = value + } + } + gatewayMap["url"] = gatewayURL + gatewayMap["token_file"] = tokenPath + delete(gatewayMap, "token") + delete(gatewayMap, "token_env") + root["gateway"] = gatewayMap + data, err := yaml.Marshal(root) + if err != nil { + return fmt.Errorf("encode satellite config: %w", err) + } + if err := config.WriteFileAtomically(path, data, 0o600); err != nil { + return fmt.Errorf("write satellite config: %w", err) + } + return nil +} + +func runGatewayHealth(cmd *cobra.Command, args []string) error { + url := strings.TrimSpace(args[0]) + req, err := http.NewRequestWithContext(cmd.Context(), http.MethodGet, url, nil) + if err != nil { + return fmt.Errorf("create health request: %w", err) + } + resp, err := (&http.Client{Timeout: 5 * time.Second}).Do(req) + if err != nil { + return fmt.Errorf("health request failed: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + return fmt.Errorf("health endpoint returned HTTP %d", resp.StatusCode) + } + fmt.Fprintln(cmd.OutOrStdout(), "ok") + return nil +} + +func runGatewayUsage(cmd *cobra.Command, _ []string) error { + stateDir, err := resolveGatewayStateDir() + if err != nil { + return err + } + records, err := gateway.ReadUsageRecords(filepath.Join(stateDir, "usage.jsonl")) + if err != nil { + return err + } + filtered := records[:0] + for _, record := range records { + if gatewayUsageClient == "" || record.ClientID == gatewayUsageClient || record.ClientName == gatewayUsageClient { + filtered = append(filtered, record) + } + } + if gatewayUsageJSON { + encoder := json.NewEncoder(cmd.OutOrStdout()) + encoder.SetIndent("", " ") + return encoder.Encode(filtered) + } + var input, output int + var cost float64 + for _, record := range filtered { + input += record.InputTokens + record.CachedInputTokens + record.CacheWriteTokens + output += record.OutputTokens + if record.CostUSD != nil { + cost += *record.CostUSD + } + fmt.Fprintf(cmd.OutOrStdout(), "%s\t%s\t%s\t%s:%s\tin=%d\tout=%d\terror=%s\n", record.CompletedAt.UTC().Format(time.RFC3339), record.ClientName, record.RequestID, record.ProviderKey, record.Model, record.InputTokens, record.OutputTokens, record.ErrorCode) + } + fmt.Fprintf(cmd.OutOrStdout(), "Total\trequests=%d\tinput=%d\toutput=%d\tcost=$%.4f\n", len(filtered), input, output, cost) + return nil +} + +func gatewayClientStore() (*gateway.ClientStore, error) { + stateDir, err := resolveGatewayStateDir() + if err != nil { + return nil, err + } + return gateway.OpenClientStore(filepath.Join(stateDir, "clients.json")) +} + +func resolveGatewayStateDir() (string, error) { + if strings.TrimSpace(gatewayStateDir) != "" { + return filepath.Clean(gatewayStateDir), nil + } + configDir, err := config.GetConfigDir() + if err != nil { + return "", err + } + return filepath.Join(configDir, "gateway"), nil +} diff --git a/cmd/gateway_enroll_test.go b/cmd/gateway_enroll_test.go new file mode 100644 index 000000000..26aed9ee0 --- /dev/null +++ b/cmd/gateway_enroll_test.go @@ -0,0 +1,119 @@ +package cmd + +import ( + "bytes" + "context" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/samsaffron/term-llm/internal/config" + "github.com/samsaffron/term-llm/internal/gateway" + "github.com/spf13/viper" +) + +func TestGatewayEnrollWritesParseableConfigAndSecureTokenByDefault(t *testing.T) { + viper.Reset() + t.Cleanup(viper.Reset) + configHome := t.TempDir() + t.Setenv("XDG_CONFIG_HOME", configHome) + + stateDir := t.TempDir() + store, err := gateway.OpenClientStore(filepath.Join(stateDir, "clients.json")) + if err != nil { + t.Fatal(err) + } + _, bootstrap, err := store.CreateEnrollment("satellite-test", gateway.Policy{AllowProviders: []string{"debug"}}, time.Minute) + if err != nil { + t.Fatal(err) + } + sealer, err := gateway.OpenStateSealer(filepath.Join(stateDir, "state.key")) + if err != nil { + t.Fatal(err) + } + server, err := gateway.NewServer(gateway.ServerConfig{Config: &config.Config{Providers: map[string]config.ProviderConfig{}}, Clients: store, Sealer: sealer}) + if err != nil { + t.Fatal(err) + } + ts := httptest.NewServer(server.Handler()) + defer ts.Close() + + oldName, oldWrite, oldTokenFile, oldPrint := gatewayEnrollName, gatewayEnrollWrite, gatewayEnrollTokenFile, gatewayEnrollPrintOnly + t.Cleanup(func() { + gatewayEnrollName, gatewayEnrollWrite, gatewayEnrollTokenFile, gatewayEnrollPrintOnly = oldName, oldWrite, oldTokenFile, oldPrint + }) + gatewayEnrollName = "satellite-test" + gatewayEnrollWrite = true + gatewayEnrollTokenFile = "" + gatewayEnrollPrintOnly = false + var output bytes.Buffer + gatewayEnrollCmd.SetOut(&output) + gatewayEnrollCmd.SetContext(t.Context()) + t.Cleanup(func() { + gatewayEnrollCmd.SetOut(nil) + gatewayEnrollCmd.SetContext(context.Background()) + }) + if err := runGatewayEnroll(gatewayEnrollCmd, []string{ts.URL, bootstrap}); err != nil { + t.Fatal(err) + } + if strings.Contains(output.String(), "tlg1_") { + t.Fatalf("default enrollment printed client token: %q", output.String()) + } + + tokenPath := filepath.Join(configHome, "term-llm", "gateway-token") + info, err := os.Stat(tokenPath) + if err != nil { + t.Fatal(err) + } + if info.Mode().Perm() != 0o600 { + t.Fatalf("token mode = %o, want 600", info.Mode().Perm()) + } + cfg, err := config.Load() + if err != nil { + t.Fatal(err) + } + if cfg.Gateway.URL != ts.URL || cfg.Gateway.TokenFile != tokenPath || cfg.Gateway.Token != "" { + t.Fatalf("enrolled gateway config = %+v", cfg.Gateway) + } + token, err := cfg.Gateway.ResolveToken() + if err != nil || !strings.HasPrefix(token, "tlg1_") { + t.Fatalf("resolved enrolled token = %q, %v", token, err) + } +} + +func TestGatewayEnrollPrintOnlyIsExplicitAndDoesNotWrite(t *testing.T) { + configHome := t.TempDir() + t.Setenv("XDG_CONFIG_HOME", configHome) + stateDir := t.TempDir() + store, _ := gateway.OpenClientStore(filepath.Join(stateDir, "clients.json")) + _, bootstrap, _ := store.CreateEnrollment("print-test", gateway.Policy{AllowProviders: []string{"debug"}}, time.Minute) + sealer, _ := gateway.OpenStateSealer(filepath.Join(stateDir, "state.key")) + server, _ := gateway.NewServer(gateway.ServerConfig{Config: &config.Config{Providers: map[string]config.ProviderConfig{}}, Clients: store, Sealer: sealer}) + ts := httptest.NewServer(server.Handler()) + defer ts.Close() + + oldName, oldWrite, oldTokenFile, oldPrint := gatewayEnrollName, gatewayEnrollWrite, gatewayEnrollTokenFile, gatewayEnrollPrintOnly + t.Cleanup(func() { + gatewayEnrollName, gatewayEnrollWrite, gatewayEnrollTokenFile, gatewayEnrollPrintOnly = oldName, oldWrite, oldTokenFile, oldPrint + }) + gatewayEnrollName, gatewayEnrollWrite, gatewayEnrollTokenFile, gatewayEnrollPrintOnly = "print-test", true, "", true + var output bytes.Buffer + gatewayEnrollCmd.SetOut(&output) + gatewayEnrollCmd.SetContext(t.Context()) + t.Cleanup(func() { + gatewayEnrollCmd.SetOut(nil) + gatewayEnrollCmd.SetContext(context.Background()) + }) + if err := runGatewayEnroll(gatewayEnrollCmd, []string{ts.URL, bootstrap}); err != nil { + t.Fatal(err) + } + if !strings.Contains(output.String(), "token: tlg1_") { + t.Fatalf("print-only output omitted token: %q", output.String()) + } + if _, err := os.Stat(filepath.Join(configHome, "term-llm", "config.yaml")); !os.IsNotExist(err) { + t.Fatalf("print-only wrote config: %v", err) + } +} diff --git a/cmd/gateway_tools_test.go b/cmd/gateway_tools_test.go new file mode 100644 index 000000000..9603765cc --- /dev/null +++ b/cmd/gateway_tools_test.go @@ -0,0 +1,28 @@ +package cmd + +import ( + "context" + "strings" + "testing" + + "github.com/samsaffron/term-llm/internal/config" +) + +func TestRequiredGatewayFetchKeepsLegibleReadURLStub(t *testing.T) { + cfg := &config.Config{Gateway: config.GatewayConfig{URL: "https://gateway.invalid", Required: true}} + if tool := newReadURLToolForConfig(cfg); tool == nil { + t.Fatal("gateway outage silently removed read_url") + } + _, err := (unavailableGatewayFetcher{err: context.DeadlineExceeded}).FetchURL(context.Background(), "https://example.com") + if err == nil || !strings.Contains(err.Error(), "gateway read_url unavailable") || !strings.Contains(err.Error(), "gateway.fetch: false") { + t.Fatalf("gateway read_url stub error = %v", err) + } +} + +func TestGatewaySearchFailureDoesNotSilentlyFallBack(t *testing.T) { + searcher := unavailableGatewaySearcher{err: context.DeadlineExceeded} + _, err := searcher.Search(context.Background(), "query", 10) + if err == nil || !strings.Contains(err.Error(), "gateway search unavailable") || !strings.Contains(err.Error(), "gateway.search: false") { + t.Fatalf("gateway search stub error = %v", err) + } +} diff --git a/cmd/gateway_usage_test.go b/cmd/gateway_usage_test.go new file mode 100644 index 000000000..b83267018 --- /dev/null +++ b/cmd/gateway_usage_test.go @@ -0,0 +1,38 @@ +package cmd + +import ( + "bytes" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/samsaffron/term-llm/internal/gateway" +) + +func TestGatewayUsageCommandReadsAttributedRecords(t *testing.T) { + oldStateDir, oldClient, oldJSON := gatewayStateDir, gatewayUsageClient, gatewayUsageJSON + t.Cleanup(func() { + gatewayStateDir, gatewayUsageClient, gatewayUsageJSON = oldStateDir, oldClient, oldJSON + }) + gatewayStateDir = t.TempDir() + gatewayUsageClient = "satellite-a" + gatewayUsageJSON = false + recorder := &gateway.JSONLUsageRecorder{Path: filepath.Join(gatewayStateDir, "usage.jsonl")} + if err := recorder.Record(gateway.UsageRecord{ + StartedAt: time.Now().Add(-time.Second), CompletedAt: time.Now(), ClientID: "client-a", ClientName: "satellite-a", + ProviderKey: "openai", Model: "gpt", RequestID: "req-1", InputTokens: 10, OutputTokens: 2, + }); err != nil { + t.Fatal(err) + } + var output bytes.Buffer + gatewayUsageCmd.SetOut(&output) + t.Cleanup(func() { gatewayUsageCmd.SetOut(nil) }) + if err := runGatewayUsage(gatewayUsageCmd, nil); err != nil { + t.Fatal(err) + } + text := output.String() + if !strings.Contains(text, "satellite-a") || !strings.Contains(text, "openai:gpt") || !strings.Contains(text, "requests=1") { + t.Fatalf("gateway usage output = %q", text) + } +} diff --git a/cmd/models.go b/cmd/models.go index 1da777bac..0083909f0 100644 --- a/cmd/models.go +++ b/cmd/models.go @@ -76,6 +76,20 @@ func runModels(cmd *cobra.Command, args []string) error { if providerName == "" { providerName = cfg.DefaultProvider } + if cfg.Gateway.Enabled() && !cfg.IsLocalProvider(providerName) { + provider, routeErr := llm.NewProviderByName(cfg, providerName, "") + if routeErr != nil { + return routeErr + } else if remote, ok := provider.(*llm.GatewayProvider); ok { + ctx, cancel := context.WithTimeout(cmd.Context(), 30*time.Second) + defer cancel() + models, listErr := remote.ListModels(ctx) + if listErr != nil { + return fmt.Errorf("failed to list gateway models: %w", listErr) + } + return outputListedModels(providerName, models, true) + } + } // Get provider config - handle built-in providers that may not be explicitly configured providerCfg, ok := cfg.Providers[providerName] @@ -225,6 +239,12 @@ func runModels(cmd *cobra.Command, args []string) error { return fmt.Errorf("failed to list models: %w", err) } + // Only these providers return or have known pricing info. + providerHasPricing := providerType == config.ProviderTypeOpenRouter || providerType == config.ProviderTypeZen || providerType == config.ProviderTypeNearAI || providerType == config.ProviderTypeSambaNova + return outputListedModels(providerName, models, providerHasPricing) +} + +func outputListedModels(providerName string, models []llm.ModelInfo, providerHasPricing bool) error { if len(models) == 0 { fmt.Println("No models found.") return nil @@ -236,12 +256,7 @@ func runModels(cmd *cobra.Command, args []string) error { return enc.Encode(models) } - // Pretty print fmt.Printf("Available models from %s:\n\n", providerName) - - // Only these providers return or have known pricing info - providerHasPricing := providerType == config.ProviderTypeOpenRouter || providerType == config.ProviderTypeZen || providerType == config.ProviderTypeNearAI || providerType == config.ProviderTypeSambaNova - for _, m := range models { if m.DisplayName != "" { fmt.Printf(" %s (%s)", m.ID, m.DisplayName) @@ -255,10 +270,8 @@ func runModels(cmd *cobra.Command, args []string) error { } // Show pricing info only if provider returns it - if providerHasPricing { - if m.InputPrice < 0 || m.OutputPrice < 0 { - fmt.Printf(" [pricing unknown]") - } else if m.InputPrice == 0 && m.OutputPrice == 0 { + if providerHasPricing && m.InputPrice >= 0 && m.OutputPrice >= 0 { + if m.InputPrice == 0 && m.OutputPrice == 0 { fmt.Printf(" [FREE]") } else { fmt.Printf(" [$%.2f/$%.2f per 1M tokens]", m.InputPrice, m.OutputPrice) diff --git a/cmd/runner.go b/cmd/runner.go index d29e39c38..1cb1f9f12 100644 --- a/cmd/runner.go +++ b/cmd/runner.go @@ -188,8 +188,18 @@ func (r *cmdRunner) prepare(ctx context.Context, req runpkg.Request, sink runpkg return nil, err } if model := strings.TrimSpace(req.Model); model != "" { - if err := applyAgentModelOverride(cfg, model); err != nil { - return nil, fmt.Errorf("apply model override %q: %w", model, err) + // A provider:model CLI selection is already concrete. In particular, + // debug:fast means the literal gateway catalog model "fast", not the + // special agent-level fast-model alias. + parts := strings.SplitN(providerFlag, ":", 2) + explicitModel := "" + if len(parts) == 2 { + explicitModel = strings.TrimSpace(parts[1]) + } + if strings.TrimSpace(providerFlag) == "" || explicitModel == "" || explicitModel != model { + if err := applyAgentModelOverride(cfg, model); err != nil { + return nil, fmt.Errorf("apply model override %q: %w", model, err) + } } } diff --git a/cmd/tools.go b/cmd/tools.go index cdc5291e4..e031e3566 100644 --- a/cmd/tools.go +++ b/cmd/tools.go @@ -1,6 +1,8 @@ package cmd import ( + "context" + "fmt" "log" "github.com/samsaffron/term-llm/internal/config" @@ -9,12 +11,29 @@ import ( "github.com/samsaffron/term-llm/internal/tools" ) +type unavailableGatewaySearcher struct{ err error } + +func (s unavailableGatewaySearcher) Search(context.Context, string, int) ([]search.Result, error) { + return nil, fmt.Errorf("gateway search unavailable: %w; check gateway URL/network/token or set gateway.search: false to use local search", s.err) +} + +type unavailableGatewayFetcher struct{ err error } + +func (f unavailableGatewayFetcher) FetchURL(context.Context, string) (string, error) { + return "", fmt.Errorf("gateway read_url unavailable: %w; check gateway URL/network/token or set gateway.fetch: false to use local fetch", f.err) +} + func defaultToolRegistry(cfg *config.Config) *llm.ToolRegistry { registry := llm.NewToolRegistry() searcher, err := search.NewSearcher(cfg) if err != nil { - log.Printf("Warning: search provider error: %v, falling back to DuckDuckGo", err) - searcher = search.NewDuckDuckGoLite(nil) + if cfg != nil && cfg.Gateway.Enabled() && cfg.Gateway.RouteSearch() { + log.Printf("Warning: gateway search unavailable: %v", err) + searcher = unavailableGatewaySearcher{err: err} + } else { + log.Printf("Warning: search provider error: %v, falling back to DuckDuckGo", err) + searcher = search.NewDuckDuckGoLite(nil) + } } registry.Register(llm.NewWebSearchTool(searcher)) if readURLTool := newReadURLToolForConfig(cfg); readURLTool != nil { @@ -24,6 +43,14 @@ func defaultToolRegistry(cfg *config.Config) *llm.ToolRegistry { } func newReadURLToolForConfig(cfg *config.Config) *llm.ReadURLTool { + if cfg.Gateway.Enabled() && cfg.Gateway.RouteFetch() { + client, err := search.NewGatewayClient(cfg.Gateway) + if err != nil { + log.Printf("Warning: gateway fetch unavailable: %v", err) + return llm.NewReadURLToolWithFetcher(unavailableGatewayFetcher{err: err}) + } + return llm.NewReadURLToolWithFetcher(client) + } switch cfg.Search.FetchProvider { case "", "jina": return llm.NewReadURLTool() diff --git a/cmd/usage.go b/cmd/usage.go index 1666ce42e..f0b4f9d6d 100644 --- a/cmd/usage.go +++ b/cmd/usage.go @@ -67,7 +67,7 @@ func init() { usageCmd.Flags().StringVar(&usageUntil, "until", "", "End date (YYYYMMDD)") usageCmd.Flags().BoolVar(&usageJSON, "json", false, "Output as JSON") usageCmd.Flags().BoolVar(&usageBreakdown, "breakdown", false, "Show per-model breakdown") - usageCmd.Flags().BoolVar(&usageIncludeExternal, "include-external", false, "Include externally-tracked term-llm usage (claude-bin, codex, gemini-cli calls)") + usageCmd.Flags().BoolVar(&usageIncludeExternal, "include-external", false, "Include externally-tracked term-llm usage (CLI-provider and gateway calls) in any provider view") usageCmd.Flags().StringVar(&usageCopilotScope, "copilot-scope", "user", "Copilot billing scope (user, org, enterprise)") usageCmd.Flags().StringVar(&usageCopilotEntity, "copilot-entity", "", "Copilot billing entity (username, organization, or enterprise slug; defaults to authenticated user for user scope)") usageCmd.Flags().IntVar(&usageCopilotYear, "year", 0, "Copilot usage year (YYYY; defaults to GitHub API current year)") diff --git a/docs-site/content/guides/inference-gateway.md b/docs-site/content/guides/inference-gateway.md new file mode 100644 index 000000000..0fefaf4bd --- /dev/null +++ b/docs-site/content/guides/inference-gateway.md @@ -0,0 +1,240 @@ +--- +title: "Inference Gateway" +weight: 18 +description: "Centralize provider credentials and web access while keeping agents, tools, approvals, and files on satellites." +kicker: "Private provider plane" +--- + +The inference gateway centralizes LLM provider credentials, OAuth/CLI homes, model discovery, web search/fetch credentials, and attributed usage. It is a remote **provider**, not a remote agent. + +A satellite still owns: + +- the engine and agent/system prompt +- session history, compaction, jobs, memory, and UI +- the local tool registry and approval policy +- its working directory and filesystem + +The gateway owns provider transport only. Satellites continue using ordinary `provider:model` selections; catalog providers are routed through the gateway automatically. + +## Protocol and trust boundary + +Gateway traffic has two authenticated provider edges: + +- the private, versioned `/g1` HTTP/SSE protocol used by term-llm satellites +- an OpenAI-compatible `POST /v1/responses` and `GET /v1/models` edge for Discourse and other provider clients + +Neither edge is a full-agent runtime. `/v1/responses` translates directly to the same provider-neutral request/event stream used by `/g1`; it does not proxy through `/g1`, start `term-llm serve`, load sessions/skills/jobs/memory, or create a gateway-side tool registry. The private protocol additionally supports ETagged catalog discovery, cancellation, authenticated synchronous satellite tool callbacks, sealed provider state, enrollment, search, fetch, and health. + +Important invariants: + +1. Provider API keys, OAuth tokens, CLI homes, and provider configuration are never returned by catalog or inference endpoints. +2. A satellite `WorkingDir` is not a wire field. Each gateway request gets a new empty directory under the gateway state-owned `runs/` root. Finished runs are removed, and stale gateway-prefixed directories are scavenged safely at startup. +3. The gateway has no satellite or external-client tool registry. Normal `/g1` tool calls return to the satellite engine. Inline CLI-provider calls on `/g1` use an authenticated callback POST and block the gateway provider until the satellite result arrives. `/v1/responses` accepts only client-defined function tools, returns function calls to the client, and never executes them; tool-bearing requests to incompatible inline-loop CLI providers are rejected before provider startup. +4. Provider resume state is opaque to satellites and authenticated with an AES-GCM gateway key. It is bound to both client ID and provider key; tampered or cross-client state is rejected. +5. Each request is recorded centrally with client, provider key, model, request/session IDs, token counters, outcome, and locally calculable cost. +6. The gateway performs provider retries. By default it makes at most three upstream attempts within 20 seconds; `--upstream-retry-attempts` and `--upstream-retry-elapsed` tighten or extend those bounds. Request cancellation and stream idle deadlines still win. The satellite gateway transport and engine do not add a second retry loop. + +Run the gateway only on a trusted private network. For traffic that is not already protected by a private overlay, service mesh, or TLS reverse proxy, configure the built-in TLS listener with both `--tls-cert` and `--tls-key`. Bearer credentials authenticate access but plaintext HTTP does not encrypt them. + +## Start a gateway + +Use the normal central term-llm config for provider credentials. Authenticate interactive providers on the gateway host before starting the service; HTTP handlers never initiate an interactive OAuth flow. + +Create a client: + +```bash +term-llm gateway client add jarvis \ + --allow-provider anthropic \ + --allow-provider openai \ + --allow-search \ + --max-concurrent-inference 2 +``` + +The token is shown once. Client records contain only a SHA-256 token hash. State defaults to `~/.config/term-llm/gateway`; override it for a service volume: + +```bash +term-llm gateway serve \ + --listen 0.0.0.0:8787 \ + --state-dir /var/lib/term-llm-gateway +``` + +Gateway state contains: + +```text +clients.json # client IDs, token hashes, policy, revocation +clients.enrollments.json # hashed one-use enrollment tokens, expiry, policy, use time +state.key # 32-byte provider-state sealing key +usage.jsonl # attributed request usage/outcomes +runs/ # state-owned ephemeral run directories; scavenged at startup +``` + +Back up `state.key` with the client database. Losing it safely invalidates existing opaque provider resume state but does not expose provider credentials. + +Manage clients: + +```bash +term-llm gateway client list +term-llm gateway client revoke jarvis +term-llm gateway usage --client jarvis +``` + +Client additions and revocations are written atomically. Active client names are unique so name-based revocation is unambiguous. To rotate a credential, revoke the existing ID or name, then run `gateway client add` again with the same name; the old token stops authenticating before the replacement is issued. A running gateway reloads the durable client file on each authentication decision, so a management CLI revocation takes effect on the next request without a polling interval or restart. + +Client policy also persists independent inference, search, and fetch controls. Safe defaults are two concurrent inference requests, two concurrent search/fetch requests, and 30 search/fetch requests per minute with a burst of five. Configure them with `--max-concurrent-inference`, `--search-rate`, `--max-concurrent-search`, `--fetch-rate`, and `--max-concurrent-fetch`. Permits are client-scoped and released when requests finish or are canceled. + +### Enrollment + +For bootstrap automation, have the gateway operator create a persisted, single-use token. Enrollment tokens default to 15 minutes, are stored only as hashes, are atomically marked used, and are bound to the requested client name and policy: + +```bash +term-llm gateway client add jarvis --enroll \ + --allow-provider anthropic \ + --allow-model 'anthropic:claude-sonnet-*' \ + --max-concurrent-inference 2 \ + --enroll-ttl 15m +``` + +Enrollment refuses unrestricted policies: at least one `--allow-provider` or `--allow-model` is required. Search, fetch, and CLI access remain off unless explicitly enabled. The default enrollment command is intentionally quiet about the new client token. It atomically updates `$XDG_CONFIG_HOME/term-llm/config.yaml` (or `~/.config/term-llm/config.yaml`) and writes the credential separately to a mode-`0600` `gateway-token` file: + +```bash +term-llm gateway enroll https://gateway:8787 tlge1_REDACTED --name jarvis +``` + +Use `--token-file PATH` to choose the credential path. For scripts that manage configuration themselves, use `--write-config=false --token-file PATH`. `--print-only` performs no writes and explicitly prints minimal YAML containing the client token; this is the only enrollment mode that prints that token. Reusing an enrollment token, using it after expiry, or requesting a different name fails. Direct `gateway client add` remains available when an operator can securely transfer the final client token. + +## Discourse AI model setup + +The Responses edge is intended for Discourse AI's existing OpenAI Responses client. In **Admin → Plugins → Discourse AI → LLMs**, create a model with: + +| Discourse field | Value | +|---|---| +| Provider | **OpenAI** (`open_ai`; do not select OpenRouter even when the gateway routes to OpenRouter) | +| API endpoint / URL | `https://gateway.example/v1/responses` — include the path exactly | +| API key | the gateway client token printed once by `term-llm gateway client add discourse ...` | +| Model name | a namespaced gateway ID such as `chatgpt/gpt-5.6-sol` or `openrouter/moonshotai/kimi-k2` | + +Discourse selects its Responses dialect only when the provider is OpenAI (or Azure) and the URL contains `/v1/responses`. The model namespace is split on its **first** slash, so provider models may contain additional slashes. Create a dedicated token with narrow policy, for example: + +```bash +term-llm gateway client add discourse \ + --allow-provider chatgpt \ + --allow-model 'chatgpt:gpt-5.6-sol' \ + --max-concurrent-inference 4 +``` + +Set these Discourse custom provider parameters where relevant: + +- `disable_native_tools: true` — required when agents may request Discourse tools. The gateway accepts Responses `type: function` tools only; it deliberately rejects hosted `web_search`, file-search, computer-use, MCP, and other gateway-side tool types. +- `reasoning_effort`: one of the efforts advertised for that namespaced model (for example `low`, `medium`, `high`, or `xhigh`). +- `service_tier`: `auto`, `flex`, or `priority` when supported by the selected provider/model. +- `disable_temperature: true` and/or `disable_top_p: true` for reasoning models that reject sampling controls. Discourse already removes both when `reasoning_effort` is configured. + +`GET /v1/models` uses the same bearer token and returns only configured models allowed by both server and client policy. Model IDs are always namespaced. Dynamic providers may route policy-allowed unlisted model IDs, but an unlisted ID cannot be advertised by the list endpoint. + +A Discourse-shaped streaming probe is: + +```bash +curl --no-buffer https://gateway.example/v1/responses \ + -H "Authorization: Bearer $TERM_LLM_GATEWAY_TOKEN" \ + -H 'Content-Type: application/json' \ + --data-binary @- <<'JSON' +{ + "model": "chatgpt/gpt-5.6-sol", + "input": [ + {"role":"developer","content":"Answer concisely."}, + {"role":"user","content":[{"type":"input_text","text":"Say hello."}]} + ], + "max_output_tokens": 256, + "reasoning": {"summary":"auto","effort":"medium"}, + "include": ["reasoning.encrypted_content"], + "stream": true +} +JSON +``` + +Each SSE record is a single `data: {json}` line followed by a blank line. A `[DONE]` sentinel is not required. The edge emits Responses-native created/in-progress, reasoning summary, text, function-call argument, output-item, and completed events. `response.completed.response.output` contains the complete reasoning, message, and function items used by Discourse for stateless replay. Client disconnect cancels the provider request and records a `canceled` gateway usage outcome. + +### Responses compatibility and deliberate exclusions + +Supported request semantics are: string or item-array `input`; developer/user/assistant messages; `input_text`, `output_text`, base64 data-URL `input_image`, and base64 data-URL `input_file`; encrypted reasoning replay; stateless `function_call` plus `function_call_output`; flat function tools; `none`/`auto`/`required`/named function tool choice; max output tokens; reasoning effort with `summary: auto`; temperature; top-p; service tier; parallel tool calls; streaming; usage; and harmless metadata/cache identity fields that do not alter execution. + +The edge clearly rejects semantics it cannot preserve: `previous_response_id` and conversation objects (send complete stateless input), background execution, automatic truncation, structured `text.format` output, nonzero logprobs, max hosted-tool-call limits, non-function/hosted tools, unsupported include values, and unknown top-level fields. It does not expose private `/g1` cancellation or sealed `ProviderState` as OpenAI protocol fields. Provider `reasoning.encrypted_content` is replay data sent only back to the selected provider; it is never accepted as or opened as gateway-sealed local provider state. + +## Satellite configuration + +The minimal configuration is: + +```yaml +gateway: + url: http://gateway:8787 + token_env: TERM_LLM_GATEWAY_TOKEN +``` + +A token can instead be supplied with `gateway.token` or a mode-`0600` `gateway.token_file`. Resolution precedence is token, token file, then environment. + +Full client controls: + +```yaml +gateway: + url: http://gateway:8787 + token_file: /run/secrets/term_llm_gateway_token + required: true + local_providers: [ollama, laptop-vllm] + search: true + fetch: true + catalog_ttl: 15m + connect_timeout: 2s + response_timeout: 5s + idle_timeout: 5m + tool_timeout: 10m +``` + +Gateway catalogs use live provider `ListModels` results where supported and fall back to configured/curated metadata or the last successful provider entry on transient refresh failures. Strict configured entries are available stale-first while each provider refreshes independently in the background; inference for one provider never waits for listing an unrelated provider. The server reloads provider config on its bounded `--catalog-ttl`. Known dynamic aggregators can explicitly permit unlisted models; set `providers..allow_unlisted_models: false` to force exact catalog membership, or `true` for another intentionally dynamic provider. The hidden `debug` provider has its own catalog type and always defaults to strict configured models rather than inheriting OpenAI-compatible behavior. Provider/model allow and deny policy is always enforced, including for an allowed unlisted model. + +Satellite catalogs are memoized per gateway identity in-process, coalesced across concurrent callers, and cached under the XDG cache directory with mode-`0600` files. Refreshes use ETags, short connect/response bounds, and stale disk fallback. Shell completion is cache-only and never waits for a gateway network request. + +Routing precedence is deterministic: + +1. `gateway.local_providers` +2. an explicit local `providers.` block +3. the configured gateway + +Once `gateway.url` is set, providers advertised by the gateway—including `debug`—route remotely from normal CLI commands. They remain local only when named in `gateway.local_providers` or given an explicit local `providers.` block. Providers not explicitly local fail closed if the gateway or catalog is unavailable; they never silently fall back to local credentials or built-ins. `gateway.required` additionally rejects a configuration that requests a gateway without setting `gateway.url`. With no `gateway.url`, behavior is unchanged. + +`gateway.search` and `gateway.fetch` default to `true` in a minimal gateway block. An explicit `false` is preserved and selects the ordinary local search/fetch configuration. When remote routing is enabled, gateway construction or outages produce a legible tool error rather than silently falling back to DuckDuckGo/Jina or removing `read_url`. + +## CLI providers are opt-in + +`claude-bin`, `grok-bin`, `cursor-bin`, and `gemini-cli` execute programs and use credentials on the gateway host. They are denied by default in two places: + +- the server must start with `--allow-cli` +- the client must be created with `gateway client add NAME --allow-cli` + +Both gates are required. CLI providers always receive a gateway-created empty temporary working directory. Their MCP/tool bridge calls are sent back to the originating authenticated satellite; they never run against a gateway-side copy of satellite tools or files. + +Opting in still gives the CLI provider access to its gateway-side account and whatever the CLI itself can reach from the gateway container. Use a dedicated Unix user/container, minimal environment, read-only root filesystem where practical, and narrow per-client provider/model policy. + +## Container/Jarvis migration + +For an existing Jarvis or `term-llm contain` satellite: + +1. Move provider API keys, OAuth stores, CLI homes, and search credentials to the gateway service. +2. Create one gateway client per container. Do not share tokens between satellites; attribution and revocation are client-scoped. +3. Replace provider secrets in the satellite with the `gateway` block and token secret. +4. Keep agent files, memory, jobs, session database, project mounts, and tool approvals in the satellite volume. +5. Add intentionally local providers such as an in-container Ollama instance to `gateway.local_providers`. +6. Verify `term-llm models --provider ` and a no-tool prompt, then test an approved local tool call. +7. Revoke old credentials from the satellite after validation. + +See `ops/gateway-compose.yaml` for a private Docker network example. Both services use the built-in `gateway health URL` probe, and the satellite waits for a healthy gateway. The gateway has no published host port; only the satellite UI is exposed. CI parses the Compose YAML and verifies this health gating without requiring a Docker daemon. + +## Operational checks + +```bash +curl http://gateway:8787/g1/health +term-llm models --provider anthropic +``` + +`/g1/health` intentionally returns only status and protocol version. Catalog/inference/search/fetch/run endpoints require the client bearer token and protocol version negotiation. + +Use `term-llm gateway usage` (optionally `--client` or `--json`) to inspect `usage.jsonl`, including failures, explicit `canceled` outcomes, and successful requests. Satellite-local term-llm usage remains visible with `tracked_externally_by: gateway`; aggregate local usage excludes that copy by default to avoid double-counting the gateway record. `term-llm usage --include-external` includes those external copies in either the all-provider or `--provider term-llm` view. Provider failures sent to satellites carry safe structured codes for API-key/OAuth authentication, rate limiting, context limits, invalid models/requests, and upstream failures, with a gateway/provider-specific action. Raw provider details are logged only on the gateway host and never include upstream bodies in satellite errors. diff --git a/docs-site/content/reference/configuration.md b/docs-site/content/reference/configuration.md index c426725cd..9fb5311f3 100644 --- a/docs-site/content/reference/configuration.md +++ b/docs-site/content/reference/configuration.md @@ -25,6 +25,22 @@ The main config file lives at: ~/.config/term-llm/config.yaml ``` +## Inference gateway + +Satellites can route catalog providers and web access through a central private gateway without moving their agent, session, tools, approvals, or filesystem: + +```yaml +gateway: + url: http://gateway:8787 + token_file: /run/secrets/term_llm_gateway_token + required: true + local_providers: [ollama] + search: true + fetch: true +``` + +`token`, `token_file`, and `token_env` are supported in that precedence order and are resolved only when a gateway operation runs. With no `gateway.url`, local provider behavior is unchanged. Once a URL is set, providers are remote by default and gateway outages fail closed unless the provider has an explicit local block or is listed in `local_providers`. `search` and `fetch` default to true; explicitly setting either to false preserves local routing. Default catalog/connect/response bounds are `15m`, `2s`, and `5s`. See [Inference Gateway](/guides/inference-gateway/) for live catalogs, one-use enrollment, per-client limits, policy, TLS, CLI-provider risks, and deployment. + ## Configuration shape A typical config has a few major parts: @@ -32,7 +48,7 @@ A typical config has a few major parts: - `default_provider` for the global LLM default - `providers` for model-specific credentials and routing - per-command blocks such as `exec`, `ask`, and `edit` -- feature-specific blocks such as `image`, `audio`, `music`, `embed`, `search`, `sessions`, `file_tracking`, `tools`, and `skills` +- feature-specific blocks such as `gateway`, `image`, `audio`, `music`, `embed`, `search`, `sessions`, `file_tracking`, `tools`, and `skills` ## Example diff --git a/internal/config/config.go b/internal/config/config.go index 8a721ef9b..6f421456f 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -3,11 +3,13 @@ package config import ( "bytes" "fmt" + "net/url" "os" "path/filepath" "reflect" "sort" "strings" + "time" mapstructure "github.com/go-viper/mapstructure/v2" "github.com/samsaffron/term-llm/internal/credentials" @@ -38,6 +40,11 @@ const ( ProviderTypeSambaNova ProviderType = "sambanova" ProviderTypeBedrock ProviderType = "bedrock" ProviderTypeOllama ProviderType = "ollama" + // ProviderTypeDebug is intentionally omitted from builtInProviderTypes so the + // development provider remains hidden from normal provider discovery. It is + // still a real type for gateway catalogs and must never inherit custom + // OpenAI-compatible model policy. + ProviderTypeDebug ProviderType = "debug" ) // builtInProviderTypes maps known provider names to their types @@ -68,6 +75,9 @@ func InferProviderType(name string, explicit ProviderType) ProviderType { if explicit != "" { return explicit } + if name == "debug" { + return ProviderTypeDebug + } if t, ok := builtInProviderTypes[name]; ok { return t } @@ -159,8 +169,9 @@ type ProviderConfig struct { UseNativeSearch *bool `mapstructure:"use_native_search"` // Model token limits (for custom/self-hosted models not in hardcoded tables) - ContextWindow int `mapstructure:"context_window"` - MaxOutputTokens int `mapstructure:"max_output_tokens"` + ContextWindow int `mapstructure:"context_window"` + MaxOutputTokens int `mapstructure:"max_output_tokens"` + AllowUnlistedModels *bool `mapstructure:"allow_unlisted_models"` // Gateway policy override for dynamic provider catalogs // OpenAI-compatible specific BaseURL string `mapstructure:"base_url"` // Base URL - /chat/completions is appended @@ -442,34 +453,185 @@ func parseEnvBool(value string) bool { } } +const ( + DefaultGatewayCatalogTTL = "15m" + DefaultGatewayConnectTimeout = "2s" + DefaultGatewayResponseTimeout = "5s" + DefaultGatewayIdleTimeout = "5m" + DefaultGatewayToolTimeout = "10m" + DefaultGatewayTokenEnv = "TERM_LLM_GATEWAY_TOKEN" + DefaultGatewayMaxResponseBytes = int64(64 << 20) +) + +// GatewayConfig configures the optional inference gateway used by satellites. +// When URL is empty provider/search/fetch behavior is unchanged. Nil Search and +// Fetch values intentionally default to remote routing; pointers preserve an +// explicit false in both YAML and programmatic configurations. +type GatewayConfig struct { + URL string `mapstructure:"url" yaml:"url,omitempty"` + Token string `mapstructure:"token" yaml:"token,omitempty"` + TokenFile string `mapstructure:"token_file" yaml:"token_file,omitempty"` + TokenEnv string `mapstructure:"token_env" yaml:"token_env,omitempty"` + LocalProviders []string `mapstructure:"local_providers" yaml:"local_providers,omitempty"` + Search *bool `mapstructure:"search" yaml:"search,omitempty"` + Fetch *bool `mapstructure:"fetch" yaml:"fetch,omitempty"` + Required bool `mapstructure:"required" yaml:"required,omitempty"` + CatalogTTL string `mapstructure:"catalog_ttl" yaml:"catalog_ttl,omitempty"` + ConnectTimeout string `mapstructure:"connect_timeout" yaml:"connect_timeout,omitempty"` + ResponseTimeout string `mapstructure:"response_timeout" yaml:"response_timeout,omitempty"` + IdleTimeout string `mapstructure:"idle_timeout" yaml:"idle_timeout,omitempty"` + ToolTimeout string `mapstructure:"tool_timeout" yaml:"tool_timeout,omitempty"` +} + +// Enabled reports whether a remote gateway is configured. +func (g GatewayConfig) Enabled() bool { return strings.TrimSpace(g.URL) != "" } + +// RouteSearch and RouteFetch default to true for a minimal gateway block while +// honoring an explicit false. Local search/fetch configuration is used when the +// corresponding method returns false. +func (g GatewayConfig) RouteSearch() bool { return g.Search == nil || *g.Search } +func (g GatewayConfig) RouteFetch() bool { return g.Fetch == nil || *g.Fetch } + +// ResolveToken resolves the satellite credential without mutating config. The +// explicit token wins, followed by token_file, then token_env. +func (g GatewayConfig) ResolveToken() (string, error) { + if token := strings.TrimSpace(g.Token); token != "" { + return token, nil + } + if path := strings.TrimSpace(g.TokenFile); path != "" { + if strings.HasPrefix(path, "~/") { + home, err := os.UserHomeDir() + if err != nil { + return "", fmt.Errorf("resolve gateway token file: %w", err) + } + path = filepath.Join(home, strings.TrimPrefix(path, "~/")) + } + data, err := os.ReadFile(path) + if err != nil { + return "", fmt.Errorf("read gateway token file: %w", err) + } + if token := strings.TrimSpace(string(data)); token != "" { + return token, nil + } + return "", fmt.Errorf("gateway token file %q is empty", path) + } + envName := strings.TrimSpace(g.TokenEnv) + if envName == "" { + envName = DefaultGatewayTokenEnv + } + if token := strings.TrimSpace(os.Getenv(envName)); token != "" { + return token, nil + } + return "", fmt.Errorf("gateway token is not configured (set gateway.token, gateway.token_file, or %s)", envName) +} + +// Validate checks gateway-only settings. It deliberately does nothing when no +// gateway URL is configured, preserving historical local behavior. +func (g GatewayConfig) Validate() error { + if !g.Enabled() { + if g.Required { + return fmt.Errorf("gateway.required requires gateway.url") + } + return nil + } + u, err := url.Parse(strings.TrimSpace(g.URL)) + if err != nil || (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" { + return fmt.Errorf("gateway.url must be an absolute http or https URL") + } + for name, value := range map[string]string{ + "catalog_ttl": g.CatalogTTL, "connect_timeout": g.ConnectTimeout, + "response_timeout": g.ResponseTimeout, "idle_timeout": g.IdleTimeout, + "tool_timeout": g.ToolTimeout, + } { + if strings.TrimSpace(value) == "" { + continue + } + d, parseErr := time.ParseDuration(value) + if parseErr != nil || d <= 0 { + return fmt.Errorf("gateway.%s must be a positive duration", name) + } + } + return nil +} + +// IsExplicitProvider reports whether a provider block came from user config. +// Programmatic Config values have no Viper presence map, so their Providers map +// is treated as explicit by construction. +func (c *Config) IsExplicitProvider(name string) bool { + if c == nil { + return false + } + name = strings.TrimSpace(name) + if c.explicitProviders != nil { + return c.explicitProviders[name] + } + _, ok := c.Providers[name] + return ok +} + +// ExplicitProviderNames returns only provider blocks that can be intentionally +// served. It excludes Viper-populated built-in defaults. +func (c *Config) ExplicitProviderNames() []string { + if c == nil { + return nil + } + names := make([]string, 0, len(c.Providers)) + for name := range c.Providers { + if c.IsExplicitProvider(name) { + names = append(names, name) + } + } + sort.Strings(names) + return names +} + +// IsLocalProvider reports whether a provider is explicitly pinned to the +// satellite. Explicit provider config and local_providers both take precedence +// over the remote catalog. +func (c *Config) IsLocalProvider(name string) bool { + if c == nil { + return false + } + for _, local := range c.Gateway.LocalProviders { + if strings.EqualFold(strings.TrimSpace(local), strings.TrimSpace(name)) { + return true + } + } + return c.IsExplicitProvider(name) +} + type Config struct { DefaultProvider string `mapstructure:"default_provider"` Providers map[string]ProviderConfig `mapstructure:"providers"` - Diagnostics DiagnosticsConfig `mapstructure:"diagnostics"` - DebugLogs DebugLogsConfig `mapstructure:"debug_logs"` - Sessions SessionsConfig `mapstructure:"sessions"` - Approval ApprovalConfig `mapstructure:"approval"` - Guardian GuardianConfig `mapstructure:"guardian"` - Exec ExecConfig `mapstructure:"exec"` - Ask AskConfig `mapstructure:"ask"` - Chat ChatConfig `mapstructure:"chat"` - Edit EditConfig `mapstructure:"edit"` - Loop LoopConfig `mapstructure:"loop"` - Image ImageConfig `mapstructure:"image"` - Audio AudioConfig `mapstructure:"audio"` - Music MusicConfig `mapstructure:"music"` - Transcription TranscriptionConfig `mapstructure:"transcription"` - Embed EmbedConfig `mapstructure:"embed"` - Search SearchConfig `mapstructure:"search"` - Reasoning ReasoningConfig `mapstructure:"reasoning"` - Theme ThemeConfig `mapstructure:"theme"` - Tools ToolsConfig `mapstructure:"tools"` - Agents AgentsConfig `mapstructure:"agents"` - Skills SkillsConfig `mapstructure:"skills"` - AgentsMd AgentsMdConfig `mapstructure:"agents_md"` - AutoCompact bool `mapstructure:"auto_compact"` - Serve ServeConfig `mapstructure:"serve"` - FileTracking FileTrackingConfig `mapstructure:"file_tracking"` + Gateway GatewayConfig `mapstructure:"gateway"` + // explicitProviders records provider blocks from the user's config file. It + // distinguishes them from Viper-populated built-in defaults for routing. + explicitProviders map[string]bool `mapstructure:"-"` + Diagnostics DiagnosticsConfig `mapstructure:"diagnostics"` + DebugLogs DebugLogsConfig `mapstructure:"debug_logs"` + Sessions SessionsConfig `mapstructure:"sessions"` + Approval ApprovalConfig `mapstructure:"approval"` + Guardian GuardianConfig `mapstructure:"guardian"` + Exec ExecConfig `mapstructure:"exec"` + Ask AskConfig `mapstructure:"ask"` + Chat ChatConfig `mapstructure:"chat"` + Edit EditConfig `mapstructure:"edit"` + Loop LoopConfig `mapstructure:"loop"` + Image ImageConfig `mapstructure:"image"` + Audio AudioConfig `mapstructure:"audio"` + Music MusicConfig `mapstructure:"music"` + Transcription TranscriptionConfig `mapstructure:"transcription"` + Embed EmbedConfig `mapstructure:"embed"` + Search SearchConfig `mapstructure:"search"` + Reasoning ReasoningConfig `mapstructure:"reasoning"` + Theme ThemeConfig `mapstructure:"theme"` + Tools ToolsConfig `mapstructure:"tools"` + Agents AgentsConfig `mapstructure:"agents"` + Skills SkillsConfig `mapstructure:"skills"` + AgentsMd AgentsMdConfig `mapstructure:"agents_md"` + AutoCompact bool `mapstructure:"auto_compact"` + Serve ServeConfig `mapstructure:"serve"` + FileTracking FileTrackingConfig `mapstructure:"file_tracking"` } // ApprovalConfig configures default approval behavior. @@ -974,6 +1136,15 @@ func Load() (*Config, error) { } applyProviderModelConfigs(&cfg, providerModelConfigsFromViper(viper.GetViper())) markReasoningConfigPresence(&cfg.Reasoning, viper.GetViper()) + cfg.explicitProviders = make(map[string]bool) + for name := range cfg.Providers { + if viper.InConfig("providers." + name) { + cfg.explicitProviders[name] = true + } + } + if err := cfg.Gateway.Validate(); err != nil { + return nil, err + } if err := cfg.ValidateApprovalModes(); err != nil { return nil, err } diff --git a/internal/config/gateway_test.go b/internal/config/gateway_test.go new file mode 100644 index 000000000..a15945afa --- /dev/null +++ b/internal/config/gateway_test.go @@ -0,0 +1,112 @@ +package config + +import ( + "os" + "path/filepath" + "strings" + "testing" + + "github.com/spf13/viper" + "gopkg.in/yaml.v3" +) + +func TestGatewayConfigLoadDefaultsValidationAndExplicitProviders(t *testing.T) { + viper.Reset() + t.Cleanup(viper.Reset) + dir := t.TempDir() + t.Setenv("XDG_CONFIG_HOME", dir) + t.Setenv("TERM_LLM_GATEWAY_TOKEN", "env-token") + configDir := filepath.Join(dir, "term-llm") + if err := os.MkdirAll(configDir, 0o700); err != nil { + t.Fatal(err) + } + data := `default_provider: zen +gateway: + url: http://gateway:8787 + required: true + local_providers: [ollama] +providers: + zen: + model: explicit-model +` + if err := os.WriteFile(filepath.Join(configDir, "config.yaml"), []byte(data), 0o600); err != nil { + t.Fatal(err) + } + cfg, err := Load() + if err != nil { + t.Fatal(err) + } + if !cfg.Gateway.Enabled() || !cfg.Gateway.RouteSearch() || !cfg.Gateway.RouteFetch() || cfg.Gateway.TokenEnv != DefaultGatewayTokenEnv { + t.Fatalf("gateway defaults not loaded: %+v", cfg.Gateway) + } + if token, err := cfg.Gateway.ResolveToken(); err != nil || token != "env-token" { + t.Fatalf("ResolveToken = %q, %v", token, err) + } + if !cfg.IsLocalProvider("zen") || !cfg.IsLocalProvider("ollama") { + t.Fatalf("explicit/local providers not pinned: explicit=%v", cfg.explicitProviders) + } + if cfg.IsLocalProvider("openai") { + t.Fatalf("Viper built-in default was mistaken for explicit local config: %v", cfg.explicitProviders) + } +} + +func TestGatewaySearchFetchExplicitFalseRoundTrips(t *testing.T) { + value := false + cfg := GatewayConfig{URL: "https://gateway.example", Search: &value, Fetch: &value} + if cfg.RouteSearch() || cfg.RouteFetch() { + t.Fatal("explicit false gateway search/fetch was overridden") + } + data, err := yaml.Marshal(cfg) + if err != nil { + t.Fatal(err) + } + text := string(data) + if !strings.Contains(text, "search: false") || !strings.Contains(text, "fetch: false") { + t.Fatalf("explicit false did not serialize: %s", text) + } + minimal := GatewayConfig{URL: "https://gateway.example"} + if !minimal.RouteSearch() || !minimal.RouteFetch() { + t.Fatal("minimal gateway config did not default search/fetch remote") + } +} + +func TestGatewayTokenResolutionIsDeferred(t *testing.T) { + cfg := GatewayConfig{URL: "https://gateway.example"} + if err := cfg.Validate(); err != nil { + t.Fatalf("unrelated commands should not resolve gateway token during config load: %v", err) + } + if _, err := cfg.ResolveToken(); err == nil { + t.Fatal("gateway operation accepted missing token") + } +} + +func TestGatewayConfigNoURLPreservesCompatibilityAndRejectsInvalid(t *testing.T) { + if err := (GatewayConfig{}).Validate(); err != nil { + t.Fatalf("empty gateway config changed local behavior: %v", err) + } + for _, cfg := range []GatewayConfig{ + {Required: true}, + {URL: "file:///tmp/gateway", Token: "x"}, + {URL: "http://gateway", Token: "x", CatalogTTL: "never"}, + } { + if err := cfg.Validate(); err == nil { + t.Fatalf("invalid gateway config accepted: %+v", cfg) + } + } +} + +func TestGatewayTokenFilePrecedence(t *testing.T) { + path := filepath.Join(t.TempDir(), "token") + if err := os.WriteFile(path, []byte(" file-token\n"), 0o600); err != nil { + t.Fatal(err) + } + t.Setenv("CUSTOM_GATEWAY_TOKEN", "env-token") + cfg := GatewayConfig{URL: "https://gateway.example", TokenFile: path, TokenEnv: "CUSTOM_GATEWAY_TOKEN"} + if token, err := cfg.ResolveToken(); err != nil || strings.TrimSpace(token) != "file-token" { + t.Fatalf("ResolveToken = %q, %v", token, err) + } + cfg.Token = "explicit" + if token, _ := cfg.ResolveToken(); token != "explicit" { + t.Fatalf("explicit token did not win: %q", token) + } +} diff --git a/internal/config/schema.go b/internal/config/schema.go index 26ec4d230..65dca0865 100644 --- a/internal/config/schema.go +++ b/internal/config/schema.go @@ -256,6 +256,20 @@ var keySpecs = []KeySpec{ optional("search.google.api_key", sensitive()), optional("search.google.cx"), + optional("gateway.url"), + optional("gateway.token", sensitive()), + optional("gateway.token_file"), + def("gateway.token_env", DefaultGatewayTokenEnv), + def("gateway.local_providers", []string{}), + def("gateway.search", true), + def("gateway.fetch", true), + def("gateway.required", false), + def("gateway.catalog_ttl", DefaultGatewayCatalogTTL), + def("gateway.connect_timeout", DefaultGatewayConnectTimeout), + def("gateway.response_timeout", DefaultGatewayResponseTimeout), + def("gateway.idle_timeout", DefaultGatewayIdleTimeout), + def("gateway.tool_timeout", DefaultGatewayToolTimeout), + def("reasoning.display", ReasoningDisplayAuto), def("reasoning.source", ReasoningSourceSummaryOrProviderSafe), def("reasoning.status", ReasoningStatusTitle), @@ -392,6 +406,7 @@ var providerFieldSpecs = []ProviderFieldSpec{ {Path: "use_native_search", Placeholder: false}, {Path: "context_window", Placeholder: 0}, {Path: "max_output_tokens", Placeholder: 0}, + {Path: "allow_unlisted_models", Placeholder: false}, {Path: "base_url"}, {Path: "url"}, {Path: "no_stream_options", Placeholder: false}, diff --git a/internal/gateway/catalog_test.go b/internal/gateway/catalog_test.go new file mode 100644 index 000000000..a1ecdfb6d --- /dev/null +++ b/internal/gateway/catalog_test.go @@ -0,0 +1,324 @@ +package gateway + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "sync" + "testing" + "time" + + "github.com/samsaffron/term-llm/internal/config" + "github.com/samsaffron/term-llm/internal/gateway/protocol" + "github.com/samsaffron/term-llm/internal/llm" + "github.com/spf13/viper" +) + +func TestCatalogCapabilitiesMatchProviderContracts(t *testing.T) { + tests := []struct { + provider config.ProviderType + search bool + fetch bool + choice bool + managed bool + inline bool + ordered bool + }{ + {config.ProviderTypeAnthropic, true, true, true, false, false, false}, + {config.ProviderTypeOpenAI, true, false, true, false, false, false}, + {config.ProviderTypeChatGPT, true, false, false, false, false, false}, + {config.ProviderTypeClaudeBin, false, false, false, true, false, false}, + {config.ProviderTypeGrokBin, true, true, false, true, true, false}, + {config.ProviderTypeCursorBin, false, false, false, true, true, true}, + {config.ProviderTypeGeminiCLI, true, false, false, false, false, false}, + } + for _, tc := range tests { + caps := catalogCapabilities(tc.provider) + if caps.NativeWebSearch != tc.search || caps.NativeWebFetch != tc.fetch || caps.SupportsToolChoice != tc.choice || caps.ManagesOwnContext != tc.managed || caps.InlineToolLoop != tc.inline || caps.OrderedInlineToolEvents != tc.ordered || !caps.ToolCalls { + t.Errorf("%s capabilities = %+v", tc.provider, caps) + } + } +} + +type catalogListProvider struct { + mu sync.RWMutex + models []llm.ModelInfo + err error +} + +func (*catalogListProvider) Name() string { return "catalog" } +func (*catalogListProvider) Credential() string { return "mock" } +func (*catalogListProvider) Capabilities() llm.Capabilities { return llm.Capabilities{ToolCalls: true} } +func (*catalogListProvider) Stream(context.Context, llm.Request) (llm.Stream, error) { + return nil, errors.New("unused") +} +func (p *catalogListProvider) ListModels(context.Context) ([]llm.ModelInfo, error) { + p.mu.RLock() + defer p.mu.RUnlock() + return append([]llm.ModelInfo(nil), p.models...), p.err +} +func (p *catalogListProvider) setModels(models []llm.ModelInfo) { + p.mu.Lock() + defer p.mu.Unlock() + p.models = append([]llm.ModelInfo(nil), models...) +} +func (p *catalogListProvider) setError(err error) { + p.mu.Lock() + defer p.mu.Unlock() + p.err = err +} + +func TestCatalogDebugTypeUsesStrictConfiguredModels(t *testing.T) { + cfg := &config.Config{Providers: map[string]config.ProviderConfig{ + "debug": {Model: "fast", Models: []string{"fast", "normal"}}, + }} + catalog, failed := buildCatalog(t.Context(), cfg, llm.NewProviderByName, false, false) + if len(failed) != 0 || len(catalog.Providers) != 1 { + t.Fatalf("debug catalog = %+v failed=%v", catalog, failed) + } + entry := catalog.Providers[0] + if entry.Type != string(config.ProviderTypeDebug) || entry.AllowUnlistedModels { + t.Fatalf("debug policy inferred incorrectly: %+v", entry) + } + if len(entry.Models) != 2 || entry.Models[0].ID != "fast" || entry.Models[1].ID != "normal" { + t.Fatalf("debug models = %+v", entry.Models) + } +} + +func TestCatalogUsesLiveModelsAndExplicitDynamicUnlistedPolicy(t *testing.T) { + provider := &catalogListProvider{models: []llm.ModelInfo{{ID: "vendor/new-model", DisplayName: "New", InputLimit: 123, InputPrice: 1.25, OutputPrice: 2.5}}} + cfg := &config.Config{Providers: map[string]config.ProviderConfig{ + "aggregator": {Type: config.ProviderTypeOpenRouter, Model: "stale-static", APIKey: "configured"}, + }} + catalog, failed := buildCatalog(t.Context(), cfg, func(*config.Config, string, string) (llm.Provider, error) { return provider, nil }, false, false) + if len(failed) != 0 || len(catalog.Providers) != 1 { + t.Fatalf("catalog = %+v failed=%v", catalog, failed) + } + entry := catalog.Providers[0] + if len(entry.Models) != 1 || entry.Models[0].ID != "vendor/new-model" || entry.Models[0].InputLimit != 123 || !entry.AllowUnlistedModels { + t.Fatalf("live catalog entry = %+v", entry) + } + if !catalogEntryAllowsModel(entry, "aggregator", "vendor/future-model") { + t.Fatal("dynamic aggregator unexpectedly denied an unlisted model") + } + policy := Policy{DenyModels: []string{"vendor/future-model"}} + if policy.Allows("aggregator", "vendor/future-model", false) { + t.Fatal("unlisted-model routing bypassed model policy") + } +} + +func TestCatalogMarksUnknownLivePricingWithoutAdvertisingFree(t *testing.T) { + provider := &catalogListProvider{models: []llm.ModelInfo{{ID: "brand-new-model"}}} + cfg := &config.Config{Providers: map[string]config.ProviderConfig{ + "openai": {Type: config.ProviderTypeOpenAI, APIKey: "configured"}, + }} + catalog, failed := buildCatalog(t.Context(), cfg, func(*config.Config, string, string) (llm.Provider, error) { return provider, nil }, false, false) + if len(failed) != 0 || len(catalog.Providers) != 1 || catalog.Providers[0].Models[0].InputPrice != -1 || catalog.Providers[0].Models[0].OutputPrice != -1 { + t.Fatalf("unknown live pricing = %+v failed=%v", catalog.Providers, failed) + } +} + +func TestCatalogFallsBackToConfiguredModelsWhenLiveListingFails(t *testing.T) { + provider := &catalogListProvider{err: errors.New("temporary model endpoint outage")} + cfg := &config.Config{Providers: map[string]config.ProviderConfig{ + "aggregator": {Type: config.ProviderTypeOpenRouter, APIKey: "configured", Models: []string{"configured-model"}}, + }} + catalog, failed := buildCatalog(t.Context(), cfg, func(*config.Config, string, string) (llm.Provider, error) { return provider, nil }, false, false) + if failed["aggregator"] == nil || len(catalog.Providers) != 1 || len(catalog.Providers[0].Models) != 1 || catalog.Providers[0].Models[0].ID != "configured-model" { + t.Fatalf("configured fallback = %+v failed=%v", catalog.Providers, failed) + } +} + +func TestCatalogAdvertisesOnlyExplicitConfiguredProvider(t *testing.T) { + viper.Reset() + t.Cleanup(viper.Reset) + dir := t.TempDir() + t.Setenv("XDG_CONFIG_HOME", dir) + configDir := filepath.Join(dir, "term-llm") + if err := os.MkdirAll(configDir, 0o700); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(configDir, "config.yaml"), []byte("default_provider: zen\nproviders:\n zen:\n model: minimax-m2.5-free\n"), 0o600); err != nil { + t.Fatal(err) + } + cfg, err := config.Load() + if err != nil { + t.Fatal(err) + } + provider := &catalogListProvider{models: []llm.ModelInfo{{ID: "minimax-m2.5-free"}}} + catalog, failed := buildCatalog(t.Context(), cfg, func(*config.Config, string, string) (llm.Provider, error) { return provider, nil }, false, false) + if len(failed) != 0 || len(catalog.Providers) != 1 || catalog.Providers[0].Key != "zen" { + t.Fatalf("catalog advertised Viper defaults: providers=%+v failed=%v explicit=%v", catalog.Providers, failed, cfg.ExplicitProviderNames()) + } +} + +func TestCatalogOmitsConfiguredButUnauthenticatedProvider(t *testing.T) { + cfg := &config.Config{Providers: map[string]config.ProviderConfig{"openai": {Type: config.ProviderTypeOpenAI, Model: "gpt"}}} + provider := &catalogListProvider{models: []llm.ModelInfo{{ID: "gpt"}}} + catalog, failed := buildCatalog(t.Context(), cfg, func(*config.Config, string, string) (llm.Provider, error) { return provider, nil }, false, false) + if len(catalog.Providers) != 0 || failed["openai"] == nil { + t.Fatalf("unauthenticated provider advertised: providers=%+v failed=%v", catalog.Providers, failed) + } +} + +type hungCatalogProvider struct { + started chan struct{} + once sync.Once +} + +func (*hungCatalogProvider) Name() string { return "hung-catalog" } +func (*hungCatalogProvider) Credential() string { return "mock" } +func (*hungCatalogProvider) Capabilities() llm.Capabilities { return llm.Capabilities{ToolCalls: true} } +func (*hungCatalogProvider) Stream(context.Context, llm.Request) (llm.Stream, error) { + return nil, errors.New("unused") +} +func (p *hungCatalogProvider) ListModels(ctx context.Context) ([]llm.ModelInfo, error) { + p.once.Do(func() { close(p.started) }) + <-ctx.Done() + return nil, ctx.Err() +} + +func TestHungUnrelatedCatalogDoesNotDelayHealthyInference(t *testing.T) { + dir := t.TempDir() + clients, _ := OpenClientStore(filepath.Join(dir, "clients.json")) + client, token, _ := clients.Add("timing-client", Policy{MaxConcurrentInference: 1}) + sealer, _ := OpenStateSealer(filepath.Join(dir, "state.key")) + strict := false + cfg := &config.Config{Providers: map[string]config.ProviderConfig{ + "healthy": {Type: config.ProviderTypeZen, Model: "healthy-model", Models: []string{"healthy-model"}, AllowUnlistedModels: &strict}, + "broken": {Type: config.ProviderTypeOpenAICompat, Model: "broken-model", Models: []string{"broken-model"}, AllowUnlistedModels: &strict}, + }} + healthy := llm.NewMockProvider("healthy").AddTurn(llm.MockTurn{Text: "healthy response"}) + hung := &hungCatalogProvider{started: make(chan struct{})} + server, err := NewServer(ServerConfig{ + Config: cfg, Clients: clients, Sealer: sealer, CatalogTTL: time.Millisecond, ModelListTimeout: 2 * time.Second, + ProviderFactory: func(_ *config.Config, name, _ string) (llm.Provider, error) { + if name == "broken" { + return hung, nil + } + return healthy, nil + }, + }) + if err != nil { + t.Fatal(err) + } + ts := httptest.NewServer(server.Handler()) + defer ts.Close() + + // The full catalog returns configured stale-first entries and starts live + // per-provider refreshes. Wait until the unrelated broken lister is blocked. + satellite := &config.Config{Gateway: config.GatewayConfig{URL: ts.URL, Token: token, CatalogTTL: "1ms", ConnectTimeout: "1s", ResponseTimeout: "1s", IdleTimeout: "1s"}, Providers: map[string]config.ProviderConfig{}} + provider, err := llm.NewGatewayProvider(satellite, "healthy", "healthy-model") + if err != nil { + t.Fatal(err) + } + select { + case <-hung.started: + case <-time.After(500 * time.Millisecond): + t.Fatal("broken catalog refresh did not start") + } + started := time.Now() + stream, err := provider.Stream(t.Context(), llm.Request{Model: "healthy-model", Messages: []llm.Message{llm.UserText("hello")}}) + if err != nil { + t.Fatal(err) + } + text, _ := collectStream(t, stream) + if text != "healthy response" { + t.Fatalf("healthy inference text = %q", text) + } + if elapsed := time.Since(started); elapsed > 750*time.Millisecond { + t.Fatalf("healthy inference waited %s for unrelated catalog; bound is 750ms", elapsed) + } + + var release func() + var ok bool + deadline := time.Now().Add(500 * time.Millisecond) + for { + release, ok = server.acquireInference(client) + if ok { + break + } + if time.Now().After(deadline) { + t.Fatal("could not reserve timing client inference slot") + } + time.Sleep(time.Millisecond) + } + defer release() + wireRequest, err := llm.EncodeGatewayRequest(llm.Request{Model: "broken-model", Messages: []llm.Message{llm.UserText("limit")}}) + if err != nil { + t.Fatal(err) + } + payload, _ := json.Marshal(protocol.InferenceRequest{Version: protocol.Version, RequestID: "req-limit", Provider: "broken", Request: wireRequest}) + req, _ := http.NewRequest(http.MethodPost, ts.URL+"/g1/inference", bytes.NewReader(payload)) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set(protocol.VersionHeader, "1") + started = time.Now() + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if resp.StatusCode != http.StatusTooManyRequests { + t.Fatalf("concurrency response = %d, want 429", resp.StatusCode) + } + if elapsed := time.Since(started); elapsed > 250*time.Millisecond { + t.Fatalf("concurrency response waited %s for hung catalog; bound is 250ms", elapsed) + } +} + +func TestServerCatalogReloadsProviderConfigAfterTTL(t *testing.T) { + dir := t.TempDir() + clients, _ := OpenClientStore(filepath.Join(dir, "clients.json")) + sealer, _ := OpenStateSealer(filepath.Join(dir, "state.key")) + current := &config.Config{Providers: map[string]config.ProviderConfig{"first": {Type: config.ProviderTypeZen, Model: "model-a"}}} + provider := &catalogListProvider{models: []llm.ModelInfo{{ID: "model-a"}}} + server, err := NewServer(ServerConfig{ + Config: current, Clients: clients, Sealer: sealer, CatalogTTL: time.Millisecond, + ConfigLoader: func() (*config.Config, error) { return current, nil }, + ProviderFactory: func(*config.Config, string, string) (llm.Provider, error) { return provider, nil }, + }) + if err != nil { + t.Fatal(err) + } + first, err := server.currentCatalog(t.Context()) + if err != nil || len(first.Providers) != 1 || first.Providers[0].Key != "first" { + t.Fatalf("first config catalog = %+v, %v", first.Providers, err) + } + current = &config.Config{Providers: map[string]config.ProviderConfig{"second": {Type: config.ProviderTypeZen, Model: "model-b"}}} + provider.setModels([]llm.ModelInfo{{ID: "model-b"}}) + time.Sleep(2 * time.Millisecond) + second, err := server.currentCatalog(t.Context()) + if err != nil || len(second.Providers) != 1 || second.Providers[0].Key != "second" || second.Providers[0].Models[0].ID != "model-b" { + t.Fatalf("refreshed config catalog = %+v, %v", second.Providers, err) + } +} + +func TestServerCatalogRefreshesAndFallsBackToStaleLiveModels(t *testing.T) { + dir := t.TempDir() + clients, _ := OpenClientStore(filepath.Join(dir, "clients.json")) + sealer, _ := OpenStateSealer(filepath.Join(dir, "state.key")) + provider := &catalogListProvider{models: []llm.ModelInfo{{ID: "model-v1"}}} + cfg := &config.Config{Providers: map[string]config.ProviderConfig{"remote": {Type: config.ProviderTypeOpenRouter, APIKey: "configured"}}} + server, err := NewServer(ServerConfig{ + Config: cfg, Clients: clients, Sealer: sealer, CatalogTTL: time.Millisecond, + ProviderFactory: func(*config.Config, string, string) (llm.Provider, error) { return provider, nil }, + }) + if err != nil { + t.Fatal(err) + } + first, err := server.currentCatalog(t.Context()) + if err != nil || first.Providers[0].Models[0].ID != "model-v1" { + t.Fatalf("first catalog = %+v, %v", first, err) + } + time.Sleep(2 * time.Millisecond) + provider.setError(errors.New("temporary secret upstream failure")) + stale, err := server.currentCatalog(t.Context()) + if err != nil || stale.Providers[0].Models[0].ID != "model-v1" { + t.Fatalf("stale catalog fallback = %+v, %v", stale, err) + } +} diff --git a/internal/gateway/error_test.go b/internal/gateway/error_test.go new file mode 100644 index 000000000..292151aea --- /dev/null +++ b/internal/gateway/error_test.go @@ -0,0 +1,72 @@ +package gateway + +import ( + "context" + "errors" + "fmt" + "net/http" + "strings" + "testing" + + "github.com/samsaffron/term-llm/internal/config" + "github.com/samsaffron/term-llm/internal/llm" + "github.com/samsaffron/term-llm/internal/providerhttp" +) + +func TestProviderErrorsAreStructuredActionableAndSafe(t *testing.T) { + tests := []struct { + name string + err error + providerType config.ProviderType + wantStatus int + wantCode string + wantAction string + }{ + {"api key", providerhttp.NewStatusErrorString("OpenAI", 401, "401 Unauthorized", nil, `raw-body api_key=super-secret /srv/private`), config.ProviderTypeOpenAI, 401, "provider_api_key_unauthenticated", "update the API key"}, + {"missing API key", errors.New("OPENAI_API_KEY is required at /srv/private"), config.ProviderTypeOpenAI, 401, "provider_api_key_unauthenticated", "update the API key"}, + {"oauth", providerhttp.NewStatusErrorString("ChatGPT", 401, "401 Unauthorized", nil, `oauth_token=super-secret`), config.ProviderTypeChatGPT, 401, "provider_oauth_unauthenticated", "gateway host"}, + {"missing oauth", errors.New("OAuth login credential missing at /srv/private"), config.ProviderTypeCopilot, 401, "provider_oauth_unauthenticated", "gateway host"}, + {"rate limit", providerhttp.NewStatusErrorString("OpenAI", 429, "429 Too Many Requests", nil, `account secret limit`), config.ProviderTypeOpenAI, 429, "provider_rate_limited", "wait and retry"}, + {"typed rate limit", &llm.RateLimitError{Message: "rate limit secret"}, config.ProviderTypeOpenAI, 429, "provider_rate_limited", "wait and retry"}, + {"context", errors.New("maximum context length exceeded; prompt /tmp/private"), config.ProviderTypeAnthropic, 400, "provider_context_limit", "compact"}, + {"model", providerhttp.NewStatusErrorString("OpenAI", 400, "400 Bad Request", nil, `model private-preview does not exist for key secret`), config.ProviderTypeOpenAI, 400, "provider_model_invalid", "catalog"}, + {"upstream", providerhttp.NewStatusErrorString("OpenAI", 503, "503 Service Unavailable", nil, `upstream raw body /etc/passwd`), config.ProviderTypeOpenAI, 502, "provider_upstream_failure", "diagnostics"}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + status, code := classifyProviderError(tc.err, string(tc.providerType)) + if status != tc.wantStatus || code != tc.wantCode { + t.Fatalf("classification = %d/%s, want %d/%s", status, code, tc.wantStatus, tc.wantCode) + } + message := safeProviderErrorMessage(code, "remote") + if !strings.Contains(message, "gateway provider") || !strings.Contains(message, tc.wantAction) { + t.Fatalf("safe message is not actionable: %q", message) + } + for _, forbidden := range []string{"super-secret", "raw-body", "/srv/", "/tmp/", "/etc/", "private-preview"} { + if strings.Contains(message, forbidden) { + t.Fatalf("safe message leaked %q: %q", forbidden, message) + } + } + }) + } +} + +func TestCancellationErrorIsClassifiedExplicitly(t *testing.T) { + status, code := classifyProviderError(fmt.Errorf("stream stopped: %w", context.Canceled), string(config.ProviderTypeOpenAI)) + if status != 499 || code != "canceled" { + t.Fatalf("cancellation classification = %d/%s", status, code) + } + if message := safeProviderErrorMessage(code, "openai"); !strings.Contains(message, "canceled") || strings.Contains(message, "upstream") { + t.Fatalf("cancellation message = %q", message) + } +} + +func TestSafeErrorFallbackIsGatewaySpecific(t *testing.T) { + status, code := classifyProviderError(errors.New("opaque transport failure"), string(config.ProviderTypeOpenAI)) + if status != http.StatusBadGateway || code != "provider_upstream_failure" { + t.Fatalf("fallback = %d/%s", status, code) + } + if got := safeProviderErrorMessage(code, "openai"); !strings.Contains(got, "gateway provider \"openai\"") || !strings.Contains(got, "retry") { + t.Fatalf("fallback message = %q", got) + } +} diff --git a/internal/gateway/execution.go b/internal/gateway/execution.go new file mode 100644 index 000000000..ada212d6a --- /dev/null +++ b/internal/gateway/execution.go @@ -0,0 +1,128 @@ +package gateway + +import ( + "context" + "errors" + "log/slog" + "net/http" + "os" + "time" + + "github.com/samsaffron/term-llm/internal/gateway/protocol" + "github.com/samsaffron/term-llm/internal/llm" +) + +// startInference is the single provider-edge admission and setup path shared by +// the private /g1 transport and the public OpenAI-compatible Responses edge. +// It applies catalog, policy, concurrency, credential, retry, state, filesystem +// isolation, and cancellation rules without instantiating an agent runtime. +func (s *Server) startInference(parent context.Context, client Client, envelope protocol.InferenceRequest, providerReq llm.Request, rejectInlineToolLoop bool) (*inferenceExecution, *inferenceRequestError) { + if providerReq.Model == "" { + if pc := s.currentConfig().GetProviderConfig(envelope.Provider); pc != nil { + providerReq.Model = pc.Model + } + } + + releaseInference, ok := s.acquireInference(client) + if !ok { + s.recordFailure(client, envelope, providerReq, "client_concurrency_limited", time.Now().UTC()) + return nil, &inferenceRequestError{ + Status: http.StatusTooManyRequests, + Code: "client_concurrency_limited", + Message: "this gateway client has reached its concurrent inference limit; " + + "wait for another request to finish", + } + } + fail := func(status int, code, message string) (*inferenceExecution, *inferenceRequestError) { + releaseInference() + return nil, &inferenceRequestError{Status: status, Code: code, Message: message} + } + + entry, found, catalogErr := s.currentCatalogProvider(parent, envelope.Provider) + if catalogErr != nil { + slog.Error("refresh gateway provider catalog for inference", "request_id", envelope.RequestID, "provider", envelope.Provider, "error", catalogErr) + return fail(http.StatusServiceUnavailable, "catalog_unavailable", "gateway provider catalog is temporarily unavailable; retry or contact the gateway operator") + } + if !found { + s.recordFailure(client, envelope, providerReq, "unknown_provider", time.Now().UTC()) + return fail(http.StatusNotFound, "unknown_provider", "provider is not available") + } + if !catalogEntryAllowsModel(entry, envelope.Provider, providerReq.Model) { + s.recordFailure(client, envelope, providerReq, "unknown_model", time.Now().UTC()) + return fail(http.StatusNotFound, "unknown_model", "model is not available for this gateway provider; choose a catalog model or ask the gateway operator to allow unlisted models") + } + if !s.cfg.Policy.Allows(envelope.Provider, providerReq.Model, entry.CLI) || !client.Policy.Allows(envelope.Provider, providerReq.Model, entry.CLI) { + s.recordFailure(client, envelope, providerReq, "policy_denied", time.Now().UTC()) + return fail(http.StatusForbidden, "policy_denied", "provider/model is denied by gateway policy; choose an allowed model or contact the gateway operator") + } + if rejectInlineToolLoop && len(providerReq.Tools) > 0 && entry.CLI && entry.Capabilities.InlineToolLoop { + s.recordFailure(client, envelope, providerReq, "incompatible_tool_request", time.Now().UTC()) + return fail(http.StatusBadRequest, "incompatible_tool_request", "this CLI provider requires an inline tool loop; Responses function tools are rejected because the gateway never executes client tools") + } + + started := time.Now().UTC() + failStarted := func(status int, code, message string) (*inferenceExecution, *inferenceRequestError) { + s.recordUsage(client, envelope, providerReq, llm.Usage{}, code, started) + return fail(status, code, message) + } + if err := nonInteractiveAuthReady(entry.Type); err != nil { + status, code := classifyProviderError(err, entry.Type) + slog.Error("gateway provider authentication unavailable", "request_id", envelope.RequestID, "provider", envelope.Provider, "error", err) + return failStarted(status, code, safeProviderErrorMessage(code, envelope.Provider)) + } + provider, err := s.cfg.ProviderFactory(s.centralConfig(), envelope.Provider, providerReq.Model) + if err != nil { + status, code := classifyProviderError(err, entry.Type) + slog.Error("create gateway provider", "request_id", envelope.RequestID, "provider", envelope.Provider, "error", err) + return failStarted(status, code, safeProviderErrorMessage(code, envelope.Provider)) + } + provider = llm.WrapWithRetry(provider, llm.RetryConfig{ + MaxAttempts: s.cfg.UpstreamRetryAttempts, + MaxElapsedTime: s.cfg.UpstreamRetryMaxElapsed, + BaseBackoff: time.Second, + MaxBackoff: 5 * time.Second, + }) + if envelope.State != "" { + plain, openErr := s.cfg.Sealer.Open(envelope.State, client.ID, envelope.Provider) + if openErr != nil { + return failStarted(http.StatusBadRequest, "invalid_state", "provider state is invalid or does not belong to this client/provider") + } + importer, ok := provider.(llm.ProviderStateImporter) + if !ok { + return failStarted(http.StatusBadRequest, "invalid_state", "provider does not accept state") + } + if err := importer.ImportProviderState(plain); err != nil { + return failStarted(http.StatusBadRequest, "invalid_state", "provider state was rejected") + } + } + + tempDir, err := s.newRunTempDir() + if err != nil { + return failStarted(http.StatusInternalServerError, "internal", "could not create isolated provider directory") + } + providerReq.WorkingDir = tempDir + ctx, cancel := context.WithCancel(parent) + stream, err := provider.Stream(ctx, providerReq) + if err != nil { + cancel() + _ = os.RemoveAll(tempDir) + status, code := classifyProviderError(err, entry.Type) + slog.Error("gateway provider request failed", "request_id", envelope.RequestID, "provider", envelope.Provider, "error", err) + return failStarted(status, code, safeProviderErrorMessage(code, envelope.Provider)) + } + + return &inferenceExecution{ + server: s, client: client, envelope: envelope, request: providerReq, entry: entry, + provider: provider, stream: stream, ctx: ctx, cancel: cancel, release: releaseInference, + tempDir: tempDir, started: started, + }, nil +} + +func (s *Server) logProviderStreamError(execution *inferenceExecution, err error) (int, string) { + status, code := classifyProviderError(err, execution.entry.Type) + if errors.Is(err, context.Canceled) || errors.Is(execution.ctx.Err(), context.Canceled) { + code = "canceled" + } + slog.Error("gateway provider stream failed", "request_id", execution.envelope.RequestID, "provider", execution.envelope.Provider, "error", err) + return status, code +} diff --git a/internal/gateway/integration_test.go b/internal/gateway/integration_test.go new file mode 100644 index 000000000..386279571 --- /dev/null +++ b/internal/gateway/integration_test.go @@ -0,0 +1,661 @@ +package gateway + +import ( + "context" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "sync" + "testing" + "time" + + "github.com/samsaffron/term-llm/internal/config" + "github.com/samsaffron/term-llm/internal/llm" + "github.com/samsaffron/term-llm/internal/search" +) + +type recordingUsage struct { + mu sync.Mutex + records []UsageRecord +} + +func (r *recordingUsage) Record(record UsageRecord) error { + r.mu.Lock() + defer r.mu.Unlock() + r.records = append(r.records, record) + return nil +} + +type echoTool struct{} + +func (echoTool) Spec() llm.ToolSpec { + return llm.ToolSpec{Name: "echo", Description: "echo", Schema: map[string]any{"type": "object", "properties": map[string]any{"text": map[string]any{"type": "string"}}}} +} +func (echoTool) Preview(json.RawMessage) string { return "echo" } +func (echoTool) Execute(_ context.Context, args json.RawMessage) (llm.ToolOutput, error) { + var payload struct { + Text string `json:"text"` + } + _ = json.Unmarshal(args, &payload) + return llm.TextOutput("satellite:" + payload.Text), nil +} + +type inlineProvider struct { + mu sync.Mutex + response llm.ToolExecutionResponse +} + +func (*inlineProvider) Name() string { return "inline" } +func (*inlineProvider) Credential() string { return "mock" } +func (*inlineProvider) Capabilities() llm.Capabilities { + return llm.Capabilities{ToolCalls: true, InlineToolLoop: true, ManagesOwnContext: true, OrderedInlineToolEvents: true} +} +func (p *inlineProvider) Stream(ctx context.Context, req llm.Request) (llm.Stream, error) { + return &inlineProviderStream{ctx: ctx, provider: p}, nil +} + +type inlineProviderStream struct { + ctx context.Context + provider *inlineProvider + step int + response chan llm.ToolExecutionResponse +} + +func (s *inlineProviderStream) Recv() (llm.Event, error) { + s.step++ + switch s.step { + case 1: + s.response = make(chan llm.ToolExecutionResponse, 1) + return llm.Event{Type: llm.EventTextDelta, Text: "before "}, nil + case 2: + return llm.Event{Type: llm.EventToolCall, ToolCallID: "inline-call", ToolName: "echo", Tool: &llm.ToolCall{ID: "inline-call", Name: "echo", Arguments: json.RawMessage(`{"text":"hello"}`)}, ToolResponse: s.response}, nil + case 3: + select { + case response := <-s.response: + s.provider.mu.Lock() + s.provider.response = response + s.provider.mu.Unlock() + if response.Err != nil { + return llm.Event{Type: llm.EventTextDelta, Text: "tool-error"}, nil + } + return llm.Event{Type: llm.EventTextDelta, Text: "after " + response.Result.Content}, nil + case <-s.ctx.Done(): + return llm.Event{}, s.ctx.Err() + } + case 4: + return llm.Event{Type: llm.EventUsage, Use: &llm.Usage{InputTokens: 10, OutputTokens: 4}}, nil + default: + return llm.Event{}, io.EOF + } +} +func (*inlineProviderStream) Close() error { return nil } + +type fakeSearcher struct{} + +func (fakeSearcher) Search(context.Context, string, int) ([]search.Result, error) { + return []search.Result{{Title: "Gateway result", URL: "https://example.com", Snippet: "central"}}, nil +} + +type fakeFetcher struct{} + +func (fakeFetcher) FetchURL(context.Context, string) (string, error) { return "central fetch", nil } + +type setupFailureProvider struct { + mu sync.Mutex + attempts int + block bool +} + +func (*setupFailureProvider) Name() string { return "setup-failure" } +func (*setupFailureProvider) Credential() string { return "mock" } +func (*setupFailureProvider) Capabilities() llm.Capabilities { return llm.Capabilities{} } +func (p *setupFailureProvider) Stream(ctx context.Context, _ llm.Request) (llm.Stream, error) { + p.mu.Lock() + p.attempts++ + p.mu.Unlock() + if p.block { + <-ctx.Done() + return nil, ctx.Err() + } + return nil, errors.New("500 Internal Server Error") +} + +func (p *setupFailureProvider) attemptCount() int { + p.mu.Lock() + defer p.mu.Unlock() + return p.attempts +} + +type gatewayFixture struct { + server *httptest.Server + gateway *Server + central *config.Config + clients *ClientStore + client Client + token string + usage *recordingUsage + provider llm.Provider +} + +func newGatewayFixture(t *testing.T, providerType config.ProviderType, provider llm.Provider, toolTimeout time.Duration) *gatewayFixture { + t.Helper() + binDir := t.TempDir() + for _, name := range []string{"claude", "grok", "cursor-agent", "gemini"} { + if err := os.WriteFile(filepath.Join(binDir, name), []byte("#!/bin/sh\nexit 0\n"), 0o700); err != nil { + t.Fatal(err) + } + } + t.Setenv("PATH", binDir+string(os.PathListSeparator)+os.Getenv("PATH")) + t.Setenv("CURSOR_API_KEY", "test-cursor-key") + dir := t.TempDir() + clients, err := OpenClientStore(filepath.Join(dir, "clients.json")) + if err != nil { + t.Fatal(err) + } + client, token, err := clients.Add("satellite-a", Policy{AllowCLI: true, AllowSearch: true, AllowFetch: true}) + if err != nil { + t.Fatal(err) + } + sealer, err := OpenStateSealer(filepath.Join(dir, "state.key")) + if err != nil { + t.Fatal(err) + } + central := &config.Config{DefaultProvider: "remote", Providers: map[string]config.ProviderConfig{"remote": {Type: providerType, Model: "model-a", Models: []string{"model-a"}, APIKey: "super-secret-provider-key"}}} + usage := &recordingUsage{} + server, err := NewServer(ServerConfig{ + Config: central, Clients: clients, Sealer: sealer, Usage: usage, + ProviderFactory: func(*config.Config, string, string) (llm.Provider, error) { return provider, nil }, + Searcher: fakeSearcher{}, FetchTool: llm.NewReadURLToolWithFetcher(fakeFetcher{}), + Policy: Policy{AllowCLI: true, AllowSearch: true, AllowFetch: true}, ToolTimeout: toolTimeout, + }) + if err != nil { + t.Fatal(err) + } + ts := httptest.NewServer(server.Handler()) + t.Cleanup(ts.Close) + return &gatewayFixture{server: ts, gateway: server, central: central, clients: clients, client: client, token: token, usage: usage, provider: provider} +} + +func (f *gatewayFixture) satelliteConfig() *config.Config { + return &config.Config{Gateway: config.GatewayConfig{URL: f.server.URL, Token: f.token, Search: gatewayBool(true), Fetch: gatewayBool(true), CatalogTTL: "1m", ConnectTimeout: "2s", ResponseTimeout: "2s", ToolTimeout: "2s"}, Providers: map[string]config.ProviderConfig{}} +} + +func gatewayBool(value bool) *bool { return &value } + +func collectStream(t *testing.T, stream llm.Stream) (string, llm.Usage) { + t.Helper() + defer stream.Close() + var text strings.Builder + var usage llm.Usage + for { + event, err := stream.Recv() + if err == io.EOF { + break + } + if err != nil { + t.Fatal(err) + } + if event.Type == llm.EventTextDelta { + text.WriteString(event.Text) + } + if event.Type == llm.EventUsage && event.Use != nil { + usage.Add(*event.Use) + } + } + return text.String(), usage +} + +func TestGatewayProviderServerSSEFidelityIsolationAndUsage(t *testing.T) { + mock := llm.NewMockProvider("central").AddTurn(llm.MockTurn{Text: "hello from central", Usage: llm.Usage{InputTokens: 7, OutputTokens: 3, CachedInputTokens: 2}}) + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, mock, time.Second) + provider, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + catalogReq, _ := http.NewRequest(http.MethodGet, fixture.server.URL+"/g1/catalog", nil) + catalogReq.Header.Set("Authorization", "Bearer "+fixture.token) + catalogReq.Header.Set("Term-LLM-Gateway-Version", "1") + catalogResp, err := http.DefaultClient.Do(catalogReq) + if err != nil { + t.Fatal(err) + } + catalogBody, _ := io.ReadAll(catalogResp.Body) + catalogResp.Body.Close() + if strings.Contains(string(catalogBody), "super-secret-provider-key") { + t.Fatal("provider credential leaked through catalog") + } + stream, err := provider.Stream(context.Background(), llm.Request{Model: "model-a", SessionID: "sess", WorkingDir: "/satellite/secret", Messages: []llm.Message{llm.UserText("hello")}}) + if err != nil { + t.Fatal(err) + } + text, usage := collectStream(t, stream) + if text != "hello from central" || usage.InputTokens != 7 || usage.CachedInputTokens != 2 { + t.Fatalf("text/usage = %q %+v", text, usage) + } + requests := mock.RecordedRequests() + if len(requests) != 1 || requests[0].WorkingDir == "" || requests[0].WorkingDir == "/satellite/secret" { + t.Fatalf("gateway working dir isolation failed: %+v", requests) + } + if _, err := os.Stat(requests[0].WorkingDir); !os.IsNotExist(err) { + t.Fatalf("ephemeral gateway working dir still exists: %s (%v)", requests[0].WorkingDir, err) + } + deadline := time.Now().Add(time.Second) + for { + fixture.usage.mu.Lock() + count := len(fixture.usage.records) + fixture.usage.mu.Unlock() + if count > 0 || time.Now().After(deadline) { + break + } + time.Sleep(time.Millisecond) + } + fixture.usage.mu.Lock() + defer fixture.usage.mu.Unlock() + if len(fixture.usage.records) != 1 || fixture.usage.records[0].ClientID != fixture.client.ID || fixture.usage.records[0].ProviderKey != "remote" || fixture.usage.records[0].RequestID == "" || fixture.usage.records[0].InputTokens != 7 { + t.Fatalf("usage attribution = %+v", fixture.usage.records) + } +} + +func TestGatewayNormalToolLoopExecutesOnSatellite(t *testing.T) { + mock := llm.NewMockProvider("central"). + AddTurn(llm.MockTurn{ToolCalls: []llm.ToolCall{{ID: "call-1", Name: "echo", Arguments: json.RawMessage(`{"text":"normal"}`)}}}). + AddTurn(llm.MockTurn{Text: "normal tool complete"}) + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, mock, time.Second) + provider, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + registry := llm.NewToolRegistry() + registry.Register(echoTool{}) + engine := llm.NewEngine(provider, registry) + stream, err := engine.Stream(context.Background(), llm.Request{Model: "model-a", Messages: []llm.Message{llm.UserText("use echo")}, Tools: []llm.ToolSpec{echoTool{}.Spec()}, MaxTurns: 3}) + if err != nil { + t.Fatal(err) + } + text, _ := collectStream(t, stream) + if text != "normal tool complete" { + t.Fatalf("normal tool loop text = %q", text) + } + requests := mock.RecordedRequests() + if len(requests) != 2 { + t.Fatalf("central provider requests = %d, want 2", len(requests)) + } + foundResult := false + for _, message := range requests[1].Messages { + for _, part := range message.Parts { + if part.ToolResult != nil && part.ToolResult.Content == "satellite:normal" { + foundResult = true + } + } + } + if !foundResult { + t.Fatalf("satellite tool result missing from continuation request: %+v", requests[1].Messages) + } +} + +func TestGatewayInlineToolCallbackExecutesOnSatellite(t *testing.T) { + inline := &inlineProvider{} + fixture := newGatewayFixture(t, config.ProviderTypeCursorBin, inline, time.Second) + provider, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + registry := llm.NewToolRegistry() + registry.Register(echoTool{}) + engine := llm.NewEngine(provider, registry) + stream, err := engine.Stream(context.Background(), llm.Request{Model: "model-a", Messages: []llm.Message{llm.UserText("use echo")}, Tools: []llm.ToolSpec{echoTool{}.Spec()}, MaxTurns: 3}) + if err != nil { + t.Fatal(err) + } + text, _ := collectStream(t, stream) + if text != "before after satellite:hello" { + t.Fatalf("inline callback text = %q", text) + } + inline.mu.Lock() + defer inline.mu.Unlock() + if inline.response.Err != nil || inline.response.Result.Content != "satellite:hello" { + t.Fatalf("central provider callback response = %+v", inline.response) + } +} + +func TestGatewayConcurrentInlineCallbacksDoNotDeadlock(t *testing.T) { + inline := &inlineProvider{} + fixture := newGatewayFixture(t, config.ProviderTypeCursorBin, inline, time.Second) + provider, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + results := make(chan string, 2) + errs := make(chan error, 2) + for i := 0; i < 2; i++ { + go func() { + registry := llm.NewToolRegistry() + registry.Register(echoTool{}) + engine := llm.NewEngine(provider, registry) + stream, err := engine.Stream(context.Background(), llm.Request{Model: "model-a", Messages: []llm.Message{llm.UserText("use echo")}, Tools: []llm.ToolSpec{echoTool{}.Spec()}, MaxTurns: 3}) + if err != nil { + errs <- err + return + } + var text strings.Builder + for { + event, recvErr := stream.Recv() + if recvErr == io.EOF { + break + } + if recvErr != nil { + errs <- recvErr + return + } + if event.Type == llm.EventTextDelta { + text.WriteString(event.Text) + } + } + _ = stream.Close() + results <- text.String() + }() + } + for i := 0; i < 2; i++ { + select { + case err := <-errs: + t.Fatal(err) + case text := <-results: + if text != "before after satellite:hello" { + t.Fatalf("concurrent callback text = %q", text) + } + case <-time.After(2 * time.Second): + t.Fatal("concurrent gateway callbacks deadlocked") + } + } +} + +func TestGatewayInlineToolCallbackTimeout(t *testing.T) { + inline := &inlineProvider{} + fixture := newGatewayFixture(t, config.ProviderTypeCursorBin, inline, 30*time.Millisecond) + provider, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + stream, err := provider.Stream(context.Background(), llm.Request{Model: "model-a", Messages: []llm.Message{llm.UserText("call")}}) + if err != nil { + t.Fatal(err) + } + first, _ := stream.Recv() + if first.Text != "before " { + t.Fatalf("first event = %+v", first) + } + callback, _ := stream.Recv() + if callback.ToolResponse == nil { + t.Fatalf("callback event missing response channel: %+v", callback) + } + next, err := stream.Recv() + if err != nil || next.Text != "tool-error" { + t.Fatalf("post-timeout event = %+v, %v", next, err) + } + _ = stream.Close() +} + +type abandonedToolResponseProvider struct{ closed chan struct{} } + +func (*abandonedToolResponseProvider) Name() string { return "abandoned-callback" } +func (*abandonedToolResponseProvider) Credential() string { return "mock" } +func (*abandonedToolResponseProvider) Capabilities() llm.Capabilities { + return llm.Capabilities{ToolCalls: true, InlineToolLoop: true} +} +func (p *abandonedToolResponseProvider) Stream(context.Context, llm.Request) (llm.Stream, error) { + return &abandonedToolResponseStream{closed: p.closed}, nil +} + +type abandonedToolResponseStream struct { + closed chan struct{} + sent bool + once sync.Once +} + +func (s *abandonedToolResponseStream) Recv() (llm.Event, error) { + if !s.sent { + s.sent = true + return llm.Event{Type: llm.EventToolCall, ToolCallID: "abandoned", ToolName: "echo", ToolResponse: make(chan llm.ToolExecutionResponse)}, nil + } + return llm.Event{}, io.EOF +} +func (s *abandonedToolResponseStream) Close() error { + s.once.Do(func() { close(s.closed) }) + return nil +} + +func TestGatewayToolResponseSendUnblocksOnCancellation(t *testing.T) { + central := &abandonedToolResponseProvider{closed: make(chan struct{})} + fixture := newGatewayFixture(t, config.ProviderTypeCursorBin, central, time.Second) + provider, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + stream, err := provider.Stream(context.Background(), llm.Request{Model: "model-a", Messages: []llm.Message{llm.UserText("call")}}) + if err != nil { + t.Fatal(err) + } + event, err := stream.Recv() + if err != nil || event.ToolResponse == nil { + t.Fatalf("callback event = %+v, %v", event, err) + } + event.ToolResponse <- llm.ToolExecutionResponse{Result: llm.TextOutput("no receiver")} + time.Sleep(20 * time.Millisecond) + if err := stream.Close(); err != nil { + t.Fatal(err) + } + select { + case <-central.closed: + case <-time.After(time.Second): + t.Fatal("gateway handler remained blocked sending ToolResponse after cancellation") + } +} + +func TestGatewayCrossClientStateAndRunAccessDenied(t *testing.T) { + mock := llm.NewMockProvider("central").AddTextResponse("one").AddTextResponse("two") + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, mock, time.Second) + provider, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + // A syntactically valid but unauthenticated state is rejected before provider execution. + if err := provider.ImportProviderState([]byte("forged-state")); err != nil { + t.Fatal(err) + } + if _, err := provider.Stream(context.Background(), llm.Request{Model: "model-a", Messages: []llm.Message{llm.UserText("x")}}); err == nil || !strings.Contains(err.Error(), "invalid_state") { + t.Fatalf("tampered state error = %v", err) + } + other, otherToken, err := fixture.clients.Add("satellite-b", Policy{AllowSearch: true, AllowFetch: true}) + if err != nil || other.ID == fixture.client.ID { + t.Fatal(err) + } + httpReq, err := http.NewRequest(http.MethodDelete, fixture.server.URL+"/g1/runs/not-owned", nil) + if err != nil { + t.Fatal(err) + } + httpReq.Header.Set("Authorization", "Bearer "+otherToken) + httpReq.Header.Set("Term-LLM-Gateway-Version", "1") + resp, err := http.DefaultClient.Do(httpReq) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if resp.StatusCode != http.StatusNotFound { + t.Fatalf("foreign run access status = %d, want 404", resp.StatusCode) + } +} + +func TestGatewayProviderStateRoundTripsSealed(t *testing.T) { + stateful := &statefulProvider{} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, stateful, time.Second) + provider, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + for i := 0; i < 2; i++ { + stream, err := provider.Stream(context.Background(), llm.Request{Model: "model-a", Messages: []llm.Message{llm.UserText("turn")}}) + if err != nil { + t.Fatal(err) + } + collectStream(t, stream) + } + stateful.mu.Lock() + defer stateful.mu.Unlock() + if stateful.imported != "state-1" { + t.Fatalf("gateway provider state import = %q, want state-1", stateful.imported) + } + sealed, ok := provider.ExportProviderState() + if !ok || strings.Contains(string(sealed), "state-2") { + t.Fatalf("satellite state is not opaque/sealed: %q, %t", sealed, ok) + } +} + +func receiveGatewayFailure(t *testing.T, stream llm.Stream) error { + t.Helper() + defer stream.Close() + for { + _, err := stream.Recv() + if err != nil { + return err + } + } +} + +func TestGatewayUpstreamPersistent500HonorsDefaultAttemptBudget(t *testing.T) { + failing := &setupFailureProvider{} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, failing, time.Second) + provider, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + started := time.Now() + stream, err := provider.Stream(t.Context(), llm.Request{Model: "model-a", Messages: []llm.Message{llm.UserText("fail")}}) + if err != nil { + t.Fatal(err) + } + if err := receiveGatewayFailure(t, stream); err == nil || !strings.Contains(err.Error(), "provider_upstream_failure") { + t.Fatalf("persistent gateway failure = %v", err) + } + if attempts := failing.attemptCount(); attempts != DefaultUpstreamRetryAttempts { + t.Fatalf("upstream attempts = %d, want %d", attempts, DefaultUpstreamRetryAttempts) + } + if elapsed := time.Since(started); elapsed > 6*time.Second { + t.Fatalf("persistent 500 took %s, want under 6s", elapsed) + } +} + +func TestGatewayUpstreamElapsedBudgetCancelsHungAttempt(t *testing.T) { + failing := &setupFailureProvider{block: true} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, failing, time.Second) + fixture.gateway.cfg.UpstreamRetryAttempts = 5 + fixture.gateway.cfg.UpstreamRetryMaxElapsed = 50 * time.Millisecond + provider, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + started := time.Now() + stream, err := provider.Stream(t.Context(), llm.Request{Model: "model-a", Messages: []llm.Message{llm.UserText("hang")}}) + if err != nil { + t.Fatal(err) + } + if err := receiveGatewayFailure(t, stream); err == nil { + t.Fatal("hung gateway upstream unexpectedly succeeded") + } + if attempts := failing.attemptCount(); attempts != 1 { + t.Fatalf("hung upstream attempts = %d, want 1", attempts) + } + if elapsed := time.Since(started); elapsed > 500*time.Millisecond { + t.Fatalf("hung upstream exceeded elapsed budget: %s", elapsed) + } +} + +func TestGatewayCancellationClosesCentralStream(t *testing.T) { + blocking := &blockingProvider{canceled: make(chan struct{})} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, blocking, time.Second) + provider, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + stream, err := provider.Stream(context.Background(), llm.Request{Model: "model-a", Messages: []llm.Message{llm.UserText("block")}}) + if err != nil { + t.Fatal(err) + } + if err := stream.Close(); err != nil { + t.Fatal(err) + } + select { + case <-blocking.canceled: + case <-time.After(time.Second): + t.Fatal("central stream was not canceled") + } + deadline := time.Now().Add(time.Second) + for { + release, ok := fixture.gateway.acquireInference(fixture.client) + if ok { + release() + break + } + if time.Now().After(deadline) { + t.Fatal("client inference permit was not released after cancellation") + } + time.Sleep(time.Millisecond) + } + for { + fixture.usage.mu.Lock() + if len(fixture.usage.records) > 0 { + record := fixture.usage.records[len(fixture.usage.records)-1] + fixture.usage.mu.Unlock() + if record.ErrorCode != "canceled" { + t.Fatalf("cancellation usage error = %q, want canceled", record.ErrorCode) + } + break + } + fixture.usage.mu.Unlock() + if time.Now().After(deadline) { + t.Fatal("cancellation usage was not recorded") + } + time.Sleep(time.Millisecond) + } +} + +func TestGatewayStreamDeathIsControlled(t *testing.T) { + dead := &streamDeathProvider{} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, dead, time.Second) + provider, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + stream, err := provider.Stream(context.Background(), llm.Request{Model: "model-a", Messages: []llm.Message{llm.UserText("die")}}) + if err != nil { + t.Fatal(err) + } + if event, err := stream.Recv(); err != nil || event.Text != "partial" { + t.Fatalf("partial event = %+v, %v", event, err) + } + if _, err := stream.Recv(); err == nil || !strings.Contains(err.Error(), "provider_upstream_failure") || !strings.Contains(err.Error(), "retry") || strings.Contains(err.Error(), "secret upstream") { + t.Fatalf("stream death error = %v", err) + } + _ = stream.Close() +} + +func TestGatewaySearchAndFetch(t *testing.T) { + mock := llm.NewMockProvider("central").AddTextResponse("unused") + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, mock, time.Second) + client, err := search.NewGatewayClient(fixture.satelliteConfig().Gateway) + if err != nil { + t.Fatal(err) + } + results, err := client.Search(context.Background(), "query", 3) + if err != nil || len(results) != 1 || results[0].Snippet != "central" { + t.Fatalf("search = %+v, %v", results, err) + } + content, err := client.FetchURL(context.Background(), "https://example.com/page") + if err != nil || content != "central fetch" { + t.Fatalf("fetch = %q, %v", content, err) + } +} diff --git a/internal/gateway/limits_test.go b/internal/gateway/limits_test.go new file mode 100644 index 000000000..3a70f5667 --- /dev/null +++ b/internal/gateway/limits_test.go @@ -0,0 +1,56 @@ +package gateway + +import "testing" + +func TestPerClientInferenceAndToolLimitsAreIsolatedAndReleased(t *testing.T) { + server := &Server{limits: make(map[string]*clientLimits)} + policy := Policy{ + MaxConcurrentInference: 1, + SearchRatePerMinute: 60, SearchBurst: 5, MaxConcurrentSearch: 1, + FetchRatePerMinute: 60, FetchBurst: 5, MaxConcurrentFetch: 1, + } + clientA := Client{ID: "a", Policy: policy} + clientB := Client{ID: "b", Policy: policy} + + releaseA, ok := server.acquireInference(clientA) + if !ok { + t.Fatal("first client A inference permit denied") + } + if _, ok := server.acquireInference(clientA); ok { + t.Fatal("client A exceeded inference concurrency cap") + } + releaseB, ok := server.acquireInference(clientB) + if !ok { + t.Fatal("client A usage leaked into client B inference cap") + } + releaseA() + if releaseAgain, ok := server.acquireInference(clientA); !ok { + t.Fatal("client A inference permit was not released") + } else { + releaseAgain() + } + releaseB() + + releaseSearchA, _, ok := server.acquireTool(clientA, true) + if !ok { + t.Fatal("first client A search permit denied") + } + if _, code, ok := server.acquireTool(clientA, true); ok || code != "search_concurrency_limited" { + t.Fatalf("second client A search = ok=%t code=%q", ok, code) + } + releaseSearchB, _, ok := server.acquireTool(clientB, true) + if !ok { + t.Fatal("client A search usage leaked into client B") + } + releaseSearchA() + releaseSearchB() + + releaseFetch, _, ok := server.acquireTool(clientA, false) + if !ok { + t.Fatal("first fetch permit denied") + } + if _, code, ok := server.acquireTool(clientA, false); ok || code != "fetch_concurrency_limited" { + t.Fatalf("second fetch = ok=%t code=%q", ok, code) + } + releaseFetch() +} diff --git a/internal/gateway/protocol/protocol.go b/internal/gateway/protocol/protocol.go new file mode 100644 index 000000000..f4f8b5589 --- /dev/null +++ b/internal/gateway/protocol/protocol.go @@ -0,0 +1,133 @@ +// Package protocol defines the private, versioned inference-gateway wire +// protocol. It intentionally contains no provider credentials or local paths. +package protocol + +import ( + "encoding/json" + "time" +) + +const ( + Version = 1 + VersionHeader = "Term-LLM-Gateway-Version" + BasePath = "/g1" +) + +type Error struct { + Code string `json:"code"` + Message string `json:"message,omitempty"` + RequestID string `json:"request_id,omitempty"` + SupportedVersions []int `json:"supported_versions,omitempty"` +} + +type InferenceRequest struct { + Version int `json:"version"` + RequestID string `json:"request_id"` + Provider string `json:"provider"` + State string `json:"state,omitempty"` + Request json.RawMessage `json:"request"` +} + +type StreamRecord struct { + Version int `json:"version"` + Type string `json:"type"` + RequestID string `json:"request_id,omitempty"` + RunID string `json:"run_id,omitempty"` + Event json.RawMessage `json:"event,omitempty"` + CallbackPath string `json:"callback_path,omitempty"` + State string `json:"state,omitempty"` + Error *Error `json:"error,omitempty"` +} + +type ToolResultRequest struct { + Version int `json:"version"` + Result json.RawMessage `json:"result"` +} + +type Health struct { + Version int `json:"version"` + Status string `json:"status"` +} + +type Catalog struct { + Version int `json:"version"` + GeneratedAt time.Time `json:"generated_at"` + Providers []CatalogEntry `json:"providers"` + Features CatalogFeatures `json:"features"` +} + +type CatalogFeatures struct { + Search bool `json:"search"` + Fetch bool `json:"fetch"` +} + +type CatalogEntry struct { + Key string `json:"key"` + Type string `json:"type"` + CLI bool `json:"cli,omitempty"` + AllowUnlistedModels bool `json:"allow_unlisted_models,omitempty"` + Capabilities Capabilities `json:"capabilities"` + Models []Model `json:"models"` +} + +type Capabilities struct { + NativeWebSearch bool `json:"native_web_search,omitempty"` + NativeWebFetch bool `json:"native_web_fetch,omitempty"` + ToolCalls bool `json:"tool_calls,omitempty"` + SupportsToolChoice bool `json:"supports_tool_choice,omitempty"` + ManagesOwnContext bool `json:"manages_own_context,omitempty"` + InlineToolLoop bool `json:"inline_tool_loop,omitempty"` + OrderedInlineToolEvents bool `json:"ordered_inline_tool_events,omitempty"` +} + +type Model struct { + ID string `json:"id"` + DisplayName string `json:"display_name,omitempty"` + Created int64 `json:"created,omitempty"` + OwnedBy string `json:"owned_by,omitempty"` + InputLimit int `json:"input_limit,omitempty"` + OutputLimit int `json:"output_limit,omitempty"` + InputPrice float64 `json:"input_price"` + OutputPrice float64 `json:"output_price"` + ReasoningEfforts []string `json:"reasoning_efforts,omitempty"` + DefaultReasoningEffort string `json:"default_reasoning_effort,omitempty"` + ReasoningModes []string `json:"reasoning_modes,omitempty"` +} + +type SearchRequest struct { + Version int `json:"version"` + Query string `json:"query"` + MaxResults int `json:"max_results"` +} + +type SearchResult struct { + Title string `json:"title"` + URL string `json:"url"` + Snippet string `json:"snippet,omitempty"` +} + +type SearchResponse struct { + Version int `json:"version"` + Results []SearchResult `json:"results"` +} + +type FetchRequest struct { + Version int `json:"version"` + URL string `json:"url"` +} + +type FetchResponse struct { + Version int `json:"version"` + Content string `json:"content"` +} + +type EnrollmentRequest struct { + Version int `json:"version"` + Name string `json:"name"` +} + +type EnrollmentResponse struct { + Version int `json:"version"` + ClientID string `json:"client_id"` + Token string `json:"token"` +} diff --git a/internal/gateway/responses.go b/internal/gateway/responses.go new file mode 100644 index 000000000..bc736048d --- /dev/null +++ b/internal/gateway/responses.go @@ -0,0 +1,695 @@ +package gateway + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strings" + "time" + + "github.com/samsaffron/term-llm/internal/gateway/protocol" + "github.com/samsaffron/term-llm/internal/llm" +) + +type responsesOutputContent struct { + Type string `json:"type"` + Text string `json:"text"` + Annotations []any `json:"annotations"` +} + +type responsesSummaryPart struct { + Type string `json:"type"` + Text string `json:"text"` +} + +type responsesOutputItem struct { + ID string + Type string + Status string + Role string + Content []responsesOutputContent + EncryptedContent string + Summary []responsesSummaryPart + CallID string + Name string + Arguments string +} + +func (item responsesOutputItem) MarshalJSON() ([]byte, error) { + wire := map[string]any{"id": item.ID, "type": item.Type} + if item.Status != "" { + wire["status"] = item.Status + } + switch item.Type { + case "reasoning": + wire["encrypted_content"] = item.EncryptedContent + wire["summary"] = item.Summary + case "message": + wire["role"] = item.Role + wire["content"] = item.Content + case "function_call": + wire["call_id"] = item.CallID + wire["name"] = item.Name + wire["arguments"] = item.Arguments + } + return json.Marshal(wire) +} + +type responsesUsageDetails struct { + CachedTokens int `json:"cached_tokens"` +} + +type responsesOutputUsageDetails struct { + ReasoningTokens int `json:"reasoning_tokens"` +} + +type responsesUsage struct { + InputTokens int `json:"input_tokens"` + InputTokensDetails responsesUsageDetails `json:"input_tokens_details"` + OutputTokens int `json:"output_tokens"` + OutputTokensDetails responsesOutputUsageDetails `json:"output_tokens_details"` + TotalTokens int `json:"total_tokens"` +} + +type responsesDocument struct { + ID string `json:"id"` + Object string `json:"object"` + CreatedAt int64 `json:"created_at"` + Status string `json:"status"` + Error any `json:"error"` + IncompleteDetails any `json:"incomplete_details"` + Instructions any `json:"instructions"` + Metadata map[string]string `json:"metadata"` + Model string `json:"model"` + Output []*responsesOutputItem `json:"output"` + ParallelToolCalls bool `json:"parallel_tool_calls"` + Temperature float64 `json:"temperature"` + ToolChoice any `json:"tool_choice"` + Tools []json.RawMessage `json:"tools"` + TopP float64 `json:"top_p"` + ServiceTier string `json:"service_tier,omitempty"` + Usage *responsesUsage `json:"usage"` +} + +type responsesReasoningOutput struct { + item *responsesOutputItem + outputIndex int + summary string + summaryParts []string + partAdded bool + done bool +} + +type responsesMessageOutput struct { + item *responsesOutputItem + outputIndex int + text string + done bool +} + +type responsesAccumulator struct { + document responsesDocument + emit func(string, map[string]any) error + sequence int + reasoning *responsesReasoningOutput + message *responsesMessageOutput + usage llm.Usage +} + +func newResponsesAccumulator(responseID string, request responsesRequest, parallel bool, emit func(string, map[string]any) error) *responsesAccumulator { + var instructions any + if request.Instructions != nil { + instructions = *request.Instructions + } + metadata := request.Metadata + if metadata == nil { + metadata = make(map[string]string) + } + temperature := 1.0 + if request.Temperature != nil { + temperature = *request.Temperature + } + topP := 1.0 + if request.TopP != nil { + topP = *request.TopP + } + toolChoice := any("auto") + if raw := bytes.TrimSpace(request.ToolChoice); len(raw) > 0 && !bytes.Equal(raw, []byte("null")) { + toolChoice = json.RawMessage(append([]byte(nil), raw...)) + } + tools := request.Tools + if tools == nil { + tools = make([]json.RawMessage, 0) + } + serviceTier := "" + if request.ServiceTier != nil { + serviceTier = strings.TrimSpace(*request.ServiceTier) + } + return &responsesAccumulator{ + document: responsesDocument{ + ID: responseID, Object: "response", CreatedAt: time.Now().Unix(), Status: "in_progress", + Error: nil, IncompleteDetails: nil, Instructions: instructions, Metadata: metadata, + Model: request.Model, Output: make([]*responsesOutputItem, 0), ParallelToolCalls: parallel, + Temperature: temperature, ToolChoice: toolChoice, Tools: tools, TopP: topP, ServiceTier: serviceTier, + }, + emit: emit, + } +} + +func (a *responsesAccumulator) event(eventType string, payload map[string]any) error { + if a.emit == nil { + return nil + } + if payload == nil { + payload = make(map[string]any) + } + payload["type"] = eventType + payload["sequence_number"] = a.sequence + a.sequence++ + return a.emit(eventType, payload) +} + +func (a *responsesAccumulator) created() error { + response := a.document + response.Output = []*responsesOutputItem{} + response.Usage = nil + if err := a.event("response.created", map[string]any{"response": response}); err != nil { + return err + } + return a.event("response.in_progress", map[string]any{"response": response}) +} + +func (a *responsesAccumulator) consume(event llm.Event) error { + switch event.Type { + case llm.EventTextDelta: + return a.addText(event.Text) + case llm.EventReasoningDelta: + return a.addReasoning(event) + case llm.EventToolCall: + if event.Tool == nil { + return nil + } + if event.ToolResponse != nil { + return &responsesWireError{Status: http.StatusBadRequest, Code: "incompatible_tool_request", Message: "provider requested gateway-side tool execution, which the Responses edge never permits"} + } + return a.addFunctionCall(*event.Tool) + case llm.EventUsage: + if event.Use != nil { + a.usage.Add(*event.Use) + } + case llm.EventError: + if event.Err != nil { + return event.Err + } + return fmt.Errorf("provider returned an error event") + case llm.EventAttemptDiscard: + if a.message != nil || a.reasoning != nil || len(a.document.Output) > 0 { + return fmt.Errorf("provider retry attempted after Responses output was committed") + } + } + return nil +} + +func (a *responsesAccumulator) addReasoning(event llm.Event) error { + displaySummary := llm.NormalizeReasoningKind(event.ReasoningKind) == llm.ReasoningKindSummary || len(event.ReasoningSummaryParts) > 0 + if !displaySummary && event.ReasoningEncryptedContent == "" { + return nil + } + if err := a.finishMessage(); err != nil { + return err + } + itemID := strings.TrimSpace(event.ReasoningItemID) + if a.reasoning != nil && itemID != "" && a.reasoning.item.ID != itemID { + if err := a.finishReasoning(); err != nil { + return err + } + } + if a.reasoning == nil { + if itemID == "" { + var err error + itemID, err = randomSecret("rs", 16) + if err != nil { + return err + } + } + item := &responsesOutputItem{ID: itemID, Type: "reasoning", Status: "in_progress", EncryptedContent: event.ReasoningEncryptedContent, Summary: []responsesSummaryPart{}} + index := len(a.document.Output) + a.document.Output = append(a.document.Output, item) + a.reasoning = &responsesReasoningOutput{item: item, outputIndex: index} + if err := a.event("response.output_item.added", map[string]any{"output_index": index, "item": item}); err != nil { + return err + } + } + state := a.reasoning + if event.ReasoningEncryptedContent != "" { + state.item.EncryptedContent = event.ReasoningEncryptedContent + } + if len(event.ReasoningSummaryParts) > 0 { + state.summaryParts = append([]string(nil), event.ReasoningSummaryParts...) + } + if displaySummary && event.Text != "" { + if !state.partAdded { + state.partAdded = true + if err := a.event("response.reasoning_summary_part.added", map[string]any{ + "item_id": state.item.ID, "output_index": state.outputIndex, "summary_index": 0, + "part": responsesSummaryPart{Type: "summary_text", Text: ""}, + }); err != nil { + return err + } + } + state.summary += event.Text + if err := a.event("response.reasoning_summary_text.delta", map[string]any{ + "item_id": state.item.ID, "output_index": state.outputIndex, "summary_index": 0, "delta": event.Text, + }); err != nil { + return err + } + } + if event.ReasoningFinal { + return a.finishReasoning() + } + return nil +} + +func (a *responsesAccumulator) finishReasoning() error { + state := a.reasoning + if state == nil || state.done { + return nil + } + state.done = true + if len(state.summaryParts) > 0 { + state.item.Summary = make([]responsesSummaryPart, 0, len(state.summaryParts)) + for _, text := range state.summaryParts { + state.item.Summary = append(state.item.Summary, responsesSummaryPart{Type: "summary_text", Text: text}) + } + } else if state.summary != "" { + state.item.Summary = []responsesSummaryPart{{Type: "summary_text", Text: state.summary}} + } + if state.partAdded { + if err := a.event("response.reasoning_summary_text.done", map[string]any{ + "item_id": state.item.ID, "output_index": state.outputIndex, "summary_index": 0, "text": state.summary, + }); err != nil { + return err + } + if err := a.event("response.reasoning_summary_part.done", map[string]any{ + "item_id": state.item.ID, "output_index": state.outputIndex, "summary_index": 0, + "part": responsesSummaryPart{Type: "summary_text", Text: state.summary}, + }); err != nil { + return err + } + } + state.item.Status = "completed" + if err := a.event("response.output_item.done", map[string]any{"output_index": state.outputIndex, "item": state.item}); err != nil { + return err + } + a.reasoning = nil + return nil +} + +func (a *responsesAccumulator) addText(delta string) error { + if delta == "" { + return nil + } + if err := a.finishReasoning(); err != nil { + return err + } + if a.message == nil { + itemID, err := randomSecret("msg", 16) + if err != nil { + return err + } + item := &responsesOutputItem{ID: itemID, Type: "message", Status: "in_progress", Role: "assistant", Content: []responsesOutputContent{}} + index := len(a.document.Output) + a.document.Output = append(a.document.Output, item) + a.message = &responsesMessageOutput{item: item, outputIndex: index} + if err := a.event("response.output_item.added", map[string]any{"output_index": index, "item": item}); err != nil { + return err + } + if err := a.event("response.content_part.added", map[string]any{ + "item_id": item.ID, "output_index": index, "content_index": 0, + "part": responsesOutputContent{Type: "output_text", Text: "", Annotations: []any{}}, + }); err != nil { + return err + } + } + a.message.text += delta + return a.event("response.output_text.delta", map[string]any{ + "item_id": a.message.item.ID, "output_index": a.message.outputIndex, "content_index": 0, + "delta": delta, "logprobs": []any{}, + }) +} + +func (a *responsesAccumulator) finishMessage() error { + state := a.message + if state == nil || state.done { + return nil + } + state.done = true + content := responsesOutputContent{Type: "output_text", Text: state.text, Annotations: []any{}} + state.item.Content = []responsesOutputContent{content} + if err := a.event("response.output_text.done", map[string]any{ + "item_id": state.item.ID, "output_index": state.outputIndex, "content_index": 0, + "text": state.text, "logprobs": []any{}, + }); err != nil { + return err + } + if err := a.event("response.content_part.done", map[string]any{ + "item_id": state.item.ID, "output_index": state.outputIndex, "content_index": 0, "part": content, + }); err != nil { + return err + } + state.item.Status = "completed" + if err := a.event("response.output_item.done", map[string]any{"output_index": state.outputIndex, "item": state.item}); err != nil { + return err + } + a.message = nil + return nil +} + +func (a *responsesAccumulator) addFunctionCall(call llm.ToolCall) error { + if err := a.finishReasoning(); err != nil { + return err + } + if err := a.finishMessage(); err != nil { + return err + } + callID := strings.TrimSpace(call.ID) + if callID == "" { + var err error + callID, err = randomSecret("call", 16) + if err != nil { + return err + } + } + itemID, err := randomSecret("fc", 16) + if err != nil { + return err + } + arguments := strings.TrimSpace(string(call.Arguments)) + if arguments == "" { + arguments = "{}" + } + item := &responsesOutputItem{ID: itemID, Type: "function_call", Status: "in_progress", CallID: callID, Name: call.Name, Arguments: ""} + index := len(a.document.Output) + a.document.Output = append(a.document.Output, item) + if err := a.event("response.output_item.added", map[string]any{"output_index": index, "item": item}); err != nil { + return err + } + if err := a.event("response.function_call_arguments.delta", map[string]any{ + "item_id": itemID, "output_index": index, "delta": arguments, + }); err != nil { + return err + } + if err := a.event("response.function_call_arguments.done", map[string]any{ + "item_id": itemID, "output_index": index, "arguments": arguments, + }); err != nil { + return err + } + item.Status = "completed" + item.Arguments = arguments + return a.event("response.output_item.done", map[string]any{"output_index": index, "item": item}) +} + +func (a *responsesAccumulator) complete() error { + if err := a.finishReasoning(); err != nil { + return err + } + if err := a.finishMessage(); err != nil { + return err + } + inputTokens := a.usage.InputTokens + a.usage.CachedInputTokens + a.usage.CacheWriteTokens + totalTokens := a.usage.ProviderTotalTokens + if totalTokens <= 0 { + totalTokens = inputTokens + a.usage.OutputTokens + } + a.document.Status = "completed" + a.document.Usage = &responsesUsage{ + InputTokens: inputTokens, InputTokensDetails: responsesUsageDetails{CachedTokens: a.usage.CachedInputTokens}, + OutputTokens: a.usage.OutputTokens, OutputTokensDetails: responsesOutputUsageDetails{ReasoningTokens: a.usage.ReasoningTokens}, + TotalTokens: totalTokens, + } + return a.event("response.completed", map[string]any{"response": &a.document}) +} + +func (a *responsesAccumulator) fail(code, message, param string) error { + if err := a.event("error", map[string]any{ + "code": code, "message": message, "param": nullableString(param), + }); err != nil { + return err + } + a.document.Status = "failed" + a.document.Error = map[string]any{"code": code, "message": message} + return a.event("response.failed", map[string]any{"response": &a.document}) +} + +func (s *Server) handleResponses(w http.ResponseWriter, r *http.Request, client Client) { + request, decodeErr := s.decodeResponsesRequest(r) + if decodeErr != nil { + s.writeResponsesError(w, decodeErr.Status, decodeErr.Code, decodeErr.Message, decodeErr.Param) + return + } + provider, model, namespaceErr := splitResponsesModel(request.Model) + if namespaceErr != nil { + s.writeResponsesError(w, namespaceErr.Status, namespaceErr.Code, namespaceErr.Message, namespaceErr.Param) + return + } + providerReq, translateErr := translateResponsesRequest(request) + if translateErr != nil { + s.writeResponsesError(w, translateErr.Status, translateErr.Code, translateErr.Message, translateErr.Param) + return + } + providerReq.Model = model + requestID, err := randomSecret("req", 16) + if err != nil { + s.writeResponsesError(w, http.StatusInternalServerError, "internal", "could not create gateway request", "") + return + } + envelope := protocol.InferenceRequest{Version: protocol.Version, RequestID: requestID, Provider: provider} + execution, requestErr := s.startInference(r.Context(), client, envelope, providerReq, true) + if requestErr != nil { + s.writeResponsesError(w, requestErr.Status, requestErr.Code, requestErr.Message, responsesErrorParam(requestErr.Code)) + return + } + errorCode := "" + defer func() { execution.close(errorCode) }() + + responseID, err := randomSecret("resp", 16) + if err != nil { + errorCode = "internal" + s.writeResponsesError(w, http.StatusInternalServerError, "internal", "could not create response", "") + return + } + var flusher http.Flusher + if request.Stream { + var ok bool + flusher, ok = w.(http.Flusher) + if !ok { + errorCode = "streaming_unsupported" + s.writeResponsesError(w, http.StatusInternalServerError, errorCode, "streaming is unavailable", "stream") + return + } + } + accumulator := newResponsesAccumulator(responseID, request, providerReq.ParallelToolCalls, nil) + var prefetched *llm.Event + streamEOF := false + if request.Stream { + event, recvErr := prefetchResponsesStream(execution.stream) + if recvErr == io.EOF { + streamEOF = true + } else if recvErr != nil { + status, code := s.logProviderStreamError(execution, recvErr) + errorCode = code + s.writeResponsesError(w, status, code, safeProviderErrorMessage(code, provider), responsesErrorParam(code)) + return + } else { + prefetched = &event + } + w.Header().Set("Content-Type", "text/event-stream") + w.Header().Set("Cache-Control", "no-cache, no-transform") + w.Header().Set("X-Accel-Buffering", "no") + w.WriteHeader(http.StatusOK) + accumulator.emit = func(eventType string, payload map[string]any) error { + if err := writeResponsesSSE(w, payload); err != nil { + return context.Canceled + } + flusher.Flush() + return nil + } + if err := accumulator.created(); err != nil { + errorCode = "canceled" + return + } + } + + for !streamEOF { + var event llm.Event + var recvErr error + if prefetched != nil { + event = *prefetched + prefetched = nil + } else { + event, recvErr = execution.stream.Recv() + } + if recvErr == io.EOF { + break + } + if recvErr != nil { + status, code := s.logProviderStreamError(execution, recvErr) + errorCode = code + if request.Stream { + _ = accumulator.fail(code, safeProviderErrorMessage(code, provider), responsesErrorParam(code)) + return + } + s.writeResponsesError(w, status, code, safeProviderErrorMessage(code, provider), responsesErrorParam(code)) + return + } + if event.Type == llm.EventUsage && event.Use != nil { + execution.addUsage(*event.Use) + } + if err := accumulator.consume(event); err != nil { + if errors.Is(err, context.Canceled) || errors.Is(r.Context().Err(), context.Canceled) { + errorCode = "canceled" + return + } + var wireErr *responsesWireError + status, code := classifyProviderError(err, execution.entry.Type) + if errors.As(err, &wireErr) { + status, code = wireErr.Status, wireErr.Code + } else { + status, code = s.logProviderStreamError(execution, err) + } + errorCode = code + message := safeProviderErrorMessage(code, provider) + param := responsesErrorParam(code) + if wireErr != nil { + message = wireErr.Message + param = wireErr.Param + } + if request.Stream { + _ = accumulator.fail(code, message, param) + return + } + s.writeResponsesError(w, status, code, message, param) + return + } + } + if err := accumulator.complete(); err != nil { + errorCode = "canceled" + return + } + if !request.Stream { + s.writeResponsesJSON(w, http.StatusOK, &accumulator.document) + } +} + +func responsesErrorParam(code string) string { + switch code { + case "unknown_provider", "unknown_model", "policy_denied": + return "model" + default: + return "" + } +} + +func (s *Server) handleResponsesModels(w http.ResponseWriter, r *http.Request, client Client) { + catalog, err := s.currentCatalog(r.Context()) + if err != nil { + s.writeResponsesError(w, http.StatusServiceUnavailable, "catalog_unavailable", "gateway model catalog is temporarily unavailable; retry or contact the gateway operator", "") + return + } + catalog = filterCatalog(catalog, s.cfg.Policy, client.Policy) + data := make([]map[string]any, 0) + for _, entry := range catalog.Providers { + for _, model := range entry.Models { + item := map[string]any{ + "id": entry.Key + "/" + model.ID, "object": "model", "created": model.Created, + "owned_by": model.OwnedBy, "provider": entry.Key, + } + if item["owned_by"] == "" { + item["owned_by"] = entry.Key + } + if model.DisplayName != "" { + item["display_name"] = model.DisplayName + } + if model.InputLimit > 0 { + item["input_limit"] = model.InputLimit + } + if model.OutputLimit > 0 { + item["output_limit"] = model.OutputLimit + } + item["input_price"] = model.InputPrice + item["output_price"] = model.OutputPrice + if len(model.ReasoningEfforts) > 0 { + item["reasoning_efforts"] = model.ReasoningEfforts + } + if model.DefaultReasoningEffort != "" { + item["default_reasoning_effort"] = model.DefaultReasoningEffort + } + if len(model.ReasoningModes) > 0 { + item["reasoning_modes"] = model.ReasoningModes + } + data = append(data, item) + } + } + s.writeResponsesJSON(w, http.StatusOK, map[string]any{"object": "list", "data": data}) +} + +func (s *Server) writeResponsesError(w http.ResponseWriter, status int, code, message, param string) { + errorType := "gateway_error" + if status == http.StatusBadRequest || status == http.StatusNotFound || status == http.StatusForbidden { + errorType = "invalid_request_error" + } else if status == http.StatusUnauthorized { + errorType = "authentication_error" + } else if status == http.StatusTooManyRequests { + errorType = "rate_limit_error" + } + s.writeResponsesJSON(w, status, map[string]any{"error": map[string]any{ + "message": message, "type": errorType, "param": nullableString(param), "code": code, + }}) +} + +func nullableString(value string) any { + if value == "" { + return nil + } + return value +} + +func (s *Server) writeResponsesJSON(w http.ResponseWriter, status int, value any) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(status) + _ = json.NewEncoder(w).Encode(value) +} + +func writeResponsesSSE(w io.Writer, payload map[string]any) error { + data, err := json.Marshal(payload) + if err != nil { + return err + } + _, err = fmt.Fprintf(w, "data: %s\n\n", data) + return err +} + +func prefetchResponsesStream(stream llm.Stream) (llm.Event, error) { + for { + event, err := stream.Recv() + if err != nil { + return llm.Event{}, err + } + if event.Type == llm.EventRetry { + // Retry notifications are not Responses output and do not prove that an + // upstream attempt started. Keep waiting without buffering output. + continue + } + if event.Type == llm.EventError { + if event.Err != nil { + return llm.Event{}, event.Err + } + return llm.Event{}, fmt.Errorf("provider returned an error event") + } + return event, nil + } +} diff --git a/internal/gateway/responses_integration_test.go b/internal/gateway/responses_integration_test.go new file mode 100644 index 000000000..9d8ac2039 --- /dev/null +++ b/internal/gateway/responses_integration_test.go @@ -0,0 +1,816 @@ +package gateway + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "path/filepath" + "strings" + "sync" + "testing" + "time" + + openai "github.com/openai/openai-go" + "github.com/openai/openai-go/option" + openairesponses "github.com/openai/openai-go/responses" + "github.com/samsaffron/term-llm/internal/config" + "github.com/samsaffron/term-llm/internal/llm" + "github.com/samsaffron/term-llm/internal/providerhttp" +) + +type responsesFixtureProvider struct { + mu sync.Mutex + requests []llm.Request + events [][]llm.Event + errors []error + started chan struct{} + canceled chan struct{} +} + +func (*responsesFixtureProvider) Name() string { return "responses-fixture" } +func (*responsesFixtureProvider) Credential() string { return "mock" } +func (*responsesFixtureProvider) Capabilities() llm.Capabilities { + return llm.Capabilities{ToolCalls: true, SupportsToolChoice: true} +} +func (p *responsesFixtureProvider) Stream(ctx context.Context, req llm.Request) (llm.Stream, error) { + p.mu.Lock() + p.requests = append(p.requests, req) + index := len(p.requests) - 1 + var events []llm.Event + if index < len(p.events) { + events = append([]llm.Event(nil), p.events[index]...) + } + var terminalErr error + if index < len(p.errors) { + terminalErr = p.errors[index] + } + started := p.started + canceled := p.canceled + p.mu.Unlock() + if started != nil { + select { + case <-started: + default: + close(started) + } + return &responsesBlockingStream{ctx: ctx, canceled: canceled}, nil + } + return &responsesSliceStream{events: events, terminalErr: terminalErr}, nil +} +func (p *responsesFixtureProvider) recorded() []llm.Request { + p.mu.Lock() + defer p.mu.Unlock() + return append([]llm.Request(nil), p.requests...) +} + +type responsesSliceStream struct { + events []llm.Event + index int + terminalErr error +} + +func (s *responsesSliceStream) Recv() (llm.Event, error) { + if s.index >= len(s.events) { + if s.terminalErr != nil { + err := s.terminalErr + s.terminalErr = nil + return llm.Event{}, err + } + return llm.Event{}, io.EOF + } + event := s.events[s.index] + s.index++ + return event, nil +} +func (*responsesSliceStream) Close() error { return nil } + +type responsesBlockingStream struct { + ctx context.Context + canceled chan struct{} + once sync.Once +} + +func (s *responsesBlockingStream) Recv() (llm.Event, error) { + <-s.ctx.Done() + s.once.Do(func() { + if s.canceled != nil { + close(s.canceled) + } + }) + return llm.Event{}, s.ctx.Err() +} +func (s *responsesBlockingStream) Close() error { + s.once.Do(func() { + if s.canceled != nil { + close(s.canceled) + } + }) + return nil +} + +func TestResponsesOfficialOpenAIClientsNonStreamingStreamingAndModels(t *testing.T) { + provider := &responsesFixtureProvider{events: [][]llm.Event{ + { + {Type: llm.EventTextDelta, Text: "hello "}, + {Type: llm.EventTextDelta, Text: "client"}, + {Type: llm.EventUsage, Use: &llm.Usage{InputTokens: 7, CachedInputTokens: 2, OutputTokens: 3, ProviderTotalTokens: 12}}, + {Type: llm.EventDone}, + }, + { + {Type: llm.EventTextDelta, Text: "streamed"}, + {Type: llm.EventDone}, + }, + }} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, provider, time.Second) + client := openai.NewClient(option.WithAPIKey(fixture.token), option.WithBaseURL(fixture.server.URL+"/v1")) + + response, err := client.Responses.New(t.Context(), openairesponses.ResponseNewParams{ + Model: "remote/model-a", + Input: openairesponses.ResponseNewParamsInputUnion{OfString: openai.String("hello")}, + Instructions: openai.String("official instructions"), + Metadata: map[string]string{"source": "official-client"}, + Temperature: openai.Float(0.25), + TopP: openai.Float(0.75), + }) + if err != nil { + t.Fatalf("official Responses client: %v", err) + } + if got := response.OutputText(); got != "hello client" { + t.Fatalf("official client output = %q", got) + } + if response.Usage.InputTokens != 9 || response.Usage.InputTokensDetails.CachedTokens != 2 || response.Usage.OutputTokens != 3 || response.Usage.TotalTokens != 12 { + t.Fatalf("official client usage = %+v", response.Usage) + } + assertOfficialResponsesDocument(t, *response) + if response.Instructions.AsString() != "official instructions" || response.Metadata["source"] != "official-client" || response.Temperature != 0.25 || response.TopP != 0.75 || response.ToolChoice.AsToolChoiceMode() != "auto" || len(response.Tools) != 0 || !response.ParallelToolCalls { + t.Fatalf("official client required response fields = %+v", response) + } + stream := client.Responses.NewStreaming(t.Context(), openairesponses.ResponseNewParams{ + Model: "remote/model-a", + Input: openairesponses.ResponseNewParamsInputUnion{OfString: openai.String("stream")}, + }) + defer stream.Close() + var streamTypes []string + var streamedText strings.Builder + var sawCreated, sawInProgress, sawCompleted, sawTextDone bool + for stream.Next() { + current := stream.Current() + streamTypes = append(streamTypes, current.Type) + switch event := current.AsAny().(type) { + case openairesponses.ResponseCreatedEvent: + sawCreated = true + assertOfficialResponsesDocument(t, event.Response) + assertOfficialResponsesDefaults(t, event.Response) + if !event.JSON.Response.Valid() || !event.JSON.SequenceNumber.Valid() || !event.JSON.Type.Valid() { + t.Fatalf("created event JSON validity = %+v", event.JSON) + } + case openairesponses.ResponseInProgressEvent: + sawInProgress = true + assertOfficialResponsesDocument(t, event.Response) + assertOfficialResponsesDefaults(t, event.Response) + case openairesponses.ResponseTextDeltaEvent: + streamedText.WriteString(event.Delta) + if !event.JSON.Logprobs.Valid() || event.Logprobs == nil { + t.Fatalf("text delta logprobs missing: raw=%s", event.RawJSON()) + } + case openairesponses.ResponseTextDoneEvent: + sawTextDone = true + if !event.JSON.Logprobs.Valid() || event.Logprobs == nil || event.Text != "streamed" { + t.Fatalf("text done event = %+v raw=%s", event, event.RawJSON()) + } + case openairesponses.ResponseCompletedEvent: + sawCompleted = true + assertOfficialResponsesDocument(t, event.Response) + assertOfficialResponsesDefaults(t, event.Response) + if event.Response.OutputText() != "streamed" || event.Response.Status != "completed" { + t.Fatalf("typed completed response = %+v", event.Response) + } + } + } + if err := stream.Err(); err != nil { + t.Fatalf("official streaming Responses client: %v", err) + } + if !sawCreated || !sawInProgress || !sawTextDone || !sawCompleted || streamedText.String() != "streamed" || !responsesEventTypesContain(streamTypes, "response.output_text.delta") || len(streamTypes) == 0 || streamTypes[len(streamTypes)-1] != "response.completed" { + t.Fatalf("official streaming events = types:%v text:%q typed:%t/%t/%t/%t", streamTypes, streamedText.String(), sawCreated, sawInProgress, sawTextDone, sawCompleted) + } + page, err := client.Models.List(t.Context()) + if err != nil { + t.Fatalf("official models client: %v", err) + } + if len(page.Data) != 1 || page.Data[0].ID != "remote/model-a" { + t.Fatalf("official models = %+v", page.Data) + } + + requests := provider.recorded() + if len(requests) != 2 || requests[0].Model != "model-a" || len(requests[0].Messages) != 2 || llm.MessageText(requests[0].Messages[0]) != "official instructions" || llm.MessageText(requests[0].Messages[1]) != "hello" || llm.MessageText(requests[1].Messages[0]) != "stream" { + t.Fatalf("translated official requests = %+v", requests) + } + fixture.usage.mu.Lock() + defer fixture.usage.mu.Unlock() + if len(fixture.usage.records) != 2 || fixture.usage.records[0].ClientID != fixture.client.ID || fixture.usage.records[0].ProviderKey != "remote" || fixture.usage.records[0].InputTokens != 7 || fixture.usage.records[0].CachedInputTokens != 2 { + t.Fatalf("Responses usage attribution = %+v", fixture.usage.records) + } +} + +func TestResponsesDiscourseShapedStreamRejectsPreOutputProviderFailures(t *testing.T) { + tests := []struct { + name string + err error + wantStatus int + wantCode string + }{ + { + name: "authentication", + err: providerhttp.NewStatusErrorString("OpenAI", http.StatusUnauthorized, "401 Unauthorized", nil, `api_key=super-secret-provider-key`), + wantStatus: http.StatusUnauthorized, + wantCode: "provider_api_key_unauthenticated", + }, + { + name: "rate limit", + err: &llm.RateLimitError{Message: "raw provider quota detail"}, + wantStatus: http.StatusTooManyRequests, + wantCode: "provider_rate_limited", + }, + { + name: "upstream", + err: providerhttp.NewStatusErrorString("OpenAI", http.StatusServiceUnavailable, "503 Service Unavailable", nil, `raw upstream body /srv/private`), + wantStatus: http.StatusBadGateway, + wantCode: "provider_upstream_failure", + }, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + provider := &responsesFixtureProvider{errors: []error{tc.err}} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, provider, time.Second) + fixture.gateway.cfg.UpstreamRetryAttempts = 1 + response := doResponsesRequest(t, fixture, map[string]any{ + "model": "remote/model-a", "input": "hello", "stream": true, + }) + if response.StatusCode != tc.wantStatus { + body, _ := io.ReadAll(response.Body) + response.Body.Close() + t.Fatalf("pre-output failure status = %d body=%s, want %d", response.StatusCode, body, tc.wantStatus) + } + if output, err := discourseShapedResponsesStream(response); err == nil { + t.Fatalf("Discourse-shaped client accepted failed stream as successful output %q", output) + } else if !strings.Contains(err.Error(), tc.wantCode) { + t.Fatalf("Discourse-shaped error = %v, want code %s", err, tc.wantCode) + } + }) + } +} + +func TestResponsesCommittedStreamUsesStandardErrorAndFailedEvents(t *testing.T) { + provider := &responsesFixtureProvider{ + events: [][]llm.Event{{{Type: llm.EventTextDelta, Text: "partial"}}}, + errors: []error{providerhttp.NewStatusErrorString("OpenAI", http.StatusServiceUnavailable, "503 Service Unavailable", nil, `raw body super-secret-provider-key /srv/private`)}, + } + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, provider, time.Second) + fixture.gateway.cfg.UpstreamRetryAttempts = 1 + response := doResponsesRequest(t, fixture, map[string]any{ + "model": "remote/model-a", "input": "hello", "stream": true, + "instructions": "keep this", "metadata": map[string]string{"client": "discourse"}, + "temperature": 0.4, "top_p": 0.6, + }) + defer response.Body.Close() + if response.StatusCode != http.StatusOK { + body, _ := io.ReadAll(response.Body) + t.Fatalf("committed stream status = %d: %s", response.StatusCode, body) + } + events := readResponsesEvents(t, response.Body) + assertResponsesEventTypes(t, events, "response.created", "response.output_text.delta", "error", "response.failed") + errorEvent := findResponsesEvent(t, events, "error") + if _, nested := errorEvent["error"]; nested || errorEvent["code"] != "provider_upstream_failure" || errorEvent["message"] == "" { + t.Fatalf("Responses error event shape = %#v", errorEvent) + } + if _, present := errorEvent["param"]; !present { + t.Fatalf("Responses error event omitted required param: %#v", errorEvent) + } + failed := findResponsesEvent(t, events, "response.failed") + failedResponse := failed["response"].(map[string]any) + for _, field := range []string{"id", "object", "created_at", "status", "error", "incomplete_details", "instructions", "metadata", "model", "output", "parallel_tool_calls", "temperature", "tool_choice", "tools", "top_p"} { + if _, present := failedResponse[field]; !present { + t.Fatalf("failed response omitted %q: %#v", field, failedResponse) + } + } + apiError := failedResponse["error"].(map[string]any) + if failedResponse["status"] != "failed" || apiError["code"] != "provider_upstream_failure" || apiError["message"] == "" { + t.Fatalf("failed response shape = %#v", failedResponse) + } + encoded, _ := json.Marshal(events) + for _, secret := range []string{"super-secret-provider-key", "raw body", "/srv/private"} { + if bytes.Contains(encoded, []byte(secret)) { + t.Fatalf("stream error leaked central diagnostic %q: %s", secret, encoded) + } + } +} + +func TestResponsesDiscourseStreamingReasoningMultimodalAndFunctionCall(t *testing.T) { + provider := &responsesFixtureProvider{events: [][]llm.Event{{ + {Type: llm.EventReasoningDelta, Text: "**Thinking**", ReasoningKind: llm.ReasoningKindSummary, ReasoningItemID: "rs_provider", ReasoningEncryptedContent: "ENC", ReasoningSummaryParts: []string{"**Thinking**"}, ReasoningFinal: true}, + {Type: llm.EventTextDelta, Text: "answer"}, + {Type: llm.EventToolCall, Tool: &llm.ToolCall{ID: "call_external", Name: "echo", Arguments: json.RawMessage(`{"string":"hello"}`)}}, + {Type: llm.EventUsage, Use: &llm.Usage{InputTokens: 10, CachedInputTokens: 4, OutputTokens: 6, ReasoningTokens: 2}}, + {Type: llm.EventDone}, + }}} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, provider, time.Second) + payload := map[string]any{ + "model": "remote/model-a", "stream": true, "max_output_tokens": 123, + "reasoning": map[string]any{"summary": "auto", "effort": "high"}, + "include": []string{"reasoning.encrypted_content"}, "temperature": 0.2, "top_p": 0.8, + "service_tier": "priority", "parallel_tool_calls": true, + "input": []any{ + map[string]any{"role": "developer", "content": "policy"}, + map[string]any{"role": "user", "content": []any{ + map[string]any{"type": "input_text", "text": "inspect"}, + map[string]any{"type": "input_image", "image_url": "data:image/png;base64,aW1n"}, + map[string]any{"type": "input_file", "filename": "notes.txt", "file_data": "data:text/plain;base64,ZmlsZQ=="}, + }}, + }, + "tools": []any{map[string]any{"type": "function", "name": "echo", "description": "echo text", "parameters": map[string]any{"type": "object", "properties": map[string]any{"string": map[string]any{"type": "string"}}, "required": []string{"string"}}}}, + "tool_choice": map[string]any{"type": "function", "name": "echo"}, + } + response := doResponsesRequest(t, fixture, payload) + defer response.Body.Close() + if response.StatusCode != http.StatusOK || response.Header.Get("Content-Type") != "text/event-stream" { + body, _ := io.ReadAll(response.Body) + t.Fatalf("stream response = %d %s: %s", response.StatusCode, response.Header.Get("Content-Type"), body) + } + events := readResponsesEvents(t, response.Body) + assertResponsesEventTypes(t, events, + "response.created", "response.in_progress", "response.output_item.added", + "response.reasoning_summary_part.added", "response.reasoning_summary_text.delta", + "response.output_item.done", "response.output_item.added", "response.output_text.delta", + "response.output_item.done", "response.output_item.added", "response.function_call_arguments.delta", + "response.function_call_arguments.done", "response.output_item.done", "response.completed", + ) + completed := findResponsesEvent(t, events, "response.completed") + completedResponse := completed["response"].(map[string]any) + output := completedResponse["output"].([]any) + if len(output) != 3 { + t.Fatalf("completed output = %#v", output) + } + reasoning := output[0].(map[string]any) + message := output[1].(map[string]any) + function := output[2].(map[string]any) + if reasoning["id"] != "rs_provider" || reasoning["encrypted_content"] != "ENC" || message["type"] != "message" || function["id"] == function["call_id"] || function["call_id"] != "call_external" { + t.Fatalf("completed items = %#v", output) + } + addedIDs := make(map[string]any) + for _, event := range events { + if event["type"] != "response.output_item.added" { + continue + } + item := event["item"].(map[string]any) + addedIDs[item["type"].(string)] = item["id"] + } + if addedIDs["reasoning"] != reasoning["id"] || addedIDs["message"] != message["id"] || addedIDs["function_call"] != function["id"] { + t.Fatalf("unstable stream item IDs: added=%#v completed=%#v", addedIDs, output) + } + usage := completedResponse["usage"].(map[string]any) + if usage["input_tokens"] != float64(14) || usage["output_tokens"] != float64(6) || usage["total_tokens"] != float64(20) { + t.Fatalf("completed usage = %#v", usage) + } + + requests := provider.recorded() + if len(requests) != 1 { + t.Fatalf("provider requests = %d", len(requests)) + } + request := requests[0] + if request.ReasoningEffort != "high" || request.MaxOutputTokens != 123 || !request.TemperatureSet || !request.TopPSet || request.ServiceTier != "priority" || !request.ParallelToolCalls || len(request.Tools) != 1 || request.Tools[0].Name != "echo" || request.ToolChoice.Mode != llm.ToolChoiceName { + t.Fatalf("translated request options = %+v", request) + } + properties, schemaOK := request.Tools[0].Schema["properties"].(map[string]interface{}) + if !schemaOK || properties["string"] == nil { + t.Fatalf("translated function schema = %#v", request.Tools[0].Schema) + } + if len(request.Messages) != 2 || len(request.Messages[1].Parts) != 3 || request.Messages[1].Parts[1].ImageData == nil || request.Messages[1].Parts[1].ImageData.Base64 != "aW1n" || request.Messages[1].Parts[2].FileData == nil || request.Messages[1].Parts[2].FileData.Filename != "notes.txt" { + t.Fatalf("translated multimodal input = %+v", request.Messages) + } +} + +func TestResponsesNonStreamingReasoningAndFunctionCall(t *testing.T) { + provider := &responsesFixtureProvider{events: [][]llm.Event{{ + {Type: llm.EventReasoningDelta, Text: "summary", ReasoningKind: llm.ReasoningKindSummary, ReasoningItemID: "rs_nonstream", ReasoningEncryptedContent: "ENC", ReasoningFinal: true}, + {Type: llm.EventToolCall, Tool: &llm.ToolCall{ID: "call_nonstream", Name: "echo", Arguments: json.RawMessage(`{"value":1}`)}}, + {Type: llm.EventUsage, Use: &llm.Usage{InputTokens: 3, OutputTokens: 2, ReasoningTokens: 1}}, + {Type: llm.EventDone}, + }}} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, provider, time.Second) + response := doResponsesRequest(t, fixture, map[string]any{ + "model": "remote/model-a", "input": "call a function", + "tools": []any{map[string]any{"type": "function", "name": "echo", "parameters": map[string]any{"type": "object"}}}, + }) + defer response.Body.Close() + var document map[string]any + if err := json.NewDecoder(response.Body).Decode(&document); err != nil { + t.Fatal(err) + } + output := document["output"].([]any) + if response.StatusCode != http.StatusOK || len(output) != 2 || output[0].(map[string]any)["id"] != "rs_nonstream" || output[1].(map[string]any)["call_id"] != "call_nonstream" { + t.Fatalf("nonstream reasoning/function response = %d %#v", response.StatusCode, document) + } +} + +func TestResponsesStatelessDiscourseFunctionOutputAndReasoningReplay(t *testing.T) { + provider := &responsesFixtureProvider{events: [][]llm.Event{{{Type: llm.EventTextDelta, Text: "continued"}, {Type: llm.EventDone}}}} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, provider, time.Second) + payload := map[string]any{ + "model": "remote/model-a", + "input": []any{ + map[string]any{"type": "reasoning", "id": "rs_old", "encrypted_content": "ENC_OLD", "summary": []any{map[string]any{"type": "summary_text", "text": "old summary"}}}, + map[string]any{"type": "message", "id": "msg_old", "role": "assistant", "content": []any{map[string]any{"type": "output_text", "text": "I will call echo"}}}, + // Discourse omits the optional output item id when provider metadata is absent. + map[string]any{"type": "function_call", "call_id": "call_old", "name": "echo", "arguments": `{"string":"old"}`}, + map[string]any{"type": "function_call_output", "call_id": "call_old", "output": "tool result"}, + }, + "tools": []any{map[string]any{"type": "function", "name": "echo", "parameters": map[string]any{"type": "object"}}}, + } + response := doResponsesRequest(t, fixture, payload) + defer response.Body.Close() + if response.StatusCode != http.StatusOK { + body, _ := io.ReadAll(response.Body) + t.Fatalf("continuation status = %d: %s", response.StatusCode, body) + } + requests := provider.recorded() + if len(requests) != 1 || len(requests[0].Messages) != 4 { + t.Fatalf("continuation request = %+v", requests) + } + if replay := requests[0].Messages[0].Parts[0].ProviderReplay; replay == nil || !bytes.Contains(replay.Raw, []byte(`"encrypted_content":"ENC_OLD"`)) { + t.Fatalf("reasoning replay = %+v", requests[0].Messages[0].Parts) + } + if replay := requests[0].Messages[1].Parts[0].ProviderReplay; replay == nil || !bytes.Contains(replay.Raw, []byte(`"id":"msg_old"`)) { + t.Fatalf("assistant message replay = %+v", requests[0].Messages[1].Parts) + } + if replay := requests[0].Messages[2].Parts[0].ProviderReplay; replay == nil || bytes.Contains(replay.Raw, []byte(`"id"`)) || !bytes.Contains(replay.Raw, []byte(`"call_id":"call_old"`)) { + t.Fatalf("function replay without provider item id = %+v", requests[0].Messages[2].Parts) + } + if call := requests[0].Messages[2].Parts[1].ToolCall; call == nil || call.ID != "call_old" { + t.Fatalf("function replay = %+v", requests[0].Messages[2].Parts) + } + if result := requests[0].Messages[3].Parts[0].ToolResult; result == nil || result.ID != "call_old" || result.Content != "tool result" { + t.Fatalf("function output = %+v", requests[0].Messages[3].Parts) + } +} + +func TestResponsesFunctionCallHistoryPreservesSuppliedItemID(t *testing.T) { + message, wireErr := decodeResponsesFunctionCall(json.RawMessage(`{"type":"function_call","id":"fc_provider","call_id":"call_provider","name":"echo","arguments":"{}"}`), 0) + if wireErr != nil { + t.Fatal(wireErr) + } + if len(message.Parts) != 2 || message.Parts[0].ProviderReplay == nil || !bytes.Contains(message.Parts[0].ProviderReplay.Raw, []byte(`"id":"fc_provider"`)) || message.Parts[1].ToolCall == nil || message.Parts[1].ToolCall.ID != "call_provider" { + t.Fatalf("function history with supplied item ID = %+v", message.Parts) + } +} + +func TestResponsesNamespacePolicyErrorsAndAllowedModels(t *testing.T) { + if provider, model, err := splitResponsesModel("openrouter/moonshotai/kimi-k2"); err != nil || provider != "openrouter" || model != "moonshotai/kimi-k2" { + t.Fatalf("first-slash namespace = %q/%q err=%v", provider, model, err) + } + provider := &responsesFixtureProvider{} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, provider, time.Second) + unauthorizedBody, _ := json.Marshal(map[string]any{"model": "remote/model-a", "input": "hello"}) + unauthorizedRequest, _ := http.NewRequest(http.MethodPost, fixture.server.URL+"/v1/responses", bytes.NewReader(unauthorizedBody)) + unauthorizedResponse, err := http.DefaultClient.Do(unauthorizedRequest) + if err != nil { + t.Fatal(err) + } + assertResponsesError(t, unauthorizedResponse, http.StatusUnauthorized, "gateway_client_unauthorized") + for _, model := range []string{"", "model-a", "/model-a", "remote/"} { + response := doResponsesRequest(t, fixture, map[string]any{"model": model, "input": "hello"}) + var body map[string]any + _ = json.NewDecoder(response.Body).Decode(&body) + response.Body.Close() + if response.StatusCode != http.StatusBadRequest || body["error"].(map[string]any)["code"] != "invalid_model_namespace" { + t.Fatalf("bad namespace %q = %d %#v", model, response.StatusCode, body) + } + } + unknown := doResponsesRequest(t, fixture, map[string]any{"model": "missing/model-a", "input": "hello"}) + assertResponsesError(t, unknown, http.StatusNotFound, "unknown_provider") + unknownModel := doResponsesRequest(t, fixture, map[string]any{"model": "remote/not-configured", "input": "hello"}) + assertResponsesError(t, unknownModel, http.StatusNotFound, "unknown_model") + + _, deniedToken, err := fixture.clients.Add("denied-token", Policy{AllowProviders: []string{"other"}}) + if err != nil { + t.Fatal(err) + } + requestBody, _ := json.Marshal(map[string]any{"model": "remote/model-a", "input": "hello"}) + req, _ := http.NewRequest(http.MethodPost, fixture.server.URL+"/v1/responses", bytes.NewReader(requestBody)) + req.Header.Set("Authorization", "Bearer "+deniedToken) + req.Header.Set("Content-Type", "application/json") + response, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + assertResponsesError(t, response, http.StatusForbidden, "policy_denied") + + modelsReq, _ := http.NewRequest(http.MethodGet, fixture.server.URL+"/v1/models", nil) + modelsReq.Header.Set("Authorization", "Bearer "+deniedToken) + modelsResponse, err := http.DefaultClient.Do(modelsReq) + if err != nil { + t.Fatal(err) + } + defer modelsResponse.Body.Close() + var models struct { + Object string `json:"object"` + Data []any `json:"data"` + } + if err := json.NewDecoder(modelsResponse.Body).Decode(&models); err != nil || models.Object != "list" || len(models.Data) != 0 { + t.Fatalf("denied models = %+v err=%v", models, err) + } +} + +func TestResponsesModelsListOnlyPolicyAllowedNamespaced(t *testing.T) { + provider := &responsesFixtureProvider{} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, provider, time.Second) + remote := fixture.central.Providers["remote"] + remote.Models = []string{"model-a", "model-b"} + fixture.central.Providers["remote"] = remote + fixture.central.Providers["other"] = config.ProviderConfig{Type: config.ProviderTypeOpenAI, Model: "other-model", Models: []string{"other-model"}, APIKey: "configured"} + _, token, err := fixture.clients.Add("model-filter", Policy{AllowProviders: []string{"remote"}, AllowModels: []string{"remote:model-a"}}) + if err != nil { + t.Fatal(err) + } + req, _ := http.NewRequest(http.MethodGet, fixture.server.URL+"/v1/models", nil) + req.Header.Set("Authorization", "Bearer "+token) + response, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + var list struct { + Object string `json:"object"` + Data []struct { + ID string `json:"id"` + Provider string `json:"provider"` + } `json:"data"` + } + if err := json.NewDecoder(response.Body).Decode(&list); err != nil { + t.Fatal(err) + } + if response.StatusCode != http.StatusOK || list.Object != "list" || len(list.Data) != 1 || list.Data[0].ID != "remote/model-a" || list.Data[0].Provider != "remote" { + t.Fatalf("policy-filtered models = %d %+v", response.StatusCode, list) + } +} + +func TestResponsesRejectsHostedToolsAndInlineCLIToolLoops(t *testing.T) { + provider := &responsesFixtureProvider{} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, provider, time.Second) + hosted := doResponsesRequest(t, fixture, map[string]any{"model": "remote/model-a", "input": "hello", "tools": []any{map[string]any{"type": "web_search"}}}) + assertResponsesError(t, hosted, http.StatusBadRequest, "unsupported_tool_type") + previousState := doResponsesRequest(t, fixture, map[string]any{"model": "remote/model-a", "input": "hello", "previous_response_id": "resp_forged"}) + assertResponsesError(t, previousState, http.StatusBadRequest, "unsupported_field") + unknownField := doResponsesRequest(t, fixture, map[string]any{"model": "remote/model-a", "input": "hello", "gateway_tools": true}) + assertResponsesError(t, unknownField, http.StatusBadRequest, "unsupported_field") + + inline := &inlineProvider{} + cliFixture := newGatewayFixture(t, config.ProviderTypeGrokBin, inline, time.Second) + incompatible := doResponsesRequest(t, cliFixture, map[string]any{ + "model": "remote/model-a", "input": "hello", + "tools": []any{map[string]any{"type": "function", "name": "echo", "parameters": map[string]any{"type": "object"}}}, + }) + assertResponsesError(t, incompatible, http.StatusBadRequest, "incompatible_tool_request") + inline.mu.Lock() + defer inline.mu.Unlock() + if inline.response.Result.Content != "" || inline.response.Err != nil { + t.Fatalf("inline provider tool loop unexpectedly ran: %+v", inline.response) + } +} + +func TestResponsesDisconnectCancelsProviderAndRecordsCanceledUsage(t *testing.T) { + provider := &responsesFixtureProvider{started: make(chan struct{}), canceled: make(chan struct{})} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, provider, time.Second) + ctx, cancel := context.WithCancel(context.Background()) + body, _ := json.Marshal(map[string]any{"model": "remote/model-a", "input": "wait", "stream": true}) + req, _ := http.NewRequestWithContext(ctx, http.MethodPost, fixture.server.URL+"/v1/responses", bytes.NewReader(body)) + req.Header.Set("Authorization", "Bearer "+fixture.token) + req.Header.Set("Content-Type", "application/json") + result := make(chan error, 1) + go func() { + resp, err := http.DefaultClient.Do(req) + if err == nil { + _, _ = bufio.NewReader(resp.Body).ReadString('\n') + _ = resp.Body.Close() + } + result <- err + }() + select { + case <-provider.started: + case <-time.After(time.Second): + t.Fatal("provider did not start") + } + cancel() + select { + case <-provider.canceled: + case <-time.After(time.Second): + t.Fatal("HTTP disconnect did not cancel provider") + } + select { + case <-result: + case <-time.After(time.Second): + t.Fatal("HTTP request did not return after cancellation") + } + deadline := time.Now().Add(time.Second) + for { + fixture.usage.mu.Lock() + if len(fixture.usage.records) > 0 { + record := fixture.usage.records[len(fixture.usage.records)-1] + fixture.usage.mu.Unlock() + if record.ErrorCode != "canceled" { + t.Fatalf("canceled usage = %+v", record) + } + break + } + fixture.usage.mu.Unlock() + if time.Now().After(deadline) { + t.Fatal("canceled usage was not recorded") + } + time.Sleep(time.Millisecond) + } +} + +func assertOfficialResponsesDocument(t *testing.T, response openairesponses.Response) { + t.Helper() + checks := []struct { + name string + valid bool + }{ + {"id", response.JSON.ID.Valid()}, + {"created_at", response.JSON.CreatedAt.Valid()}, + {"error", response.JSON.Error.Raw() != ""}, + {"incomplete_details", response.JSON.IncompleteDetails.Raw() != ""}, + {"instructions", response.JSON.Instructions.Raw() != ""}, + {"metadata", response.JSON.Metadata.Valid()}, + {"model", response.JSON.Model.Valid()}, + {"object", response.JSON.Object.Valid()}, + {"output", response.JSON.Output.Valid()}, + {"parallel_tool_calls", response.JSON.ParallelToolCalls.Valid()}, + {"temperature", response.JSON.Temperature.Valid()}, + {"tool_choice", response.JSON.ToolChoice.Valid()}, + {"tools", response.JSON.Tools.Valid()}, + {"top_p", response.JSON.TopP.Valid()}, + } + for _, check := range checks { + if !check.valid { + t.Fatalf("official Responses document omitted required field %q: %s", check.name, response.RawJSON()) + } + } + if response.ID == "" || response.Object != "response" || response.Model != "remote/model-a" || response.Output == nil || response.Metadata == nil || response.Tools == nil { + t.Fatalf("official typed Responses document fields = %+v", response) + } +} + +func assertOfficialResponsesDefaults(t *testing.T, response openairesponses.Response) { + t.Helper() + if response.JSON.Instructions.Raw() != "null" || len(response.Metadata) != 0 || response.Temperature != 1 || response.TopP != 1 || response.ToolChoice.AsToolChoiceMode() != "auto" || len(response.Tools) != 0 || !response.ParallelToolCalls { + t.Fatalf("official Responses defaults = %+v raw=%s", response, response.RawJSON()) + } +} + +func responsesEventTypesContain(types []string, want string) bool { + for _, eventType := range types { + if eventType == want { + return true + } + } + return false +} + +func doResponsesRequest(t *testing.T, fixture *gatewayFixture, payload any) *http.Response { + t.Helper() + body, err := json.Marshal(payload) + if err != nil { + t.Fatal(err) + } + req, err := http.NewRequest(http.MethodPost, fixture.server.URL+"/v1/responses", bytes.NewReader(body)) + if err != nil { + t.Fatal(err) + } + req.Header.Set("Authorization", "Bearer "+fixture.token) + req.Header.Set("Content-Type", "application/json") + response, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + return response +} + +func discourseShapedResponsesStream(response *http.Response) (string, error) { + defer response.Body.Close() + if response.StatusCode < 200 || response.StatusCode >= 300 { + var body struct { + Error struct { + Code string `json:"code"` + Message string `json:"message"` + } `json:"error"` + } + if err := json.NewDecoder(response.Body).Decode(&body); err != nil { + return "", fmt.Errorf("Responses HTTP %d: invalid error body: %w", response.StatusCode, err) + } + return "", fmt.Errorf("Responses HTTP %d: %s: %s", response.StatusCode, body.Error.Code, body.Error.Message) + } + var output strings.Builder + scanner := bufio.NewScanner(response.Body) + for scanner.Scan() { + line := scanner.Text() + if !strings.HasPrefix(line, "data: ") { + continue + } + var event struct { + Type string `json:"type"` + Delta string `json:"delta"` + } + if json.Unmarshal([]byte(strings.TrimPrefix(line, "data: ")), &event) == nil && event.Type == "response.output_text.delta" { + output.WriteString(event.Delta) + } + } + if err := scanner.Err(); err != nil { + return "", err + } + // Discourse's Net::HTTP edge trusts the HTTP status and consumes text deltas; + // it cannot reinterpret an HTTP 200 with no deltas as an upstream rejection. + return output.String(), nil +} + +func readResponsesEvents(t *testing.T, reader io.Reader) []map[string]any { + t.Helper() + var events []map[string]any + scanner := bufio.NewScanner(reader) + for scanner.Scan() { + line := scanner.Text() + if strings.HasPrefix(line, "event:") || line == "" { + continue + } + if !strings.HasPrefix(line, "data: ") { + t.Fatalf("non-Discourse SSE line %q", line) + } + var event map[string]any + if err := json.Unmarshal([]byte(strings.TrimPrefix(line, "data: ")), &event); err != nil { + t.Fatalf("decode SSE event: %v", err) + } + events = append(events, event) + } + if err := scanner.Err(); err != nil { + t.Fatal(err) + } + return events +} + +func assertResponsesEventTypes(t *testing.T, events []map[string]any, required ...string) { + t.Helper() + positions := make(map[string][]int) + for i, event := range events { + positions[event["type"].(string)] = append(positions[event["type"].(string)], i) + } + last := -1 + for _, eventType := range required { + found := -1 + for _, position := range positions[eventType] { + if position > last { + found = position + break + } + } + if found < 0 { + t.Fatalf("event %q not found after position %d; events=%#v", eventType, last, events) + } + last = found + } +} + +func findResponsesEvent(t *testing.T, events []map[string]any, eventType string) map[string]any { + t.Helper() + for _, event := range events { + if event["type"] == eventType { + return event + } + } + t.Fatalf("event %q not found", eventType) + return nil +} + +func assertResponsesError(t *testing.T, response *http.Response, status int, code string) { + t.Helper() + defer response.Body.Close() + var body struct { + Error struct { + Code string `json:"code"` + Message string `json:"message"` + } `json:"error"` + } + if err := json.NewDecoder(response.Body).Decode(&body); err != nil { + t.Fatal(err) + } + if response.StatusCode != status || body.Error.Code != code || body.Error.Message == "" { + t.Fatalf("Responses error = %d %+v, want %d/%s", response.StatusCode, body.Error, status, code) + } + for _, forbidden := range []string{"super-secret-provider-key", filepath.Clean(t.TempDir())} { + if forbidden != "" && strings.Contains(body.Error.Message, forbidden) { + t.Fatalf("unsafe error message = %q", body.Error.Message) + } + } +} diff --git a/internal/gateway/responses_wire.go b/internal/gateway/responses_wire.go new file mode 100644 index 000000000..f8d6cb267 --- /dev/null +++ b/internal/gateway/responses_wire.go @@ -0,0 +1,517 @@ +package gateway + +import ( + "bytes" + "encoding/base64" + "encoding/json" + "fmt" + "io" + "math" + "mime" + "net/http" + "strings" + + "github.com/samsaffron/term-llm/internal/llm" +) + +type responsesRequest struct { + Model string `json:"model"` + Input json.RawMessage `json:"input"` + Instructions *string `json:"instructions"` + Metadata map[string]string `json:"metadata"` + MaxOutputTokens *int `json:"max_output_tokens,omitempty"` + Stream bool `json:"stream,omitempty"` + Reasoning json.RawMessage `json:"reasoning,omitempty"` + Include []string `json:"include,omitempty"` + Temperature *float64 `json:"temperature,omitempty"` + TopP *float64 `json:"top_p,omitempty"` + ServiceTier *string `json:"service_tier,omitempty"` + ParallelToolCalls *bool `json:"parallel_tool_calls,omitempty"` + Tools []json.RawMessage `json:"tools,omitempty"` + ToolChoice json.RawMessage `json:"tool_choice,omitempty"` +} + +type responsesWireError struct { + Status int + Code string + Message string + Param string +} + +func (e *responsesWireError) Error() string { + if e == nil { + return "invalid Responses request" + } + return e.Message +} + +func (s *Server) decodeResponsesRequest(r *http.Request) (responsesRequest, *responsesWireError) { + data, err := io.ReadAll(io.LimitReader(r.Body, s.cfg.MaxBodyBytes+1)) + if err != nil { + return responsesRequest{}, invalidResponsesRequest("invalid_json", "could not read request body", "") + } + if int64(len(data)) > s.cfg.MaxBodyBytes { + return responsesRequest{}, &responsesWireError{Status: http.StatusRequestEntityTooLarge, Code: "request_too_large", Message: "request body exceeds the configured gateway limit"} + } + var fields map[string]json.RawMessage + dec := json.NewDecoder(bytes.NewReader(data)) + if err := dec.Decode(&fields); err != nil { + return responsesRequest{}, invalidResponsesRequest("invalid_json", "request body must be valid JSON", "") + } + if err := dec.Decode(&struct{}{}); err != io.EOF { + return responsesRequest{}, invalidResponsesRequest("invalid_json", "request body must contain one JSON object", "") + } + if fields == nil { + return responsesRequest{}, invalidResponsesRequest("invalid_json", "request body must be a JSON object", "") + } + + known := map[string]bool{ + "model": true, "input": true, "instructions": true, "max_output_tokens": true, + "stream": true, "reasoning": true, "include": true, "temperature": true, + "top_p": true, "service_tier": true, "parallel_tool_calls": true, "tools": true, + "tool_choice": true, "metadata": true, "user": true, "store": true, + "stream_options": true, "prompt_cache_key": true, "prompt_cache_retention": true, + "safety_identifier": true, "background": true, "previous_response_id": true, + "conversation": true, "text": true, "truncation": true, "max_tool_calls": true, + "top_logprobs": true, + } + for name := range fields { + if !known[name] { + return responsesRequest{}, invalidResponsesRequest("unsupported_field", fmt.Sprintf("field %q is not supported by this Responses edge", name), name) + } + } + if raw := fields["background"]; len(raw) > 0 && !bytes.Equal(bytes.TrimSpace(raw), []byte("false")) && !bytes.Equal(bytes.TrimSpace(raw), []byte("null")) { + return responsesRequest{}, invalidResponsesRequest("unsupported_field", "background responses are not supported", "background") + } + for _, name := range []string{"previous_response_id", "conversation", "text", "max_tool_calls"} { + if raw := bytes.TrimSpace(fields[name]); len(raw) > 0 && !bytes.Equal(raw, []byte("null")) && !bytes.Equal(raw, []byte(`""`)) && !bytes.Equal(raw, []byte("{}")) { + return responsesRequest{}, invalidResponsesRequest("unsupported_field", fmt.Sprintf("field %q is not supported; send complete stateless input instead", name), name) + } + } + if raw := fields["truncation"]; len(raw) > 0 { + var value string + if json.Unmarshal(raw, &value) != nil || (value != "" && value != "disabled") { + return responsesRequest{}, invalidResponsesRequest("unsupported_field", "only truncation=disabled is supported", "truncation") + } + } + if raw := fields["top_logprobs"]; len(raw) > 0 { + var value int + if json.Unmarshal(raw, &value) != nil || value != 0 { + return responsesRequest{}, invalidResponsesRequest("unsupported_field", "top_logprobs is not supported", "top_logprobs") + } + } + + var request responsesRequest + if err := json.Unmarshal(data, &request); err != nil { + return responsesRequest{}, invalidResponsesRequest("invalid_json", "request fields have invalid types", "") + } + request.Model = strings.TrimSpace(request.Model) + if request.Model == "" { + return responsesRequest{}, invalidResponsesRequest("invalid_model_namespace", "model must use provider/model namespace", "model") + } + if len(bytes.TrimSpace(request.Input)) == 0 || bytes.Equal(bytes.TrimSpace(request.Input), []byte("null")) { + return responsesRequest{}, invalidResponsesRequest("invalid_input", "input is required", "input") + } + if request.MaxOutputTokens != nil && *request.MaxOutputTokens <= 0 { + return responsesRequest{}, invalidResponsesRequest("invalid_request", "max_output_tokens must be positive", "max_output_tokens") + } + if err := validateResponsesFloat(request.Temperature, "temperature", 0, 2); err != nil { + return responsesRequest{}, err + } + if err := validateResponsesFloat(request.TopP, "top_p", 0, 1); err != nil { + return responsesRequest{}, err + } + if request.ServiceTier != nil { + switch strings.TrimSpace(*request.ServiceTier) { + case "", "auto", "default", "flex", "scale", "priority": + default: + return responsesRequest{}, invalidResponsesRequest("invalid_service_tier", "service_tier must be auto, default, flex, scale, or priority", "service_tier") + } + } + for _, include := range request.Include { + if include != "reasoning.encrypted_content" { + return responsesRequest{}, invalidResponsesRequest("unsupported_include", fmt.Sprintf("include value %q is not supported", include), "include") + } + } + return request, nil +} + +func validateResponsesFloat(value *float64, name string, minValue, maxValue float64) *responsesWireError { + if value == nil { + return nil + } + if math.IsNaN(*value) || math.IsInf(*value, 0) || *value < minValue || *value > maxValue { + return invalidResponsesRequest("invalid_request", fmt.Sprintf("%s must be between %g and %g", name, minValue, maxValue), name) + } + return nil +} + +func invalidResponsesRequest(code, message, param string) *responsesWireError { + return &responsesWireError{Status: http.StatusBadRequest, Code: code, Message: message, Param: param} +} + +func splitResponsesModel(value string) (string, string, *responsesWireError) { + provider, model, found := strings.Cut(strings.TrimSpace(value), "/") + provider = strings.TrimSpace(provider) + model = strings.TrimSpace(model) + if !found || provider == "" || model == "" { + return "", "", invalidResponsesRequest("invalid_model_namespace", "model must use provider/model namespace with non-empty provider and model", "model") + } + return provider, model, nil +} + +func translateResponsesRequest(request responsesRequest) (llm.Request, *responsesWireError) { + messages, inputErr := decodeResponsesInput(request.Input) + if inputErr != nil { + return llm.Request{}, inputErr + } + if request.Instructions != nil && strings.TrimSpace(*request.Instructions) != "" { + messages = append([]llm.Message{{Role: llm.RoleDeveloper, Parts: []llm.Part{{Type: llm.PartText, Text: *request.Instructions}}}}, messages...) + } + tools, toolErr := decodeResponsesTools(request.Tools) + if toolErr != nil { + return llm.Request{}, toolErr + } + choice, choiceErr := decodeResponsesToolChoice(request.ToolChoice, tools) + if choiceErr != nil { + return llm.Request{}, choiceErr + } + effort, reasoningErr := decodeResponsesReasoning(request.Reasoning) + if reasoningErr != nil { + return llm.Request{}, reasoningErr + } + out := llm.Request{ + Messages: messages, Tools: tools, ToolChoice: choice, ReasoningEffort: effort, + DisableExternalWebFetch: true, ParallelToolCalls: true, + } + if request.ParallelToolCalls != nil { + out.ParallelToolCalls = *request.ParallelToolCalls + } + if request.MaxOutputTokens != nil { + out.MaxOutputTokens = *request.MaxOutputTokens + } + if request.Temperature != nil { + out.Temperature = float32(*request.Temperature) + out.TemperatureSet = true + } + if request.TopP != nil { + out.TopP = float32(*request.TopP) + out.TopPSet = true + } + if request.ServiceTier != nil { + out.ServiceTier = strings.TrimSpace(*request.ServiceTier) + out.ServiceTierSet = true + } + return out, nil +} + +func decodeResponsesReasoning(raw json.RawMessage) (string, *responsesWireError) { + if len(bytes.TrimSpace(raw)) == 0 || bytes.Equal(bytes.TrimSpace(raw), []byte("null")) { + return "", nil + } + var fields map[string]json.RawMessage + if err := json.Unmarshal(raw, &fields); err != nil { + return "", invalidResponsesRequest("invalid_reasoning", "reasoning must be an object", "reasoning") + } + for name := range fields { + if name != "summary" && name != "effort" { + return "", invalidResponsesRequest("unsupported_field", fmt.Sprintf("reasoning.%s is not supported", name), "reasoning."+name) + } + } + var value struct { + Summary string `json:"summary"` + Effort string `json:"effort"` + } + if err := json.Unmarshal(raw, &value); err != nil { + return "", invalidResponsesRequest("invalid_reasoning", "reasoning fields have invalid types", "reasoning") + } + if value.Summary != "" && value.Summary != "auto" { + return "", invalidResponsesRequest("unsupported_reasoning_summary", "only reasoning.summary=auto is supported", "reasoning.summary") + } + return strings.TrimSpace(value.Effort), nil +} + +func decodeResponsesTools(rawTools []json.RawMessage) ([]llm.ToolSpec, *responsesWireError) { + tools := make([]llm.ToolSpec, 0, len(rawTools)) + seen := make(map[string]bool) + for i, raw := range rawTools { + var fields map[string]json.RawMessage + if err := json.Unmarshal(raw, &fields); err != nil { + return nil, invalidResponsesRequest("invalid_tool", fmt.Sprintf("tools[%d] must be an object", i), "tools") + } + var tool struct { + Type string `json:"type"` + Name string `json:"name"` + Description string `json:"description"` + Parameters map[string]interface{} `json:"parameters"` + Strict bool `json:"strict"` + } + if err := json.Unmarshal(raw, &tool); err != nil || tool.Type != "function" { + return nil, invalidResponsesRequest("unsupported_tool_type", "only flat Responses function tools are supported; gateway-side hosted/search tools are disabled", fmt.Sprintf("tools[%d].type", i)) + } + for name := range fields { + if name != "type" && name != "name" && name != "description" && name != "parameters" && name != "strict" { + return nil, invalidResponsesRequest("unsupported_tool_field", fmt.Sprintf("tools[%d].%s is not supported", i, name), "tools") + } + } + tool.Name = strings.TrimSpace(tool.Name) + if tool.Name == "" || tool.Parameters == nil || seen[tool.Name] { + return nil, invalidResponsesRequest("invalid_tool", "function tools require unique non-empty names and parameters", "tools") + } + seen[tool.Name] = true + tools = append(tools, llm.ToolSpec{Name: tool.Name, Description: tool.Description, Schema: tool.Parameters, Strict: tool.Strict}) + } + return tools, nil +} + +func decodeResponsesToolChoice(raw json.RawMessage, tools []llm.ToolSpec) (llm.ToolChoice, *responsesWireError) { + raw = bytes.TrimSpace(raw) + if len(raw) == 0 || bytes.Equal(raw, []byte("null")) { + return llm.ToolChoice{}, nil + } + var mode string + if json.Unmarshal(raw, &mode) == nil { + switch mode { + case "none": + return llm.ToolChoice{Mode: llm.ToolChoiceNone}, nil + case "auto": + return llm.ToolChoice{Mode: llm.ToolChoiceAuto}, nil + case "required": + return llm.ToolChoice{Mode: llm.ToolChoiceRequired}, nil + default: + return llm.ToolChoice{}, invalidResponsesRequest("invalid_tool_choice", "tool_choice must be none, auto, required, or a named function", "tool_choice") + } + } + var named struct { + Type string `json:"type"` + Name string `json:"name"` + } + if err := json.Unmarshal(raw, &named); err != nil || named.Type != "function" || strings.TrimSpace(named.Name) == "" { + return llm.ToolChoice{}, invalidResponsesRequest("invalid_tool_choice", "named tool_choice must be {type:function,name:...}", "tool_choice") + } + for _, tool := range tools { + if tool.Name == named.Name { + return llm.ToolChoice{Mode: llm.ToolChoiceName, Name: named.Name}, nil + } + } + return llm.ToolChoice{}, invalidResponsesRequest("invalid_tool_choice", "named tool_choice does not match a supplied function tool", "tool_choice") +} + +func decodeResponsesInput(raw json.RawMessage) ([]llm.Message, *responsesWireError) { + var text string + if json.Unmarshal(raw, &text) == nil { + return []llm.Message{{Role: llm.RoleUser, Parts: []llm.Part{{Type: llm.PartText, Text: text}}}}, nil + } + var items []json.RawMessage + if err := json.Unmarshal(raw, &items); err != nil { + return nil, invalidResponsesRequest("invalid_input", "input must be a string or an array of Responses input items", "input") + } + messages := make([]llm.Message, 0, len(items)) + for i, rawItem := range items { + message, err := decodeResponsesInputItem(rawItem, i) + if err != nil { + return nil, err + } + messages = append(messages, message) + } + return messages, nil +} + +func decodeResponsesInputItem(raw json.RawMessage, index int) (llm.Message, *responsesWireError) { + var header struct { + Type string `json:"type"` + Role string `json:"role"` + } + if err := json.Unmarshal(raw, &header); err != nil { + return llm.Message{}, invalidResponsesRequest("invalid_input", fmt.Sprintf("input[%d] must be an object", index), "input") + } + if header.Type == "" && header.Role != "" { + header.Type = "message" + } + switch header.Type { + case "message": + return decodeResponsesMessage(raw, index, header.Role) + case "reasoning": + return decodeResponsesReasoningItem(raw, index) + case "function_call": + return decodeResponsesFunctionCall(raw, index) + case "function_call_output": + return decodeResponsesFunctionOutput(raw, index) + default: + return llm.Message{}, invalidResponsesRequest("unsupported_input_type", fmt.Sprintf("input[%d] type %q is not supported", index, header.Type), "input") + } +} + +func decodeResponsesMessage(raw json.RawMessage, index int, roleName string) (llm.Message, *responsesWireError) { + role, ok := mapResponsesRole(roleName) + if !ok { + return llm.Message{}, invalidResponsesRequest("invalid_role", fmt.Sprintf("input[%d] role must be developer, user, or assistant", index), "input") + } + var item struct { + ID string `json:"id,omitempty"` + Content json.RawMessage `json:"content"` + } + if err := json.Unmarshal(raw, &item); err != nil || len(item.Content) == 0 { + return llm.Message{}, invalidResponsesRequest("invalid_content", fmt.Sprintf("input[%d] message content is required", index), "input") + } + parts, err := decodeResponsesContent(item.Content, role, index) + if err != nil { + return llm.Message{}, err + } + message := llm.Message{Role: role, Parts: parts} + if role == llm.RoleAssistant && strings.TrimSpace(item.ID) != "" { + sanitized, marshalErr := json.Marshal(map[string]any{"type": "message", "id": strings.TrimSpace(item.ID), "role": "assistant", "content": responsesReplayContent(parts)}) + if marshalErr == nil { + message.Parts = append([]llm.Part{{Type: llm.PartProviderReplay, ProviderReplay: &llm.ProviderReplayItem{Raw: sanitized}}}, message.Parts...) + } + } + return message, nil +} + +func mapResponsesRole(value string) (llm.Role, bool) { + switch value { + case "developer": + return llm.RoleDeveloper, true + case "user": + return llm.RoleUser, true + case "assistant": + return llm.RoleAssistant, true + default: + return "", false + } +} + +func decodeResponsesContent(raw json.RawMessage, role llm.Role, itemIndex int) ([]llm.Part, *responsesWireError) { + var text string + if json.Unmarshal(raw, &text) == nil { + return []llm.Part{{Type: llm.PartText, Text: text}}, nil + } + var content []struct { + Type string `json:"type"` + Text string `json:"text,omitempty"` + ImageURL string `json:"image_url,omitempty"` + Detail string `json:"detail,omitempty"` + Filename string `json:"filename,omitempty"` + FileData string `json:"file_data,omitempty"` + } + if err := json.Unmarshal(raw, &content); err != nil { + return nil, invalidResponsesRequest("invalid_content", fmt.Sprintf("input[%d].content must be a string or array", itemIndex), "input") + } + parts := make([]llm.Part, 0, len(content)) + for partIndex, part := range content { + switch part.Type { + case "input_text", "output_text", "text": + parts = append(parts, llm.Part{Type: llm.PartText, Text: part.Text}) + case "input_image": + if role == llm.RoleAssistant { + return nil, invalidResponsesRequest("invalid_content", "assistant input_image content is not supported", "input") + } + mediaType, encoded, size, err := decodeResponsesDataURL(part.ImageURL) + if err != nil || !strings.HasPrefix(mediaType, "image/") { + return nil, invalidResponsesRequest("invalid_image", fmt.Sprintf("input[%d].content[%d] must contain a base64 image data URL", itemIndex, partIndex), "input") + } + _ = size + parts = append(parts, llm.Part{Type: llm.PartImage, ImageData: &llm.ToolImageData{MediaType: mediaType, Base64: encoded, Detail: part.Detail}}) + case "input_file": + if role == llm.RoleAssistant { + return nil, invalidResponsesRequest("invalid_content", "assistant input_file content is not supported", "input") + } + mediaType, encoded, size, err := decodeResponsesDataURL(part.FileData) + if err != nil { + return nil, invalidResponsesRequest("invalid_file", fmt.Sprintf("input[%d].content[%d] must contain a base64 file data URL", itemIndex, partIndex), "input") + } + parts = append(parts, llm.Part{Type: llm.PartFile, FileData: &llm.ToolFileData{MediaType: mediaType, Base64: encoded, Filename: part.Filename, SizeBytes: int64(size)}}) + default: + return nil, invalidResponsesRequest("unsupported_content_type", fmt.Sprintf("input[%d].content[%d] type %q is not supported", itemIndex, partIndex, part.Type), "input") + } + } + return parts, nil +} + +func decodeResponsesDataURL(value string) (string, string, int, error) { + if !strings.HasPrefix(value, "data:") { + return "", "", 0, fmt.Errorf("not a data URL") + } + header, payload, found := strings.Cut(strings.TrimPrefix(value, "data:"), ",") + if !found || !strings.HasSuffix(strings.ToLower(header), ";base64") { + return "", "", 0, fmt.Errorf("not base64") + } + mediaType := strings.TrimSpace(strings.TrimSuffix(header, ";base64")) + if parsed, _, err := mime.ParseMediaType(mediaType); err == nil { + mediaType = parsed + } else { + return "", "", 0, err + } + decoded, err := base64.StdEncoding.DecodeString(payload) + if err != nil { + return "", "", 0, err + } + return mediaType, payload, len(decoded), nil +} + +func decodeResponsesReasoningItem(raw json.RawMessage, index int) (llm.Message, *responsesWireError) { + var item struct { + ID string `json:"id"` + EncryptedContent string `json:"encrypted_content"` + Summary []struct { + Type string `json:"type"` + Text string `json:"text"` + } `json:"summary"` + } + if err := json.Unmarshal(raw, &item); err != nil || strings.TrimSpace(item.ID) == "" { + return llm.Message{}, invalidResponsesRequest("invalid_reasoning", fmt.Sprintf("input[%d] reasoning item requires a non-empty id", index), "input") + } + summaries := make([]string, 0, len(item.Summary)) + for _, summary := range item.Summary { + if summary.Type != "summary_text" { + return llm.Message{}, invalidResponsesRequest("invalid_reasoning", "reasoning summary parts must use type summary_text", "input") + } + summaries = append(summaries, summary.Text) + } + part := llm.Part{Type: llm.PartText, ReasoningItemID: strings.TrimSpace(item.ID), ReasoningEncryptedContent: item.EncryptedContent, ReasoningSummaryParts: summaries, ReasoningContent: strings.Join(summaries, "\n\n"), ReasoningKind: llm.ReasoningKindSummary} + if len(summaries) == 0 && item.EncryptedContent != "" { + part.ReasoningKind = llm.ReasoningKindEncrypted + } + sanitized, _ := json.Marshal(map[string]any{"type": "reasoning", "id": strings.TrimSpace(item.ID), "encrypted_content": item.EncryptedContent, "summary": item.Summary}) + return llm.Message{Role: llm.RoleAssistant, Parts: []llm.Part{{Type: llm.PartProviderReplay, ProviderReplay: &llm.ProviderReplayItem{Raw: sanitized}}, part}}, nil +} + +func decodeResponsesFunctionCall(raw json.RawMessage, index int) (llm.Message, *responsesWireError) { + var item struct { + ID string `json:"id"` + CallID string `json:"call_id"` + Name string `json:"name"` + Arguments string `json:"arguments"` + } + if err := json.Unmarshal(raw, &item); err != nil || strings.TrimSpace(item.CallID) == "" || strings.TrimSpace(item.Name) == "" || !json.Valid([]byte(item.Arguments)) { + return llm.Message{}, invalidResponsesRequest("invalid_function_call", fmt.Sprintf("input[%d] function_call requires call_id, name, and JSON arguments", index), "input") + } + replay := map[string]any{"type": "function_call", "call_id": item.CallID, "name": item.Name, "arguments": item.Arguments} + if itemID := strings.TrimSpace(item.ID); itemID != "" { + replay["id"] = itemID + } + sanitized, _ := json.Marshal(replay) + call := &llm.ToolCall{ID: item.CallID, Name: item.Name, Arguments: json.RawMessage(item.Arguments)} + return llm.Message{Role: llm.RoleAssistant, Parts: []llm.Part{{Type: llm.PartProviderReplay, ProviderReplay: &llm.ProviderReplayItem{Raw: sanitized}}, {Type: llm.PartToolCall, ToolCall: call}}}, nil +} + +func decodeResponsesFunctionOutput(raw json.RawMessage, index int) (llm.Message, *responsesWireError) { + var item struct { + CallID string `json:"call_id"` + Output string `json:"output"` + } + if err := json.Unmarshal(raw, &item); err != nil || strings.TrimSpace(item.CallID) == "" { + return llm.Message{}, invalidResponsesRequest("invalid_function_call_output", fmt.Sprintf("input[%d] function_call_output requires call_id and string output", index), "input") + } + return llm.Message{Role: llm.RoleTool, Parts: []llm.Part{{Type: llm.PartToolResult, ToolResult: &llm.ToolResult{ID: item.CallID, Content: item.Output}}}}, nil +} + +func responsesReplayContent(parts []llm.Part) []map[string]any { + content := make([]map[string]any, 0, len(parts)) + for _, part := range parts { + if part.Type == llm.PartText { + content = append(content, map[string]any{"type": "output_text", "text": part.Text}) + } + } + return content +} diff --git a/internal/gateway/seal.go b/internal/gateway/seal.go new file mode 100644 index 000000000..476ad8791 --- /dev/null +++ b/internal/gateway/seal.go @@ -0,0 +1,85 @@ +package gateway + +import ( + "crypto/aes" + "crypto/cipher" + "crypto/rand" + "encoding/base64" + "encoding/json" + "fmt" + "io" + "os" + "path/filepath" + + "github.com/samsaffron/term-llm/internal/gateway/protocol" +) + +type sealedProviderState struct { + Version int `json:"version"` + ClientID string `json:"client_id"` + Provider string `json:"provider"` + State []byte `json:"state"` +} + +type StateSealer struct{ aead cipher.AEAD } + +func OpenStateSealer(path string) (*StateSealer, error) { + key, err := os.ReadFile(path) + if err != nil { + if !os.IsNotExist(err) { + return nil, fmt.Errorf("read gateway state key: %w", err) + } + key = make([]byte, 32) + if _, err := io.ReadFull(rand.Reader, key); err != nil { + return nil, fmt.Errorf("generate gateway state key: %w", err) + } + if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil { + return nil, err + } + if err := os.WriteFile(path, key, 0o600); err != nil { + return nil, fmt.Errorf("write gateway state key: %w", err) + } + } + if len(key) != 32 { + return nil, fmt.Errorf("gateway state key must be 32 bytes") + } + block, err := aes.NewCipher(key) + if err != nil { + return nil, err + } + aead, err := cipher.NewGCM(block) + if err != nil { + return nil, err + } + return &StateSealer{aead: aead}, nil +} + +func (s *StateSealer) Seal(clientID, provider string, state []byte) (string, error) { + plain, err := json.Marshal(sealedProviderState{Version: protocol.Version, ClientID: clientID, Provider: provider, State: state}) + if err != nil { + return "", err + } + nonce := make([]byte, s.aead.NonceSize()) + if _, err := io.ReadFull(rand.Reader, nonce); err != nil { + return "", err + } + ciphertext := s.aead.Seal(nil, nonce, plain, []byte("term-llm-gateway-state-v1")) + return base64.RawURLEncoding.EncodeToString(append(nonce, ciphertext...)), nil +} + +func (s *StateSealer) Open(blob, clientID, provider string) ([]byte, error) { + raw, err := base64.RawURLEncoding.DecodeString(blob) + if err != nil || len(raw) < s.aead.NonceSize() { + return nil, fmt.Errorf("invalid sealed provider state") + } + nonce, ciphertext := raw[:s.aead.NonceSize()], raw[s.aead.NonceSize():] + plain, err := s.aead.Open(nil, nonce, ciphertext, []byte("term-llm-gateway-state-v1")) + if err != nil { + return nil, fmt.Errorf("invalid sealed provider state") + } + var state sealedProviderState + if err := json.Unmarshal(plain, &state); err != nil || state.Version != protocol.Version || state.ClientID != clientID || state.Provider != provider { + return nil, fmt.Errorf("sealed provider state does not belong to this client/provider") + } + return append([]byte(nil), state.State...), nil +} diff --git a/internal/gateway/security_test.go b/internal/gateway/security_test.go new file mode 100644 index 000000000..e75d9cfcf --- /dev/null +++ b/internal/gateway/security_test.go @@ -0,0 +1,414 @@ +package gateway + +import ( + "bytes" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "sync" + "testing" + "time" + + "github.com/samsaffron/term-llm/internal/config" + "github.com/samsaffron/term-llm/internal/gateway/protocol" + "github.com/samsaffron/term-llm/internal/llm" +) + +func TestClientStoreHashesAuthenticatesPoliciesAndRevokes(t *testing.T) { + path := filepath.Join(t.TempDir(), "clients.json") + store, err := OpenClientStore(path) + if err != nil { + t.Fatal(err) + } + policy := Policy{AllowProviders: []string{"openai"}, DenyModels: []string{"danger*"}, AllowSearch: true, AllowFetch: true, MaxConcurrentInference: 1, SearchRatePerMinute: 7, MaxConcurrentSearch: 1, FetchRatePerMinute: 9, MaxConcurrentFetch: 1} + client, token, err := store.Add("satellite-a", policy) + if err != nil { + t.Fatal(err) + } + data, _ := os.ReadFile(path) + if string(data) == "" || strings.Contains(string(data), token) { + t.Fatal("plaintext client token was persisted") + } + got, ok := store.Authenticate(token) + if !ok || got.ID != client.ID { + t.Fatalf("authentication failed: %+v %t", got, ok) + } + if !client.Policy.Allows("openai", "gpt", false) || client.Policy.Allows("anthropic", "claude", false) || client.Policy.Allows("openai", "danger-model", false) || client.Policy.Allows("openai", "gpt", true) { + t.Fatal("client provider/model/CLI policy was not enforced") + } + reopened, err := OpenClientStore(path) + if err != nil { + t.Fatal(err) + } + persisted, ok := reopened.Authenticate(token) + if !ok || persisted.Policy.MaxConcurrentInference != 1 || persisted.Policy.SearchRatePerMinute != 7 || persisted.Policy.MaxConcurrentSearch != 1 || persisted.Policy.FetchRatePerMinute != 9 || persisted.Policy.MaxConcurrentFetch != 1 { + t.Fatalf("persisted client limits = %+v, authenticated=%t", persisted.Policy, ok) + } + if err := store.Revoke(client.ID); err != nil { + t.Fatal(err) + } + if _, ok := store.Authenticate(token); ok { + t.Fatal("revoked token authenticated") + } +} + +func TestClientStoreEnforcesUniqueActiveNamesAndRotation(t *testing.T) { + path := filepath.Join(t.TempDir(), "clients.json") + store, err := OpenClientStore(path) + if err != nil { + t.Fatal(err) + } + original, originalToken, err := store.Add("satellite-a", Policy{}) + if err != nil { + t.Fatal(err) + } + if _, _, err := store.Add("satellite-a", Policy{}); err == nil || !strings.Contains(err.Error(), "revoke it before rotating") { + t.Fatalf("duplicate active name error = %v", err) + } + if authenticated, ok := store.Authenticate(originalToken); !ok || authenticated.ID != original.ID { + t.Fatalf("duplicate rejection disturbed original credential: %+v authenticated=%t", authenticated, ok) + } + if err := store.Revoke("satellite-a"); err != nil { + t.Fatal(err) + } + if _, ok := store.Authenticate(originalToken); ok { + t.Fatal("name-based revocation did not immediately disable original token") + } + rotated, rotatedToken, err := store.Add("satellite-a", Policy{}) + if err != nil { + t.Fatalf("add after revoke rotation: %v", err) + } + if rotated.ID == original.ID { + t.Fatal("rotation reused client identity") + } + if authenticated, ok := store.Authenticate(rotatedToken); !ok || authenticated.ID != rotated.ID { + t.Fatalf("rotated credential did not authenticate immediately: %+v authenticated=%t", authenticated, ok) + } +} + +func TestClientStoreRevokeReloadsAuthoritativeMultiStoreWithoutOverwritingAdds(t *testing.T) { + path := filepath.Join(t.TempDir(), "clients.json") + adminA, err := OpenClientStore(path) + if err != nil { + t.Fatal(err) + } + victim, victimToken, err := adminA.Add("victim", Policy{}) + if err != nil { + t.Fatal(err) + } + staleAdmin, err := OpenClientStore(path) + if err != nil { + t.Fatal(err) + } + _, survivorToken, err := adminA.Add("survivor", Policy{}) + if err != nil { + t.Fatal(err) + } + if err := staleAdmin.Revoke(victim.ID); err != nil { + t.Fatal(err) + } + authoritative, err := OpenClientStore(path) + if err != nil { + t.Fatal(err) + } + if _, ok := authoritative.Authenticate(victimToken); ok { + t.Fatal("revoked token remained active") + } + if survivor, ok := authoritative.Authenticate(survivorToken); !ok || survivor.Name != "survivor" { + t.Fatalf("stale-store revoke overwrote concurrent addition: %+v authenticated=%t", survivor, ok) + } +} + +func TestClientStoreConcurrentMultiStoreAddsAreSerialized(t *testing.T) { + path := filepath.Join(t.TempDir(), "clients.json") + stores := make([]*ClientStore, 2) + for i := range stores { + var err error + stores[i], err = OpenClientStore(path) + if err != nil { + t.Fatal(err) + } + } + start := make(chan struct{}) + tokens := make([]string, 2) + errs := make([]error, 2) + var wg sync.WaitGroup + for i := range stores { + wg.Add(1) + go func(i int) { + defer wg.Done() + <-start + _, tokens[i], errs[i] = stores[i].Add(fmt.Sprintf("concurrent-%d", i), Policy{}) + }(i) + } + close(start) + wg.Wait() + for _, err := range errs { + if err != nil { + t.Fatal(err) + } + } + authoritative, err := OpenClientStore(path) + if err != nil { + t.Fatal(err) + } + for i, token := range tokens { + client, ok := authoritative.Authenticate(token) + if !ok || client.Name != fmt.Sprintf("concurrent-%d", i) { + t.Fatalf("concurrent client %d = %+v authenticated=%t", i, client, ok) + } + } +} + +func TestRunningGatewayObservesAddAndRevokeFromSecondStoreImmediately(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "clients.json") + runtimeStore, err := OpenClientStore(path) + if err != nil { + t.Fatal(err) + } + adminStore, err := OpenClientStore(path) + if err != nil { + t.Fatal(err) + } + sealer, err := OpenStateSealer(filepath.Join(dir, "state.key")) + if err != nil { + t.Fatal(err) + } + server, err := NewServer(ServerConfig{Config: &config.Config{Providers: map[string]config.ProviderConfig{}}, Clients: runtimeStore, Sealer: sealer}) + if err != nil { + t.Fatal(err) + } + ts := httptest.NewServer(server.Handler()) + defer ts.Close() + + client, token, err := adminStore.Add("live-admin-change", Policy{}) + if err != nil { + t.Fatal(err) + } + request := func() int { + req, reqErr := http.NewRequest(http.MethodGet, ts.URL+"/g1/catalog", nil) + if reqErr != nil { + t.Fatal(reqErr) + } + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set(protocol.VersionHeader, "1") + resp, reqErr := http.DefaultClient.Do(req) + if reqErr != nil { + t.Fatal(reqErr) + } + resp.Body.Close() + return resp.StatusCode + } + if status := request(); status != http.StatusOK { + t.Fatalf("new client status = %d, want 200", status) + } + started := time.Now() + if err := adminStore.Revoke(client.ID); err != nil { + t.Fatal(err) + } + if status := request(); status != http.StatusUnauthorized { + t.Fatalf("revoked client status = %d, want 401", status) + } + if elapsed := time.Since(started); elapsed > 250*time.Millisecond { + t.Fatalf("revocation observation took %s; authentication has no polling interval", elapsed) + } +} + +func TestGatewayRunTempRootScavengesOnlyOwnedPrefixDirectories(t *testing.T) { + root := filepath.Join(t.TempDir(), "gateway-runs") + if err := os.MkdirAll(filepath.Join(root, "run-stale"), 0o700); err != nil { + t.Fatal(err) + } + if err := os.MkdirAll(filepath.Join(root, "operator-content"), 0o700); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(root, "run-not-a-directory"), []byte("keep"), 0o600); err != nil { + t.Fatal(err) + } + store, _ := OpenClientStore(filepath.Join(t.TempDir(), "clients.json")) + sealer, _ := OpenStateSealer(filepath.Join(t.TempDir(), "state.key")) + if _, err := NewServer(ServerConfig{Config: &config.Config{Providers: map[string]config.ProviderConfig{}}, Clients: store, Sealer: sealer, RunTempRoot: root}); err != nil { + t.Fatal(err) + } + if _, err := os.Stat(filepath.Join(root, "run-stale")); !os.IsNotExist(err) { + t.Fatalf("stale gateway run directory remains: %v", err) + } + for _, name := range []string{"operator-content", "run-not-a-directory"} { + if _, err := os.Stat(filepath.Join(root, name)); err != nil { + t.Fatalf("safe scavenging removed %s: %v", name, err) + } + } + if info, err := os.Stat(root); err != nil || info.Mode().Perm() != 0o700 { + t.Fatalf("run temp root mode = %v, %v", info, err) + } +} + +func TestStateSealerRoundTripTamperAndCrossClient(t *testing.T) { + sealer, err := OpenStateSealer(filepath.Join(t.TempDir(), "state.key")) + if err != nil { + t.Fatal(err) + } + blob, err := sealer.Seal("client-a", "claude-bin", []byte("gateway-local-state")) + if err != nil { + t.Fatal(err) + } + plain, err := sealer.Open(blob, "client-a", "claude-bin") + if err != nil || string(plain) != "gateway-local-state" { + t.Fatalf("round trip = %q, %v", plain, err) + } + tampered := []byte(blob) + if tampered[len(tampered)/2] == 'A' { + tampered[len(tampered)/2] = 'B' + } else { + tampered[len(tampered)/2] = 'A' + } + for _, tc := range []struct{ blob, client, provider string }{ + {string(tampered), "client-a", "claude-bin"}, + {blob, "client-b", "claude-bin"}, + {blob, "client-a", "grok-bin"}, + } { + if _, err := sealer.Open(tc.blob, tc.client, tc.provider); err == nil { + t.Fatalf("accepted tampered/foreign state: %+v", tc) + } + } +} + +func TestEnrollmentCreatesUniqueAuthenticatedClient(t *testing.T) { + dir := t.TempDir() + store, _ := OpenClientStore(filepath.Join(dir, "clients.json")) + policy := Policy{AllowProviders: []string{"openai"}, AllowSearch: true, MaxConcurrentInference: 1} + enrollment, bootstrap, err := store.CreateEnrollment("new-satellite", policy, time.Minute) + if err != nil { + t.Fatal(err) + } + if enrollment.ExpiresAt.Sub(enrollment.CreatedAt) != time.Minute { + t.Fatalf("enrollment TTL = %s", enrollment.ExpiresAt.Sub(enrollment.CreatedAt)) + } + sealer, _ := OpenStateSealer(filepath.Join(dir, "state.key")) + server, err := NewServer(ServerConfig{Config: &config.Config{Providers: map[string]config.ProviderConfig{}}, Clients: store, Sealer: sealer, Policy: Policy{}}) + if err != nil { + t.Fatal(err) + } + payload, _ := json.Marshal(protocol.EnrollmentRequest{Version: protocol.Version, Name: "new-satellite"}) + enroll := func(body []byte) *httptest.ResponseRecorder { + req := httptest.NewRequest(http.MethodPost, "/g1/enroll", bytes.NewReader(body)) + req.Header.Set("Authorization", "Bearer "+bootstrap) + rr := httptest.NewRecorder() + server.Handler().ServeHTTP(rr, req) + return rr + } + unsupportedPayload, _ := json.Marshal(protocol.EnrollmentRequest{Version: protocol.Version + 1, Name: "new-satellite"}) + unsupported := enroll(unsupportedPayload) + var versionError protocol.Error + if err := json.Unmarshal(unsupported.Body.Bytes(), &versionError); err != nil { + t.Fatal(err) + } + if unsupported.Code != http.StatusUpgradeRequired || unsupported.Header().Get(protocol.VersionHeader) != "1" || versionError.Code != "unsupported_version" || len(versionError.SupportedVersions) != 1 || versionError.SupportedVersions[0] != protocol.Version { + t.Fatalf("unsupported enrollment version = %d headers=%v body=%+v", unsupported.Code, unsupported.Header(), versionError) + } + rr := enroll(payload) + if rr.Code != http.StatusCreated { + t.Fatalf("enrollment status = %d: %s", rr.Code, rr.Body.String()) + } + var enrolled protocol.EnrollmentResponse + if err := json.Unmarshal(rr.Body.Bytes(), &enrolled); err != nil { + t.Fatal(err) + } + client, ok := store.Authenticate(enrolled.Token) + if !ok || client.ID != enrolled.ClientID || client.Name != "new-satellite" || len(client.Policy.AllowProviders) != 1 || client.Policy.MaxConcurrentInference != 1 { + t.Fatalf("enrolled client = %+v, authenticated=%t", client, ok) + } + if second := enroll(payload); second.Code != http.StatusUnauthorized { + t.Fatalf("reused enrollment status = %d, want 401", second.Code) + } + for _, path := range []string{filepath.Join(dir, "clients.json"), filepath.Join(dir, "clients.enrollments.json")} { + if strings.Contains(string(mustRead(t, path)), enrolled.Token) || strings.Contains(string(mustRead(t, path)), bootstrap) { + t.Fatalf("plaintext token persisted in %s", path) + } + } +} + +func TestEnrollmentRejectsUnrestrictedAndExpiredTokens(t *testing.T) { + store, err := OpenClientStore(filepath.Join(t.TempDir(), "clients.json")) + if err != nil { + t.Fatal(err) + } + if _, _, err := store.CreateEnrollment("unsafe", Policy{}, time.Minute); err == nil || !strings.Contains(err.Error(), "requires --allow-provider or --allow-model") { + t.Fatalf("unrestricted enrollment error = %v", err) + } + _, token, err := store.CreateEnrollment("short", Policy{AllowProviders: []string{"openai"}}, time.Millisecond) + if err != nil { + t.Fatal(err) + } + time.Sleep(2 * time.Millisecond) + if _, _, err := store.ConsumeEnrollment(token, "short"); err == nil || !strings.Contains(err.Error(), "expired") { + t.Fatalf("expired enrollment error = %v", err) + } +} + +func mustRead(t *testing.T, path string) []byte { + t.Helper() + data, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + return data +} + +func TestOversizedGatewayBodyReturnsSingleStructured413(t *testing.T) { + dir := t.TempDir() + store, _ := OpenClientStore(filepath.Join(dir, "clients.json")) + sealer, _ := OpenStateSealer(filepath.Join(dir, "state.key")) + server, err := NewServer(ServerConfig{Config: &config.Config{Providers: map[string]config.ProviderConfig{}}, Clients: store, Sealer: sealer, MaxBodyBytes: 8}) + if err != nil { + t.Fatal(err) + } + rr := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/", strings.NewReader(`{"version":1,"query":"too long"}`)) + var target map[string]any + if server.decodeJSON(rr, req, &target) { + t.Fatal("oversized body decoded") + } + if rr.Code != http.StatusRequestEntityTooLarge || strings.Count(rr.Body.String(), `"code"`) != 1 { + t.Fatalf("oversized response = %d %q", rr.Code, rr.Body.String()) + } +} + +func TestRunCancellationIsClientScopedAndVersioned(t *testing.T) { + dir := t.TempDir() + store, _ := OpenClientStore(filepath.Join(dir, "clients.json")) + clientA, tokenA, _ := store.Add("a", Policy{}) + _, tokenB, _ := store.Add("b", Policy{}) + sealer, _ := OpenStateSealer(filepath.Join(dir, "state.key")) + server, err := NewServer(ServerConfig{Config: &config.Config{Providers: map[string]config.ProviderConfig{}}, Clients: store, Sealer: sealer, Policy: Policy{}}) + if err != nil { + t.Fatal(err) + } + canceled := false + server.runs["run-a"] = &runState{clientID: clientA.ID, cancel: func() { canceled = true }, callbacks: make(map[string]chan llm.ToolExecutionResponse)} + + request := func(token, version string) *httptest.ResponseRecorder { + req := httptest.NewRequest(http.MethodDelete, "/g1/runs/run-a", nil) + req.Header.Set("Authorization", "Bearer "+token) + if version != "" { + req.Header.Set("Term-LLM-Gateway-Version", version) + } + rr := httptest.NewRecorder() + server.Handler().ServeHTTP(rr, req) + return rr + } + if rr := request(tokenB, "1"); rr.Code != http.StatusNotFound || canceled { + t.Fatalf("foreign cancellation = %d canceled=%t", rr.Code, canceled) + } + if rr := request(tokenA, ""); rr.Code != http.StatusUpgradeRequired || canceled { + t.Fatalf("version negotiation = %d canceled=%t", rr.Code, canceled) + } + if rr := request(tokenA, "1"); rr.Code != http.StatusNoContent || !canceled { + t.Fatalf("owner cancellation = %d canceled=%t", rr.Code, canceled) + } +} diff --git a/internal/gateway/server.go b/internal/gateway/server.go new file mode 100644 index 000000000..a55bec33e --- /dev/null +++ b/internal/gateway/server.go @@ -0,0 +1,1415 @@ +package gateway + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "log/slog" + "net/http" + "os" + "os/exec" + "path" + "path/filepath" + "sort" + "strings" + "sync" + "sync/atomic" + "time" + + "github.com/samsaffron/term-llm/internal/config" + "github.com/samsaffron/term-llm/internal/credentials" + "github.com/samsaffron/term-llm/internal/gateway/protocol" + "github.com/samsaffron/term-llm/internal/llm" + "github.com/samsaffron/term-llm/internal/providerhttp" + "github.com/samsaffron/term-llm/internal/search" + "golang.org/x/sync/singleflight" + "golang.org/x/time/rate" +) + +const ( + defaultMaxBodyBytes = 64 << 20 + DefaultUpstreamRetryAttempts = 3 + DefaultUpstreamRetryElapsed = 20 * time.Second +) + +type ProviderFactory func(*config.Config, string, string) (llm.Provider, error) +type ConfigLoader func() (*config.Config, error) + +type ServerConfig struct { + Config *config.Config + ConfigLoader ConfigLoader + Clients *ClientStore + Sealer *StateSealer + Usage UsageRecorder + ProviderFactory ProviderFactory + Searcher search.Searcher + FetchTool *llm.ReadURLTool + Policy Policy + MaxBodyBytes int64 + IdleTimeout time.Duration + ToolTimeout time.Duration + CatalogTTL time.Duration + ModelListTimeout time.Duration + UpstreamRetryAttempts int + UpstreamRetryMaxElapsed time.Duration + RunTempRoot string +} + +type runState struct { + clientID string + cancel context.CancelFunc + mu sync.Mutex + callbacks map[string]chan llm.ToolExecutionResponse +} + +type clientLimits struct { + inference chan struct{} + search chan struct{} + fetch chan struct{} + searchRPS *rate.Limiter + fetchRPS *rate.Limiter +} + +type inferenceExecution struct { + server *Server + client Client + envelope protocol.InferenceRequest + request llm.Request + entry protocol.CatalogEntry + provider llm.Provider + stream llm.Stream + ctx context.Context + cancel context.CancelFunc + release func() + tempDir string + started time.Time + total llm.Usage + finish sync.Once +} + +type inferenceRequestError struct { + Status int + Code string + Message string +} + +func (e *inferenceExecution) addUsage(use llm.Usage) { e.total.Add(use) } + +func (e *inferenceExecution) close(errorCode string) { + if e == nil { + return + } + e.finish.Do(func() { + if errorCode == "" && errors.Is(e.ctx.Err(), context.Canceled) { + errorCode = "canceled" + } + _ = e.stream.Close() + e.cancel() + if e.tempDir != "" { + _ = os.RemoveAll(e.tempDir) + } + if e.release != nil { + e.release() + } + e.server.recordUsage(e.client, e.envelope, e.request, e.total, errorCode, e.started) + }) +} + +type Server struct { + cfg ServerConfig + + configMu sync.RWMutex + config *config.Config + + catalogMu sync.RWMutex + catalog protocol.Catalog + catalogFetchedAt time.Time + catalogProviderFetched map[string]time.Time + catalogRefresh singleflight.Group + configRefreshMu sync.Mutex + configFetchedAt time.Time + + runsMu sync.RWMutex + runs map[string]*runState + + limitsMu sync.Mutex + limits map[string]*clientLimits +} + +func NewServer(cfg ServerConfig) (*Server, error) { + if cfg.Config == nil || cfg.Clients == nil || cfg.Sealer == nil { + return nil, fmt.Errorf("gateway server requires config, client store, and state sealer") + } + if cfg.ProviderFactory == nil { + cfg.ProviderFactory = llm.NewProviderByName + } + if cfg.MaxBodyBytes <= 0 { + cfg.MaxBodyBytes = defaultMaxBodyBytes + } + if cfg.IdleTimeout <= 0 { + cfg.IdleTimeout = 5 * time.Minute + } + if cfg.ToolTimeout <= 0 { + cfg.ToolTimeout = 10 * time.Minute + } + if cfg.CatalogTTL <= 0 { + cfg.CatalogTTL = 5 * time.Minute + } + if cfg.ModelListTimeout <= 0 { + cfg.ModelListTimeout = 5 * time.Second + } + if cfg.UpstreamRetryAttempts <= 0 { + cfg.UpstreamRetryAttempts = DefaultUpstreamRetryAttempts + } + if cfg.UpstreamRetryMaxElapsed <= 0 { + cfg.UpstreamRetryMaxElapsed = DefaultUpstreamRetryElapsed + } + if strings.TrimSpace(cfg.RunTempRoot) != "" { + cfg.RunTempRoot = filepath.Clean(cfg.RunTempRoot) + if err := prepareRunTempRoot(cfg.RunTempRoot); err != nil { + return nil, err + } + } + return &Server{ + cfg: cfg, config: cfg.Config, configFetchedAt: time.Now().UTC(), + catalog: protocol.Catalog{Version: protocol.Version}, catalogProviderFetched: make(map[string]time.Time), + runs: make(map[string]*runState), limits: make(map[string]*clientLimits), + }, nil +} + +func (s *Server) Handler() http.Handler { return http.HandlerFunc(s.serveHTTP) } + +func (s *Server) serveHTTP(w http.ResponseWriter, r *http.Request) { + clean := path.Clean(r.URL.Path) + if clean == "/v1/responses" || clean == "/v1/models" { + client, ok := s.authenticateResponses(w, r) + if !ok { + return + } + switch { + case clean == "/v1/responses" && r.Method == http.MethodPost: + s.handleResponses(w, r, client) + case clean == "/v1/models" && r.Method == http.MethodGet: + s.handleResponsesModels(w, r, client) + default: + w.Header().Set("Allow", map[string]string{"/v1/responses": http.MethodPost, "/v1/models": http.MethodGet}[clean]) + s.writeResponsesError(w, http.StatusMethodNotAllowed, "method_not_allowed", "HTTP method is not supported for this endpoint", "") + } + return + } + if clean == "/g1/health" && r.Method == http.MethodGet { + s.writeJSON(w, http.StatusOK, protocol.Health{Version: protocol.Version, Status: "ok"}) + return + } + if clean == "/g1/enroll" && r.Method == http.MethodPost { + s.handleEnroll(w, r) + return + } + if !s.checkVersion(w, r) { + return + } + client, ok := s.authenticate(w, r) + if !ok { + return + } + switch { + case clean == "/g1/catalog" && r.Method == http.MethodGet: + s.handleCatalog(w, r, client) + case clean == "/g1/inference" && r.Method == http.MethodPost: + s.handleInference(w, r, client) + case clean == "/g1/search" && r.Method == http.MethodPost: + s.handleSearch(w, r, client) + case clean == "/g1/fetch" && r.Method == http.MethodPost: + s.handleFetch(w, r, client) + case strings.HasPrefix(clean, "/g1/runs/"): + s.handleRun(w, r, client, clean) + default: + s.writeError(w, http.StatusNotFound, "not_found", "gateway endpoint not found", "") + } +} + +func (s *Server) checkVersion(w http.ResponseWriter, r *http.Request) bool { + if r.Header.Get(protocol.VersionHeader) != "1" { + s.writeUnsupportedVersion(w, "") + return false + } + return true +} + +func (s *Server) authenticate(w http.ResponseWriter, r *http.Request) (Client, bool) { + client, ok := s.authenticateBearer(r) + if !ok { + s.writeError(w, http.StatusUnauthorized, "gateway_client_unauthorized", "gateway client credential is missing or invalid; update gateway.token/token_file on the satellite", "") + return Client{}, false + } + return client, true +} + +func (s *Server) authenticateResponses(w http.ResponseWriter, r *http.Request) (Client, bool) { + client, ok := s.authenticateBearer(r) + if !ok { + s.writeResponsesError(w, http.StatusUnauthorized, "gateway_client_unauthorized", "gateway client bearer credential is missing, invalid, or revoked", "") + return Client{}, false + } + return client, true +} + +func (s *Server) authenticateBearer(r *http.Request) (Client, bool) { + authorization := r.Header.Get("Authorization") + if !strings.HasPrefix(authorization, "Bearer ") { + return Client{}, false + } + token := strings.TrimSpace(strings.TrimPrefix(authorization, "Bearer ")) + if token == "" { + return Client{}, false + } + return s.cfg.Clients.Authenticate(token) +} + +func (s *Server) handleEnroll(w http.ResponseWriter, r *http.Request) { + authorization := r.Header.Get("Authorization") + if !strings.HasPrefix(authorization, "Bearer ") { + s.writeError(w, http.StatusUnauthorized, "unauthorized", "invalid or expired gateway enrollment token", "") + return + } + var req protocol.EnrollmentRequest + if !s.decodeJSON(w, r, &req) { + return + } + if req.Version != protocol.Version { + s.writeUnsupportedVersion(w, "") + return + } + token := strings.TrimSpace(strings.TrimPrefix(authorization, "Bearer ")) + client, clientToken, err := s.cfg.Clients.ConsumeEnrollment(token, req.Name) + if err != nil { + slog.Warn("gateway enrollment rejected", "reason", err) + s.writeError(w, http.StatusUnauthorized, "enrollment_rejected", "invalid, expired, or already-used gateway enrollment token; ask the gateway operator for a new token", "") + return + } + s.writeJSON(w, http.StatusCreated, protocol.EnrollmentResponse{Version: protocol.Version, ClientID: client.ID, Token: clientToken}) +} + +func (s *Server) handleCatalog(w http.ResponseWriter, r *http.Request, client Client) { + catalog, err := s.currentCatalog(r.Context()) + if err != nil { + slog.Error("refresh gateway catalog", "error", err) + s.writeError(w, http.StatusServiceUnavailable, "catalog_unavailable", "gateway catalog is temporarily unavailable; retry or check provider configuration on the gateway", "") + return + } + catalog = filterCatalog(catalog, s.cfg.Policy, client.Policy) + data, err := json.Marshal(catalog) + if err != nil { + s.writeError(w, http.StatusInternalServerError, "internal", "catalog unavailable", "") + return + } + sum := sha256.Sum256(data) + etag := `"` + hex.EncodeToString(sum[:]) + `"` + w.Header().Set("ETag", etag) + w.Header().Set("Cache-Control", "private, max-age=0, must-revalidate") + if r.Header.Get("If-None-Match") == etag { + w.WriteHeader(http.StatusNotModified) + return + } + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + _, _ = w.Write(data) +} + +func (s *Server) handleInference(w http.ResponseWriter, r *http.Request, client Client) { + var envelope protocol.InferenceRequest + if !s.decodeJSON(w, r, &envelope) { + return + } + if envelope.Version != protocol.Version { + s.writeUnsupportedVersion(w, envelope.RequestID) + return + } + if strings.TrimSpace(envelope.RequestID) == "" || strings.TrimSpace(envelope.Provider) == "" { + s.writeError(w, http.StatusBadRequest, "invalid_request", "request_id and provider are required", envelope.RequestID) + return + } + providerReq, err := llm.DecodeGatewayRequest(envelope.Request) + if err != nil { + s.writeError(w, http.StatusBadRequest, "invalid_request", err.Error(), envelope.RequestID) + return + } + execution, requestErr := s.startInference(r.Context(), client, envelope, providerReq, false) + if requestErr != nil { + s.writeError(w, requestErr.Status, requestErr.Code, requestErr.Message, envelope.RequestID) + return + } + errorCode := "" + defer func() { execution.close(errorCode) }() + providerReq = execution.request + provider := execution.provider + stream := execution.stream + ctx := execution.ctx + entry := execution.entry + + runID, err := randomSecret("run", 16) + if err != nil { + errorCode = "internal" + s.writeError(w, http.StatusInternalServerError, "internal", "could not create gateway run", envelope.RequestID) + return + } + run := &runState{clientID: client.ID, cancel: execution.cancel, callbacks: make(map[string]chan llm.ToolExecutionResponse)} + s.runsMu.Lock() + s.runs[runID] = run + s.runsMu.Unlock() + defer func() { + s.runsMu.Lock() + delete(s.runs, runID) + s.runsMu.Unlock() + }() + var lastActivity atomic.Int64 + lastActivity.Store(time.Now().UnixNano()) + go s.watchIdle(ctx, execution.cancel, &lastActivity) + + flusher, ok := w.(http.Flusher) + if !ok { + s.writeError(w, http.StatusInternalServerError, "streaming_unsupported", "streaming is unavailable", envelope.RequestID) + return + } + w.Header().Set("Content-Type", "text/event-stream") + w.Header().Set("Cache-Control", "no-cache, no-transform") + w.Header().Set("X-Accel-Buffering", "no") + w.WriteHeader(http.StatusOK) + if !writeSSE(w, protocol.StreamRecord{Version: protocol.Version, Type: "run", RequestID: envelope.RequestID, RunID: runID}) { + return + } + flusher.Flush() + + for { + event, recvErr := stream.Recv() + lastActivity.Store(time.Now().UnixNano()) + if recvErr == io.EOF { + break + } + if recvErr != nil { + _, errorCode = s.logProviderStreamError(execution, recvErr) + _ = writeSSE(w, protocol.StreamRecord{Version: protocol.Version, Type: "error", RequestID: envelope.RequestID, RunID: runID, Error: &protocol.Error{Code: errorCode, Message: safeProviderErrorMessage(errorCode, envelope.Provider), RequestID: envelope.RequestID}}) + flusher.Flush() + break + } + if event.Type == llm.EventUsage && event.Use != nil { + execution.addUsage(*event.Use) + } + var publicError *protocol.Error + if event.Err != nil { + _, classified := classifyProviderError(event.Err, entry.Type) + publicError = &protocol.Error{Code: classified, Message: safeProviderErrorMessage(classified, envelope.Provider), RequestID: envelope.RequestID} + slog.Error("gateway provider event failed", "request_id", envelope.RequestID, "provider", envelope.Provider, "event_type", event.Type, "error", event.Err) + if event.Type == llm.EventError { + errorCode = classified + _ = writeSSE(w, protocol.StreamRecord{Version: protocol.Version, Type: "error", RequestID: envelope.RequestID, RunID: runID, Error: publicError}) + flusher.Flush() + break + } + } + wireEvent, encodeErr := llm.EncodeGatewayEvent(event, publicError) + if encodeErr != nil { + errorCode = "encoding_error" + break + } + record := protocol.StreamRecord{Version: protocol.Version, Type: "event", RequestID: envelope.RequestID, RunID: runID, Event: wireEvent} + if event.ToolResponse != nil { + callbackID, idErr := randomSecret("callback", 8) + if idErr != nil { + errorCode = "internal" + _ = writeSSE(w, protocol.StreamRecord{Version: protocol.Version, Type: "error", RequestID: envelope.RequestID, RunID: runID, Error: &protocol.Error{Code: errorCode, Message: "gateway callback unavailable", RequestID: envelope.RequestID}}) + flusher.Flush() + break + } + callback := make(chan llm.ToolExecutionResponse, 1) + run.mu.Lock() + run.callbacks[callbackID] = callback + run.mu.Unlock() + record.Type = "tool_callback" + record.CallbackPath = "/g1/runs/" + runID + "/tools/" + callbackID + if !writeSSE(w, record) { + return + } + flusher.Flush() + // A pending satellite callback is active work, not provider idleness. + lastActivity.Store(time.Now().Add(s.cfg.ToolTimeout).UnixNano()) + timer := time.NewTimer(s.cfg.ToolTimeout) + var response llm.ToolExecutionResponse + select { + case response = <-callback: + if !timer.Stop() { + select { + case <-timer.C: + default: + } + } + case <-timer.C: + response.Err = fmt.Errorf("satellite tool callback timed out") + case <-ctx.Done(): + timer.Stop() + return + } + select { + case event.ToolResponse <- response: + case <-ctx.Done(): + return + } + lastActivity.Store(time.Now().UnixNano()) + run.mu.Lock() + delete(run.callbacks, callbackID) + run.mu.Unlock() + continue + } + if !writeSSE(w, record) { + return + } + flusher.Flush() + } + if errorCode == "" && errors.Is(ctx.Err(), context.Canceled) { + errorCode = "canceled" + } + if exporter, ok := provider.(llm.ProviderStateExporter); ok { + if plain, valid := exporter.ExportProviderState(); valid { + if sealed, sealErr := s.cfg.Sealer.Seal(client.ID, envelope.Provider, plain); sealErr == nil { + _ = writeSSE(w, protocol.StreamRecord{Version: protocol.Version, Type: "state", RequestID: envelope.RequestID, RunID: runID, State: sealed}) + flusher.Flush() + } + } + } + if errorCode == "" { + _ = writeSSE(w, protocol.StreamRecord{Version: protocol.Version, Type: "done", RequestID: envelope.RequestID, RunID: runID}) + flusher.Flush() + } +} + +func (s *Server) watchIdle(ctx context.Context, cancel context.CancelFunc, last *atomic.Int64) { + interval := s.cfg.IdleTimeout / 4 + if interval > time.Second { + interval = time.Second + } + if interval <= 0 { + interval = 100 * time.Millisecond + } + ticker := time.NewTicker(interval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + if time.Since(time.Unix(0, last.Load())) > s.cfg.IdleTimeout { + cancel() + return + } + } + } +} + +func (s *Server) handleRun(w http.ResponseWriter, r *http.Request, client Client, clean string) { + parts := strings.Split(strings.TrimPrefix(clean, "/g1/runs/"), "/") + if len(parts) < 1 || parts[0] == "" { + s.writeError(w, http.StatusNotFound, "not_found", "run not found", "") + return + } + s.runsMu.RLock() + run := s.runs[parts[0]] + s.runsMu.RUnlock() + if run == nil || run.clientID != client.ID { + s.writeError(w, http.StatusNotFound, "not_found", "run not found", "") + return + } + if r.Method == http.MethodDelete && len(parts) == 1 { + run.cancel() + w.WriteHeader(http.StatusNoContent) + return + } + if r.Method == http.MethodPost && len(parts) == 3 && parts[1] == "tools" { + var req protocol.ToolResultRequest + if !s.decodeJSON(w, r, &req) || req.Version != protocol.Version { + return + } + response, err := llm.DecodeGatewayToolResponse(req.Result) + if err != nil { + s.writeError(w, http.StatusBadRequest, "invalid_tool_result", err.Error(), "") + return + } + run.mu.Lock() + callback := run.callbacks[parts[2]] + run.mu.Unlock() + if callback == nil { + s.writeError(w, http.StatusNotFound, "not_found", "tool callback not found", "") + return + } + select { + case callback <- response: + w.WriteHeader(http.StatusNoContent) + case <-r.Context().Done(): + } + return + } + s.writeError(w, http.StatusNotFound, "not_found", "run endpoint not found", "") +} + +func (s *Server) handleSearch(w http.ResponseWriter, r *http.Request, client Client) { + if s.cfg.Searcher == nil || !s.cfg.Policy.AllowSearch || !client.Policy.AllowSearch { + s.writeError(w, http.StatusForbidden, "policy_denied", "gateway search is unavailable", "") + return + } + release, code, ok := s.acquireTool(client, true) + if !ok { + s.writeError(w, http.StatusTooManyRequests, code, "gateway search limit reached for this client; wait and retry", "") + return + } + defer release() + var req protocol.SearchRequest + if !s.decodeJSON(w, r, &req) || req.Version != protocol.Version { + return + } + if req.MaxResults <= 0 || req.MaxResults > 100 { + req.MaxResults = 20 + } + results, err := s.cfg.Searcher.Search(r.Context(), req.Query, req.MaxResults) + if err != nil { + s.writeError(w, http.StatusBadGateway, "search_failed", "gateway search failed upstream; retry, then check gateway-side search configuration", "") + return + } + out := make([]protocol.SearchResult, 0, len(results)) + for _, result := range results { + out = append(out, protocol.SearchResult{Title: result.Title, URL: result.URL, Snippet: result.Snippet}) + } + s.writeJSON(w, http.StatusOK, protocol.SearchResponse{Version: protocol.Version, Results: out}) +} + +func (s *Server) handleFetch(w http.ResponseWriter, r *http.Request, client Client) { + if s.cfg.FetchTool == nil || !s.cfg.Policy.AllowFetch || !client.Policy.AllowFetch { + s.writeError(w, http.StatusForbidden, "policy_denied", "gateway fetch is unavailable", "") + return + } + release, code, ok := s.acquireTool(client, false) + if !ok { + s.writeError(w, http.StatusTooManyRequests, code, "gateway fetch limit reached for this client; wait and retry", "") + return + } + defer release() + var req protocol.FetchRequest + if !s.decodeJSON(w, r, &req) || req.Version != protocol.Version { + return + } + args, _ := json.Marshal(map[string]string{"url": req.URL}) + output, err := s.cfg.FetchTool.Execute(r.Context(), args) + if err != nil { + s.writeError(w, http.StatusBadGateway, "fetch_failed", "gateway fetch failed upstream; retry, then check gateway-side fetch configuration", "") + return + } + s.writeJSON(w, http.StatusOK, protocol.FetchResponse{Version: protocol.Version, Content: output.Content}) +} + +func prepareRunTempRoot(root string) error { + info, err := os.Lstat(root) + if err != nil { + if !os.IsNotExist(err) { + return fmt.Errorf("inspect gateway run temp root: %w", err) + } + if err := os.MkdirAll(root, 0o700); err != nil { + return fmt.Errorf("create gateway run temp root: %w", err) + } + info, err = os.Lstat(root) + if err != nil { + return fmt.Errorf("inspect created gateway run temp root: %w", err) + } + } + if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() { + return fmt.Errorf("gateway run temp root must be a real directory: %s", root) + } + if err := os.Chmod(root, 0o700); err != nil { + return fmt.Errorf("secure gateway run temp root: %w", err) + } + entries, err := os.ReadDir(root) + if err != nil { + return fmt.Errorf("scan gateway run temp root: %w", err) + } + for _, entry := range entries { + // This state-owned directory is reserved for gateway run directories. The + // prefix and directory checks retain unrelated operator-created content. + if !entry.IsDir() || !strings.HasPrefix(entry.Name(), "run-") || strings.Contains(entry.Name(), string(os.PathSeparator)) { + continue + } + if err := os.RemoveAll(filepath.Join(root, entry.Name())); err != nil { + return fmt.Errorf("remove stale gateway run directory %q: %w", entry.Name(), err) + } + } + return nil +} + +func (s *Server) newRunTempDir() (string, error) { + if s.cfg.RunTempRoot == "" { + return os.MkdirTemp("", "term-llm-gateway-run-*") + } + return os.MkdirTemp(s.cfg.RunTempRoot, "run-") +} + +func (s *Server) centralConfig() *config.Config { + clone := *s.currentConfig() + clone.Gateway = config.GatewayConfig{} + return &clone +} + +func (s *Server) currentConfig() *config.Config { + s.configMu.RLock() + defer s.configMu.RUnlock() + return s.config +} + +func (s *Server) decodeJSON(w http.ResponseWriter, r *http.Request, target any) bool { + data, err := io.ReadAll(io.LimitReader(r.Body, s.cfg.MaxBodyBytes+1)) + if err != nil { + s.writeError(w, http.StatusBadRequest, "invalid_json", "could not read gateway request body", "") + return false + } + if int64(len(data)) > s.cfg.MaxBodyBytes { + s.writeError(w, http.StatusRequestEntityTooLarge, "request_too_large", "gateway request body exceeds the configured limit", "") + return false + } + dec := json.NewDecoder(bytes.NewReader(data)) + dec.DisallowUnknownFields() + if err := dec.Decode(target); err != nil { + s.writeError(w, http.StatusBadRequest, "invalid_json", "invalid gateway request body", "") + return false + } + if err := dec.Decode(&struct{}{}); !errors.Is(err, io.EOF) { + s.writeError(w, http.StatusBadRequest, "invalid_json", "gateway request must contain one JSON value", "") + return false + } + return true +} + +func (s *Server) writeUnsupportedVersion(w http.ResponseWriter, requestID string) { + s.writeJSON(w, http.StatusUpgradeRequired, protocol.Error{ + Code: "unsupported_version", + Message: "gateway protocol version 1 is required", + RequestID: requestID, + SupportedVersions: []int{protocol.Version}, + }) +} + +func (s *Server) writeError(w http.ResponseWriter, status int, code, message, requestID string) { + s.writeJSON(w, status, protocol.Error{Code: code, Message: message, RequestID: requestID}) +} + +func (s *Server) writeJSON(w http.ResponseWriter, status int, value any) { + w.Header().Set("Content-Type", "application/json") + w.Header().Set(protocol.VersionHeader, "1") + w.WriteHeader(status) + _ = json.NewEncoder(w).Encode(value) +} + +func writeSSE(w io.Writer, record protocol.StreamRecord) bool { + data, err := json.Marshal(record) + if err != nil { + return false + } + _, err = fmt.Fprintf(w, "event: gateway\ndata: %s\n\n", data) + return err == nil +} + +func (s *Server) recordFailure(client Client, envelope protocol.InferenceRequest, req llm.Request, code string, started time.Time) { + s.recordUsage(client, envelope, req, llm.Usage{}, code, started) +} + +func (s *Server) recordUsage(client Client, envelope protocol.InferenceRequest, req llm.Request, use llm.Usage, code string, started time.Time) { + if s.cfg.Usage == nil { + return + } + record := UsageRecord{ + StartedAt: started.UTC(), CompletedAt: time.Now().UTC(), ClientID: client.ID, ClientName: client.Name, + ProviderKey: envelope.Provider, Model: req.Model, RequestID: envelope.RequestID, SessionID: req.SessionID, + InputTokens: use.InputTokens, OutputTokens: use.OutputTokens, CachedInputTokens: use.CachedInputTokens, + CacheWriteTokens: use.CacheWriteTokens, ReasoningTokens: use.ReasoningTokens, ErrorCode: code, + } + record.CostUSD = estimateUsageCost(envelope.Provider, req.Model, use) + if err := s.cfg.Usage.Record(record); err != nil { + slog.Error("record gateway usage", "error", err) + } +} + +type providerCatalogResult struct { + entry protocol.CatalogEntry + found bool +} + +func (s *Server) refreshConfigIfStale() error { + if s.cfg.ConfigLoader == nil { + return nil + } + s.configMu.RLock() + fetchedAt := s.configFetchedAt + s.configMu.RUnlock() + if !fetchedAt.IsZero() && time.Since(fetchedAt) < s.cfg.CatalogTTL { + return nil + } + s.configRefreshMu.Lock() + defer s.configRefreshMu.Unlock() + s.configMu.RLock() + fetchedAt = s.configFetchedAt + s.configMu.RUnlock() + if !fetchedAt.IsZero() && time.Since(fetchedAt) < s.cfg.CatalogTTL { + return nil + } + loaded, err := s.cfg.ConfigLoader() + if err != nil { + // Keep the last known-good configuration and avoid coupling inference to a + // transient local reload failure. + s.configMu.Lock() + s.configFetchedAt = time.Now().UTC() + s.configMu.Unlock() + return fmt.Errorf("reload gateway config: %w", err) + } + if loaded != nil { + clone := *loaded + clone.Gateway = config.GatewayConfig{} + s.configMu.Lock() + s.config = &clone + s.configFetchedAt = time.Now().UTC() + s.configMu.Unlock() + } + return nil +} + +func (s *Server) cachedCatalogProvider(name string) (protocol.CatalogEntry, time.Time, bool) { + s.catalogMu.RLock() + defer s.catalogMu.RUnlock() + for _, entry := range s.catalog.Providers { + if entry.Key == name { + return entry, s.catalogProviderFetched[name], true + } + } + return protocol.CatalogEntry{}, time.Time{}, false +} + +func (s *Server) storeCatalogProviderAt(entry protocol.CatalogEntry, fetchedAt time.Time) { + s.catalogMu.Lock() + defer s.catalogMu.Unlock() + replaced := false + for i := range s.catalog.Providers { + if s.catalog.Providers[i].Key == entry.Key { + s.catalog.Providers[i] = entry + replaced = true + break + } + } + if !replaced { + s.catalog.Providers = append(s.catalog.Providers, entry) + } + sort.Slice(s.catalog.Providers, func(i, j int) bool { return s.catalog.Providers[i].Key < s.catalog.Providers[j].Key }) + now := time.Now().UTC() + s.catalog.Version = protocol.Version + s.catalog.GeneratedAt = now + s.catalog.Features = protocol.CatalogFeatures{Search: s.cfg.Searcher != nil, Fetch: s.cfg.FetchTool != nil} + s.catalogFetchedAt = now + s.catalogProviderFetched[entry.Key] = fetchedAt +} + +func (s *Server) storeCatalogProvider(entry protocol.CatalogEntry) { + s.storeCatalogProviderAt(entry, time.Now().UTC()) +} + +func (s *Server) providerConfigured(name string) bool { + cfg := s.currentConfig() + _, ok := cfg.Providers[name] + return ok && cfg.IsExplicitProvider(name) +} + +func configuredCatalogEntry(cfg *config.Config, name string) (protocol.CatalogEntry, bool) { + if cfg == nil || !cfg.IsExplicitProvider(name) { + return protocol.CatalogEntry{}, false + } + pc, ok := cfg.Providers[name] + if !ok { + return protocol.CatalogEntry{}, false + } + providerType := config.InferProviderType(name, pc.Type) + entry := protocol.CatalogEntry{ + Key: name, Type: string(providerType), CLI: isCLIProvider(providerType), + AllowUnlistedModels: allowUnlistedModels(providerType, pc), + Capabilities: catalogCapabilities(providerType), + } + if err := nonInteractiveAuthReady(entry.Type); err != nil || !providerCredentialReady(cfg, name, providerType) { + return protocol.CatalogEntry{}, false + } + ids := append([]string(nil), pc.Models...) + if pc.Model != "" { + ids = append(ids, pc.Model) + } + seen := make(map[string]bool) + for _, id := range ids { + id = strings.TrimSpace(id) + if id == "" || seen[id] { + continue + } + seen[id] = true + inputPrice, outputPrice := -1.0, -1.0 + if in, out, known := llm.PricingForProviderModel(name, id); known { + inputPrice, outputPrice = in, out + } + entry.Models = append(entry.Models, protocol.Model{ + ID: id, InputLimit: llm.InputLimitForProviderModel(name, id), OutputLimit: llm.OutputLimitForModel(id), + InputPrice: inputPrice, OutputPrice: outputPrice, ReasoningEfforts: llm.ReasoningEffortsForProviderModel(name, id), + }) + } + return entry, len(entry.Models) > 0 +} + +func (s *Server) refreshCatalogProvider(name string) <-chan singleflight.Result { + return s.catalogRefresh.DoChan(name, func() (any, error) { + if !s.providerConfigured(name) { + return providerCatalogResult{}, nil + } + refreshCtx, cancel := context.WithTimeout(context.Background(), s.cfg.ModelListTimeout) + defer cancel() + entry, err := buildCatalogEntry(refreshCtx, cloneConfigForCatalog(s.currentConfig()), s.cfg.ProviderFactory, name) + if err != nil { + slog.Error("refresh gateway provider catalog", "provider", name, "error", err) + // A configured fallback entry remains safe to use even when live model + // listing failed. Otherwise retain a last-known-good provider entry. + if len(entry.Models) == 0 && !entry.AllowUnlistedModels { + if stale, _, ok := s.cachedCatalogProvider(name); ok { + return providerCatalogResult{entry: stale, found: true}, nil + } + return providerCatalogResult{}, err + } + } + s.storeCatalogProvider(entry) + return providerCatalogResult{entry: entry, found: true}, nil + }) +} + +// currentCatalogProvider refreshes only the provider required for inference. +// Once an entry exists it is stale-first: the request uses it immediately while +// a bounded refresh proceeds in the background. +func (s *Server) currentCatalogProvider(ctx context.Context, name string) (protocol.CatalogEntry, bool, error) { + if err := s.refreshConfigIfStale(); err != nil { + slog.Error("reload gateway config", "error", err) + } + if !s.providerConfigured(name) { + return protocol.CatalogEntry{}, false, nil + } + cached, fetchedAt, ok := s.cachedCatalogProvider(name) + if ok && time.Since(fetchedAt) < s.cfg.CatalogTTL { + return cached, true, nil + } + if !ok { + if configured, valid := configuredCatalogEntry(s.currentConfig(), name); valid { + cached, ok = configured, true + // A zero fetch time makes this an immediately usable stale entry while + // live metadata is refreshed asynchronously. + s.storeCatalogProviderAt(configured, time.Time{}) + } + } + result := s.refreshCatalogProvider(name) + if ok { + // The channel must remain consumable so singleflight can complete, but this + // request never waits for unrelated or stale model-list I/O. + return cached, true, nil + } + select { + case <-ctx.Done(): + return protocol.CatalogEntry{}, false, ctx.Err() + case loaded := <-result: + if loaded.Err != nil { + return protocol.CatalogEntry{}, false, loaded.Err + } + value := loaded.Val.(providerCatalogResult) + return value.entry, value.found, nil + } +} + +func (s *Server) currentCatalog(ctx context.Context) (protocol.Catalog, error) { + if err := s.refreshConfigIfStale(); err != nil { + slog.Error("reload gateway config", "error", err) + } + cfg := s.currentConfig() + names := cfg.ExplicitProviderNames() + type result struct { + entry protocol.CatalogEntry + found bool + err error + } + results := make(chan result, len(names)) + for _, name := range names { + name := name + go func() { + entry, found, err := s.currentCatalogProvider(ctx, name) + results <- result{entry: entry, found: found, err: err} + }() + } + catalog := protocol.Catalog{Version: protocol.Version, GeneratedAt: time.Now().UTC(), Features: protocol.CatalogFeatures{Search: s.cfg.Searcher != nil, Fetch: s.cfg.FetchTool != nil}} + var firstErr error + for range names { + result := <-results + if result.err != nil { + if firstErr == nil { + firstErr = result.err + } + continue + } + if result.found { + catalog.Providers = append(catalog.Providers, result.entry) + } + } + sort.Slice(catalog.Providers, func(i, j int) bool { return catalog.Providers[i].Key < catalog.Providers[j].Key }) + if len(catalog.Providers) == 0 && len(names) > 0 && firstErr != nil { + return protocol.Catalog{}, firstErr + } + return catalog, nil +} + +func buildCatalog(ctx context.Context, cfg *config.Config, factory ProviderFactory, hasSearch, hasFetch bool) (protocol.Catalog, map[string]error) { + catalog := protocol.Catalog{Version: protocol.Version, GeneratedAt: time.Now().UTC(), Features: protocol.CatalogFeatures{Search: hasSearch, Fetch: hasFetch}} + failed := make(map[string]error) + if cfg == nil { + return catalog, failed + } + type result struct { + entry protocol.CatalogEntry + err error + } + names := cfg.ExplicitProviderNames() + results := make(chan result, len(names)) + var wg sync.WaitGroup + for _, name := range names { + name := name + wg.Add(1) + go func() { + defer wg.Done() + providerConfig := cloneConfigForCatalog(cfg) + entry, err := buildCatalogEntry(ctx, providerConfig, factory, name) + results <- result{entry: entry, err: err} + }() + } + wg.Wait() + close(results) + for result := range results { + if result.err != nil { + failed[result.entry.Key] = result.err + if len(result.entry.Models) == 0 && !result.entry.AllowUnlistedModels { + continue + } + } + catalog.Providers = append(catalog.Providers, result.entry) + } + sort.Slice(catalog.Providers, func(i, j int) bool { return catalog.Providers[i].Key < catalog.Providers[j].Key }) + return catalog, failed +} + +func cloneConfigForCatalog(cfg *config.Config) *config.Config { + clone := *cfg + clone.Providers = make(map[string]config.ProviderConfig, len(cfg.Providers)) + for name, provider := range cfg.Providers { + clone.Providers[name] = provider + } + return &clone +} + +func buildCatalogEntry(ctx context.Context, cfg *config.Config, factory ProviderFactory, key string) (protocol.CatalogEntry, error) { + pc, ok := cfg.Providers[key] + if !ok { + return protocol.CatalogEntry{Key: key}, fmt.Errorf("provider configuration disappeared") + } + providerType := config.InferProviderType(key, pc.Type) + entry := protocol.CatalogEntry{Key: key, Type: string(providerType), CLI: isCLIProvider(providerType), AllowUnlistedModels: allowUnlistedModels(providerType, pc)} + if err := nonInteractiveAuthReady(entry.Type); err != nil { + return entry, err + } + provider, err := factory(cfg, key, pc.Model) + if err != nil { + return entry, err + } + if !providerCredentialReady(cfg, key, providerType) { + return entry, fmt.Errorf("provider credential is not configured") + } + entry.Capabilities = llm.CapabilitiesToGatewayProtocol(provider.Capabilities()) + if entry.Capabilities == (protocol.Capabilities{}) { + entry.Capabilities = catalogCapabilities(providerType) + } + + ids := append([]string(nil), pc.Models...) + liveModels := []llm.ModelInfo(nil) + var listErr error + if lister, ok := provider.(interface { + ListModels(context.Context) ([]llm.ModelInfo, error) + }); ok { + liveModels, err = lister.ListModels(ctx) + if err != nil && !errors.Is(err, llm.ErrListModelsUnsupported) { + status, code := classifyProviderError(err, entry.Type) + if status == http.StatusUnauthorized || code == "provider_api_key_unauthenticated" || code == "provider_oauth_unauthenticated" { + return entry, fmt.Errorf("list live models: %w", err) + } + listErr = fmt.Errorf("list live models; using fallback: %w", err) + liveModels = nil + } + } + if len(liveModels) > 0 { + for _, model := range liveModels { + if strings.TrimSpace(model.ID) == "" { + continue + } + inputPrice, outputPrice := catalogLivePricing(key, providerType, model) + entry.Models = append(entry.Models, protocol.Model{ + ID: model.ID, DisplayName: model.DisplayName, Created: model.Created, OwnedBy: model.OwnedBy, + InputLimit: model.InputLimit, InputPrice: inputPrice, OutputPrice: outputPrice, + ReasoningEfforts: model.ReasoningEfforts, DefaultReasoningEffort: model.DefaultReasoningEffort, ReasoningModes: model.ReasoningModes, + }) + } + } else { + if len(ids) == 0 { + ids = llm.ResolveProviderModelIDs(key) + } + if pc.Model != "" { + ids = append(ids, pc.Model) + } + seen := make(map[string]bool) + for _, id := range ids { + id = strings.TrimSpace(id) + if id == "" || seen[id] { + continue + } + seen[id] = true + inputPrice, outputPrice := -1.0, -1.0 + if in, out, known := llm.PricingForProviderModel(key, id); known { + inputPrice, outputPrice = in, out + } + entry.Models = append(entry.Models, protocol.Model{ID: id, InputLimit: llm.InputLimitForProviderModel(key, id), OutputLimit: llm.OutputLimitForModel(id), InputPrice: inputPrice, OutputPrice: outputPrice, ReasoningEfforts: llm.ReasoningEffortsForProviderModel(key, id)}) + } + } + if len(entry.Models) == 0 && !entry.AllowUnlistedModels { + if listErr != nil { + return entry, listErr + } + return entry, fmt.Errorf("provider has no discoverable or configured models") + } + return entry, listErr +} + +func catalogLivePricing(provider string, providerType config.ProviderType, model llm.ModelInfo) (float64, float64) { + if model.InputPrice != 0 || model.OutputPrice != 0 { + return model.InputPrice, model.OutputPrice + } + if input, output, known := llm.PricingForProviderModel(provider, model.ID); known { + return input, output + } + // These catalog APIs explicitly report free models as zero prices. For APIs + // that do not return pricing, zero is ambiguous and must not be advertised as + // free to satellites. + switch providerType { + case config.ProviderTypeOpenRouter, config.ProviderTypeZen, config.ProviderTypeVenice, + config.ProviderTypeNearAI, config.ProviderTypeSambaNova: + return 0, 0 + default: + return -1, -1 + } +} + +func allowUnlistedModels(providerType config.ProviderType, pc config.ProviderConfig) bool { + if pc.AllowUnlistedModels != nil { + return *pc.AllowUnlistedModels + } + switch providerType { + case config.ProviderTypeOpenRouter, config.ProviderTypeOpenAICompat, config.ProviderTypeVLLM, + config.ProviderTypeOllama, config.ProviderTypeVenice, config.ProviderTypeNearAI, + config.ProviderTypeSambaNova, config.ProviderTypeZen: + return true + default: + return false + } +} + +func catalogCapabilities(providerType config.ProviderType) protocol.Capabilities { + caps := protocol.Capabilities{ToolCalls: true} + switch providerType { + case config.ProviderTypeAnthropic, config.ProviderTypeBedrock: + caps.NativeWebSearch = true + caps.NativeWebFetch = true + caps.SupportsToolChoice = true + case config.ProviderTypeOpenAI, config.ProviderTypeOpenRouter, config.ProviderTypeGemini, + config.ProviderTypeXAI, config.ProviderTypeOpenAICompat, config.ProviderTypeVLLM, + config.ProviderTypeZen, config.ProviderTypeVenice, config.ProviderTypeNearAI, + config.ProviderTypeSambaNova: + caps.SupportsToolChoice = true + if providerType == config.ProviderTypeOpenAI || providerType == config.ProviderTypeGemini || providerType == config.ProviderTypeXAI { + caps.NativeWebSearch = true + } + case config.ProviderTypeChatGPT: + caps.NativeWebSearch = true + case config.ProviderTypeClaudeBin: + caps.ManagesOwnContext = true + case config.ProviderTypeGrokBin: + caps.NativeWebSearch = true + caps.NativeWebFetch = true + caps.ManagesOwnContext = true + caps.InlineToolLoop = true + case config.ProviderTypeCursorBin: + caps.ManagesOwnContext = true + caps.InlineToolLoop = true + caps.OrderedInlineToolEvents = true + case config.ProviderTypeGeminiCLI: + caps.NativeWebSearch = true + } + return caps +} + +func nonInteractiveAuthReady(providerType string) error { + var binary string + switch config.ProviderType(providerType) { + case config.ProviderTypeChatGPT: + creds, err := credentials.GetChatGPTCredentials() + if err != nil || creds == nil || creds.IsExpired() { + return fmt.Errorf("provider is not authenticated on the gateway; run `term-llm auth login chatgpt` on the gateway host") + } + case config.ProviderTypeCopilot: + creds, err := credentials.GetCopilotCredentials() + if err != nil || creds == nil || creds.IsExpired() { + return fmt.Errorf("provider is not authenticated on the gateway; run `term-llm auth login copilot` on the gateway host") + } + case config.ProviderTypeGeminiCLI: + binary = "gemini" + creds, err := credentials.GetGeminiOAuthCredentials() + if err != nil || creds == nil { + return fmt.Errorf("provider is not authenticated on the gateway; run Gemini CLI login on the gateway host") + } + case config.ProviderTypeClaudeBin: + binary = "claude" + case config.ProviderTypeGrokBin: + binary = "grok" + case config.ProviderTypeCursorBin: + binary = "cursor-agent" + if !llm.CursorBinHasCredentials() { + return fmt.Errorf("provider is not authenticated on the gateway; run `cursor-agent login` on the gateway host") + } + } + if binary != "" { + if _, err := exec.LookPath(binary); err != nil { + return fmt.Errorf("provider executable %q is not available on the gateway host", binary) + } + } + return nil +} + +func providerCredentialReady(cfg *config.Config, key string, providerType config.ProviderType) bool { + pc := cfg.GetProviderConfig(key) + if pc == nil { + return false + } + switch providerType { + case config.ProviderTypeOpenAI, config.ProviderTypeGemini, config.ProviderTypeOpenRouter, + config.ProviderTypeXAI, config.ProviderTypeVenice, config.ProviderTypeNearAI, + config.ProviderTypeSambaNova: + return strings.TrimSpace(pc.ResolvedAPIKey) != "" || strings.TrimSpace(pc.APIKey) != "" + default: + return true + } +} + +func isCLIProvider(providerType config.ProviderType) bool { + switch providerType { + case config.ProviderTypeClaudeBin, config.ProviderTypeGrokBin, config.ProviderTypeCursorBin, config.ProviderTypeGeminiCLI: + return true + default: + return false + } +} + +func filterCatalog(catalog protocol.Catalog, serverPolicy, clientPolicy Policy) protocol.Catalog { + out := catalog + out.Providers = nil + for _, entry := range catalog.Providers { + if !serverPolicy.AllowsProvider(entry.Key, entry.CLI) || !clientPolicy.AllowsProvider(entry.Key, entry.CLI) { + continue + } + copyEntry := entry + copyEntry.Models = nil + for _, model := range entry.Models { + if serverPolicy.Allows(entry.Key, model.ID, entry.CLI) && clientPolicy.Allows(entry.Key, model.ID, entry.CLI) { + copyEntry.Models = append(copyEntry.Models, model) + } + } + if len(copyEntry.Models) > 0 || entry.AllowUnlistedModels { + out.Providers = append(out.Providers, copyEntry) + } + } + out.Features.Search = out.Features.Search && serverPolicy.AllowSearch && clientPolicy.AllowSearch + out.Features.Fetch = out.Features.Fetch && serverPolicy.AllowFetch && clientPolicy.AllowFetch + return out +} + +func catalogEntryAllowsModel(entry protocol.CatalogEntry, provider, model string) bool { + if model == "" { + return false + } + base, _ := llm.BaseModelAndEffortForProvider(provider, model) + for _, candidate := range entry.Models { + if candidate.ID == model || candidate.ID == base { + return true + } + } + return entry.AllowUnlistedModels +} + +func catalogProvider(catalog protocol.Catalog, key string) (protocol.CatalogEntry, bool) { + for _, entry := range catalog.Providers { + if entry.Key == key { + return entry, true + } + } + return protocol.CatalogEntry{}, false +} + +func (s *Server) limitsFor(client Client) *clientLimits { + s.limitsMu.Lock() + defer s.limitsMu.Unlock() + if existing := s.limits[client.ID]; existing != nil { + return existing + } + policy := client.Policy + searchRate := policy.SearchRatePerMinute + if searchRate <= 0 { + searchRate = DefaultSearchRatePerMinute + } + searchBurst := policy.SearchBurst + if searchBurst <= 0 { + searchBurst = DefaultSearchBurst + } + searchConcurrency := policy.MaxConcurrentSearch + if searchConcurrency <= 0 { + searchConcurrency = DefaultMaxConcurrentSearch + } + fetchRate := policy.FetchRatePerMinute + if fetchRate <= 0 { + fetchRate = DefaultFetchRatePerMinute + } + fetchBurst := policy.FetchBurst + if fetchBurst <= 0 { + fetchBurst = DefaultFetchBurst + } + fetchConcurrency := policy.MaxConcurrentFetch + if fetchConcurrency <= 0 { + fetchConcurrency = DefaultMaxConcurrentFetch + } + limits := &clientLimits{ + inference: make(chan struct{}, policy.InferenceConcurrency()), + search: make(chan struct{}, searchConcurrency), + fetch: make(chan struct{}, fetchConcurrency), + searchRPS: rate.NewLimiter(rate.Every(time.Minute/time.Duration(searchRate)), searchBurst), + fetchRPS: rate.NewLimiter(rate.Every(time.Minute/time.Duration(fetchRate)), fetchBurst), + } + s.limits[client.ID] = limits + return limits +} + +func (s *Server) acquireInference(client Client) (func(), bool) { + sem := s.limitsFor(client).inference + select { + case sem <- struct{}{}: + return func() { <-sem }, true + default: + return func() {}, false + } +} + +func (s *Server) acquireTool(client Client, searchRequest bool) (func(), string, bool) { + limits := s.limitsFor(client) + sem, limiter, prefix := limits.fetch, limits.fetchRPS, "fetch" + if searchRequest { + sem, limiter, prefix = limits.search, limits.searchRPS, "search" + } + if !limiter.Allow() { + return func() {}, prefix + "_rate_limited", false + } + select { + case sem <- struct{}{}: + return func() { <-sem }, "", true + default: + return func() {}, prefix + "_concurrency_limited", false + } +} + +func classifyProviderError(err error, providerType string) (int, string) { + if err == nil { + return http.StatusBadGateway, "provider_upstream_failure" + } + if errors.Is(err, context.Canceled) { + return 499, "canceled" + } + var statusErr *providerhttp.StatusError + statusCode := 0 + if errors.As(err, &statusErr) { + statusCode = statusErr.StatusCode + } else if status, ok := err.(interface{ HTTPStatusCode() int }); ok { + statusCode = status.HTTPStatusCode() + } + message := strings.ToLower(err.Error()) + if statusCode == http.StatusUnauthorized || statusCode == http.StatusForbidden { + switch config.ProviderType(providerType) { + case config.ProviderTypeChatGPT, config.ProviderTypeCopilot, config.ProviderTypeGeminiCLI: + return http.StatusUnauthorized, "provider_oauth_unauthenticated" + default: + return http.StatusUnauthorized, "provider_api_key_unauthenticated" + } + } + var rateLimit *llm.RateLimitError + if statusCode == http.StatusTooManyRequests || errors.As(err, &rateLimit) || strings.Contains(message, "rate limit") { + return http.StatusTooManyRequests, "provider_rate_limited" + } + oauthProvider := config.ProviderType(providerType) == config.ProviderTypeChatGPT || config.ProviderType(providerType) == config.ProviderTypeCopilot || config.ProviderType(providerType) == config.ProviderTypeGeminiCLI || config.ProviderType(providerType) == config.ProviderTypeCursorBin + if oauthProvider && (strings.Contains(message, "oauth") || strings.Contains(message, "login") || strings.Contains(message, "not authenticated") || strings.Contains(message, "credential")) { + return http.StatusUnauthorized, "provider_oauth_unauthenticated" + } + if strings.Contains(message, "api key") || strings.Contains(message, "api_key") || (strings.Contains(message, "credential") && (strings.Contains(message, "missing") || strings.Contains(message, "invalid") || strings.Contains(message, "required"))) { + return http.StatusUnauthorized, "provider_api_key_unauthenticated" + } + for _, marker := range []string{"context length", "context_length", "too many tokens", "prompt is too long", "request too large", "maximum context"} { + if strings.Contains(message, marker) { + return http.StatusBadRequest, "provider_context_limit" + } + } + modelTextError := strings.Contains(message, "model") && (strings.Contains(message, "not found") || strings.Contains(message, "invalid") || strings.Contains(message, "unknown") || strings.Contains(message, "unsupported") || strings.Contains(message, "does not exist")) + if statusCode == http.StatusNotFound || modelTextError || ((statusCode == http.StatusBadRequest || statusCode == http.StatusUnprocessableEntity) && strings.Contains(message, "model")) { + return http.StatusBadRequest, "provider_model_invalid" + } + if statusCode >= http.StatusInternalServerError || statusCode == http.StatusRequestTimeout || statusCode == http.StatusBadGateway || statusCode == http.StatusServiceUnavailable || statusCode == http.StatusGatewayTimeout { + return http.StatusBadGateway, "provider_upstream_failure" + } + if statusCode >= http.StatusBadRequest { + return http.StatusBadRequest, "provider_request_invalid" + } + return http.StatusBadGateway, "provider_upstream_failure" +} + +func safeProviderErrorMessage(code, provider string) string { + provider = strings.TrimSpace(provider) + if provider == "" { + provider = "selected" + } + switch code { + case "canceled": + return "gateway request was canceled" + case "provider_api_key_unauthenticated": + return fmt.Sprintf("gateway provider %q rejected its API credential; update the API key on the gateway host", provider) + case "provider_oauth_unauthenticated": + return fmt.Sprintf("gateway provider %q needs OAuth authentication; run the provider login command on the gateway host", provider) + case "provider_rate_limited": + return fmt.Sprintf("gateway provider %q is rate limited; wait and retry or reduce concurrent requests", provider) + case "provider_context_limit": + return fmt.Sprintf("gateway provider %q rejected the request context; shorten or compact the conversation", provider) + case "provider_model_invalid": + return fmt.Sprintf("gateway provider %q rejected the model; choose a model from the gateway catalog", provider) + case "provider_request_invalid": + return fmt.Sprintf("gateway provider %q rejected the request; check model options and retry", provider) + default: + return fmt.Sprintf("gateway provider %q failed upstream; retry, then check gateway-side diagnostics if it persists", provider) + } +} diff --git a/internal/gateway/store.go b/internal/gateway/store.go new file mode 100644 index 000000000..a2f162317 --- /dev/null +++ b/internal/gateway/store.go @@ -0,0 +1,433 @@ +package gateway + +import ( + "crypto/rand" + "crypto/sha256" + "crypto/subtle" + "encoding/hex" + "encoding/json" + "fmt" + "os" + "path/filepath" + "sort" + "strings" + "sync" + "time" +) + +const ( + DefaultMaxConcurrentInference = 2 + DefaultSearchRatePerMinute = 30 + DefaultSearchBurst = 5 + DefaultMaxConcurrentSearch = 2 + DefaultFetchRatePerMinute = 30 + DefaultFetchBurst = 5 + DefaultMaxConcurrentFetch = 2 + DefaultEnrollmentTTL = 15 * time.Minute + MaxEnrollmentTTL = 24 * time.Hour +) + +type Policy struct { + AllowProviders []string `json:"allow_providers,omitempty"` + DenyProviders []string `json:"deny_providers,omitempty"` + AllowModels []string `json:"allow_models,omitempty"` + DenyModels []string `json:"deny_models,omitempty"` + AllowCLI bool `json:"allow_cli,omitempty"` + AllowSearch bool `json:"allow_search"` + AllowFetch bool `json:"allow_fetch"` + MaxConcurrentInference int `json:"max_concurrent_inference,omitempty"` + SearchRatePerMinute int `json:"search_rate_per_minute,omitempty"` + SearchBurst int `json:"search_burst,omitempty"` + MaxConcurrentSearch int `json:"max_concurrent_search,omitempty"` + FetchRatePerMinute int `json:"fetch_rate_per_minute,omitempty"` + FetchBurst int `json:"fetch_burst,omitempty"` + MaxConcurrentFetch int `json:"max_concurrent_fetch,omitempty"` +} + +type Client struct { + ID string `json:"id"` + Name string `json:"name"` + TokenHash string `json:"token_hash"` + CreatedAt time.Time `json:"created_at"` + RevokedAt time.Time `json:"revoked_at,omitempty"` + Policy Policy `json:"policy"` +} + +type Enrollment struct { + Name string `json:"name"` + TokenHash string `json:"token_hash"` + CreatedAt time.Time `json:"created_at"` + ExpiresAt time.Time `json:"expires_at"` + UsedAt time.Time `json:"used_at,omitempty"` + Policy Policy `json:"policy"` +} + +type ClientStore struct { + path string + enrollmentPath string + mu sync.RWMutex + clients []Client + enrollments []Enrollment +} + +func OpenClientStore(path string) (*ClientStore, error) { + store := &ClientStore{path: path, enrollmentPath: strings.TrimSuffix(path, filepath.Ext(path)) + ".enrollments.json"} + if err := readJSONFile(path, &store.clients); err != nil { + return nil, fmt.Errorf("read gateway clients: %w", err) + } + if err := readJSONFile(store.enrollmentPath, &store.enrollments); err != nil { + return nil, fmt.Errorf("read gateway enrollments: %w", err) + } + return store, nil +} + +func readJSONFile(path string, target any) error { + data, err := os.ReadFile(path) + if err != nil { + if os.IsNotExist(err) { + return nil + } + return err + } + if err := json.Unmarshal(data, target); err != nil { + return fmt.Errorf("decode %s: %w", filepath.Base(path), err) + } + return nil +} + +func (s *ClientStore) Add(name string, policy Policy) (Client, string, error) { + name = strings.TrimSpace(name) + if name == "" { + return Client{}, "", fmt.Errorf("client name is required") + } + s.mu.Lock() + defer s.mu.Unlock() + unlock, err := lockGatewayClientStore(s.path + ".lock") + if err != nil { + return Client{}, "", err + } + defer func() { _ = unlock() }() + var clients []Client + if err := readJSONFile(s.path, &clients); err != nil { + return Client{}, "", fmt.Errorf("reload gateway clients: %w", err) + } + s.clients = clients + return s.addLocked(name, policy) +} + +func (s *ClientStore) addLocked(name string, policy Policy) (Client, string, error) { + if s.activeClientNameLocked(name) { + return Client{}, "", fmt.Errorf("active gateway client %q already exists; revoke it before rotating credentials", name) + } + token, err := randomSecret("tlg1", 32) + if err != nil { + return Client{}, "", err + } + clientID, err := randomSecret("client", 12) + if err != nil { + return Client{}, "", err + } + client := Client{ID: clientID, Name: name, TokenHash: hashToken(token), CreatedAt: time.Now().UTC(), Policy: policy} + s.clients = append(s.clients, client) + if err := s.saveClientsLocked(); err != nil { + s.clients = s.clients[:len(s.clients)-1] + return Client{}, "", err + } + return client, token, nil +} + +// CreateEnrollment creates a persisted, single-use bootstrap token bound to a +// client name, restricted policy, and short expiry. +func (s *ClientStore) CreateEnrollment(name string, policy Policy, ttl time.Duration) (Enrollment, string, error) { + name = strings.TrimSpace(name) + if name == "" { + return Enrollment{}, "", fmt.Errorf("enrollment client name is required") + } + if len(policy.AllowProviders) == 0 && len(policy.AllowModels) == 0 { + return Enrollment{}, "", fmt.Errorf("enrollment requires --allow-provider or --allow-model") + } + if ttl <= 0 { + ttl = DefaultEnrollmentTTL + } + if ttl > MaxEnrollmentTTL { + return Enrollment{}, "", fmt.Errorf("enrollment TTL must not exceed %s", MaxEnrollmentTTL) + } + token, err := randomSecret("tlge1", 32) + if err != nil { + return Enrollment{}, "", err + } + now := time.Now().UTC() + enrollment := Enrollment{Name: name, TokenHash: hashToken(token), CreatedAt: now, ExpiresAt: now.Add(ttl), Policy: policy} + s.mu.Lock() + defer s.mu.Unlock() + unlock, err := lockGatewayClientStore(s.path + ".lock") + if err != nil { + return Enrollment{}, "", err + } + defer func() { _ = unlock() }() + var clients []Client + if err := readJSONFile(s.path, &clients); err != nil { + return Enrollment{}, "", fmt.Errorf("reload gateway clients: %w", err) + } + s.clients = clients + if s.activeClientNameLocked(name) { + return Enrollment{}, "", fmt.Errorf("active gateway client %q already exists; revoke it before rotating credentials", name) + } + var enrollments []Enrollment + if err := readJSONFile(s.enrollmentPath, &enrollments); err != nil { + return Enrollment{}, "", fmt.Errorf("reload gateway enrollments: %w", err) + } + s.enrollments = enrollments + s.enrollments = append(s.enrollments, enrollment) + if err := s.saveEnrollmentsLocked(); err != nil { + s.enrollments = s.enrollments[:len(s.enrollments)-1] + return Enrollment{}, "", err + } + return enrollment, token, nil +} + +// ConsumeEnrollment atomically marks a valid token used before minting its +// per-client credential. A token can never create two clients, including across +// process restarts. +func (s *ClientStore) ConsumeEnrollment(token, requestedName string) (Client, string, error) { + hash := hashToken(strings.TrimSpace(token)) + requestedName = strings.TrimSpace(requestedName) + now := time.Now().UTC() + s.mu.Lock() + defer s.mu.Unlock() + unlock, err := lockGatewayClientStore(s.path + ".lock") + if err != nil { + return Client{}, "", err + } + defer func() { _ = unlock() }() + var enrollments []Enrollment + if err := readJSONFile(s.enrollmentPath, &enrollments); err != nil { + return Client{}, "", fmt.Errorf("reload gateway enrollments: %w", err) + } + s.enrollments = enrollments + var clients []Client + if err := readJSONFile(s.path, &clients); err != nil { + return Client{}, "", fmt.Errorf("reload gateway clients: %w", err) + } + s.clients = clients + for i := range s.enrollments { + enrollment := &s.enrollments[i] + if subtle.ConstantTimeCompare([]byte(hash), []byte(enrollment.TokenHash)) != 1 { + continue + } + if !enrollment.UsedAt.IsZero() { + return Client{}, "", fmt.Errorf("enrollment token has already been used") + } + if !now.Before(enrollment.ExpiresAt) { + return Client{}, "", fmt.Errorf("enrollment token has expired") + } + if requestedName != "" && requestedName != enrollment.Name { + return Client{}, "", fmt.Errorf("enrollment token is bound to client %q", enrollment.Name) + } + if s.activeClientNameLocked(enrollment.Name) { + return Client{}, "", fmt.Errorf("active gateway client %q already exists; revoke it before rotating credentials", enrollment.Name) + } + enrollment.UsedAt = now + if err := s.saveEnrollmentsLocked(); err != nil { + enrollment.UsedAt = time.Time{} + return Client{}, "", err + } + client, clientToken, err := s.addLocked(enrollment.Name, enrollment.Policy) + if err != nil { + enrollment.UsedAt = time.Time{} + if rollbackErr := s.saveEnrollmentsLocked(); rollbackErr != nil { + return Client{}, "", fmt.Errorf("create enrolled client: %v; restore enrollment token: %w", err, rollbackErr) + } + return Client{}, "", err + } + return client, clientToken, nil + } + return Client{}, "", fmt.Errorf("invalid enrollment token") +} + +// Authenticate reloads the durable client file before every decision. Management +// commands intentionally open a separate ClientStore, so request-time reloads make +// additions and revocations visible to a running gateway as soon as the atomic +// rename reaches the filesystem; there is no polling interval or process restart. +func (s *ClientStore) Authenticate(token string) (Client, bool) { + hash := hashToken(strings.TrimSpace(token)) + s.mu.Lock() + defer s.mu.Unlock() + var clients []Client + if err := readJSONFile(s.path, &clients); err != nil { + // Authentication fails closed if the durable store cannot be read. Atomic + // writers ensure a valid old or new file is normally observed. + return Client{}, false + } + s.clients = clients + for _, client := range s.clients { + if !client.RevokedAt.IsZero() { + continue + } + if subtle.ConstantTimeCompare([]byte(hash), []byte(client.TokenHash)) == 1 { + return client, true + } + } + return Client{}, false +} + +func (s *ClientStore) List() []Client { + s.mu.RLock() + defer s.mu.RUnlock() + out := append([]Client(nil), s.clients...) + sort.Slice(out, func(i, j int) bool { return out[i].CreatedAt.Before(out[j].CreatedAt) }) + return out +} + +func (s *ClientStore) Revoke(idOrName string) error { + idOrName = strings.TrimSpace(idOrName) + if idOrName == "" { + return fmt.Errorf("gateway client ID or name is required") + } + s.mu.Lock() + defer s.mu.Unlock() + unlock, err := lockGatewayClientStore(s.path + ".lock") + if err != nil { + return err + } + defer func() { _ = unlock() }() + var clients []Client + if err := readJSONFile(s.path, &clients); err != nil { + return fmt.Errorf("reload gateway clients: %w", err) + } + s.clients = clients + + // IDs take precedence over names. Name-based revocation covers every active + // legacy match, while new writes enforce one active client per name. + for i := range s.clients { + if s.clients[i].ID == idOrName { + if !s.clients[i].RevokedAt.IsZero() { + return nil + } + previous := s.clients[i].RevokedAt + s.clients[i].RevokedAt = time.Now().UTC() + if err := s.saveClientsLocked(); err != nil { + s.clients[i].RevokedAt = previous + return err + } + return nil + } + } + + matched := false + changed := false + previous := append([]Client(nil), s.clients...) + now := time.Now().UTC() + for i := range s.clients { + if s.clients[i].Name != idOrName { + continue + } + matched = true + if s.clients[i].RevokedAt.IsZero() { + s.clients[i].RevokedAt = now + changed = true + } + } + if !matched { + return fmt.Errorf("gateway client %q not found", idOrName) + } + if !changed { + return nil + } + if err := s.saveClientsLocked(); err != nil { + s.clients = previous + return err + } + return nil +} + +func (s *ClientStore) activeClientNameLocked(name string) bool { + for _, client := range s.clients { + if client.Name == name && client.RevokedAt.IsZero() { + return true + } + } + return false +} + +func (s *ClientStore) saveClientsLocked() error { + return writeSecureJSON(s.path, s.clients, "gateway clients") +} + +func (s *ClientStore) saveEnrollmentsLocked() error { + return writeSecureJSON(s.enrollmentPath, s.enrollments, "gateway enrollments") +} + +func writeSecureJSON(path string, value any, label string) error { + if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil { + return fmt.Errorf("create gateway state directory: %w", err) + } + data, err := json.MarshalIndent(value, "", " ") + if err != nil { + return err + } + tmp := path + ".tmp" + if err := os.WriteFile(tmp, data, 0o600); err != nil { + return fmt.Errorf("write %s: %w", label, err) + } + if err := os.Rename(tmp, path); err != nil { + return fmt.Errorf("replace %s: %w", label, err) + } + return nil +} + +func hashToken(token string) string { + sum := sha256.Sum256([]byte(token)) + return hex.EncodeToString(sum[:]) +} + +func randomSecret(prefix string, size int) (string, error) { + raw := make([]byte, size) + if _, err := rand.Read(raw); err != nil { + return "", fmt.Errorf("generate secure gateway identifier: %w", err) + } + return prefix + "_" + hex.EncodeToString(raw), nil +} + +func (p Policy) Allows(provider, model string, cli bool) bool { + if !p.AllowsProvider(provider, cli) { + return false + } + if matchesPolicy(p.DenyModels, provider+":"+model) || matchesPolicy(p.DenyModels, model) { + return false + } + if len(p.AllowModels) > 0 && !matchesPolicy(p.AllowModels, provider+":"+model) && !matchesPolicy(p.AllowModels, model) { + return false + } + return true +} + +func (p Policy) AllowsProvider(provider string, cli bool) bool { + if cli && !p.AllowCLI { + return false + } + if matchesPolicy(p.DenyProviders, provider) { + return false + } + return len(p.AllowProviders) == 0 || matchesPolicy(p.AllowProviders, provider) +} + +func (p Policy) InferenceConcurrency() int { + if p.MaxConcurrentInference > 0 { + return p.MaxConcurrentInference + } + return DefaultMaxConcurrentInference +} + +func matchesPolicy(patterns []string, value string) bool { + for _, pattern := range patterns { + pattern = strings.TrimSpace(pattern) + if pattern == "*" || pattern == value { + return true + } + if strings.HasSuffix(pattern, "*") && strings.HasPrefix(value, strings.TrimSuffix(pattern, "*")) { + return true + } + } + return false +} diff --git a/internal/gateway/store_lock_unix.go b/internal/gateway/store_lock_unix.go new file mode 100644 index 000000000..795ca3f27 --- /dev/null +++ b/internal/gateway/store_lock_unix.go @@ -0,0 +1,35 @@ +//go:build !windows + +package gateway + +import ( + "fmt" + "os" + "path/filepath" + "syscall" +) + +func lockGatewayClientStore(lockPath string) (func() error, error) { + if err := os.MkdirAll(filepath.Dir(lockPath), 0o700); err != nil { + return nil, fmt.Errorf("create gateway state directory: %w", err) + } + file, err := os.OpenFile(lockPath, os.O_CREATE|os.O_RDWR, 0o600) + if err != nil { + return nil, fmt.Errorf("open gateway client lock: %w", err) + } + if err := syscall.Flock(int(file.Fd()), syscall.LOCK_EX); err != nil { + _ = file.Close() + return nil, fmt.Errorf("acquire gateway client lock: %w", err) + } + return func() error { + unlockErr := syscall.Flock(int(file.Fd()), syscall.LOCK_UN) + closeErr := file.Close() + if unlockErr != nil { + return fmt.Errorf("unlock gateway client store: %w", unlockErr) + } + if closeErr != nil { + return fmt.Errorf("close gateway client lock: %w", closeErr) + } + return nil + }, nil +} diff --git a/internal/gateway/store_lock_windows.go b/internal/gateway/store_lock_windows.go new file mode 100644 index 000000000..5ffa52445 --- /dev/null +++ b/internal/gateway/store_lock_windows.go @@ -0,0 +1,37 @@ +//go:build windows + +package gateway + +import ( + "fmt" + "os" + "path/filepath" + + "golang.org/x/sys/windows" +) + +func lockGatewayClientStore(lockPath string) (func() error, error) { + if err := os.MkdirAll(filepath.Dir(lockPath), 0o700); err != nil { + return nil, fmt.Errorf("create gateway state directory: %w", err) + } + file, err := os.OpenFile(lockPath, os.O_CREATE|os.O_RDWR, 0o600) + if err != nil { + return nil, fmt.Errorf("open gateway client lock: %w", err) + } + var overlapped windows.Overlapped + if err := windows.LockFileEx(windows.Handle(file.Fd()), windows.LOCKFILE_EXCLUSIVE_LOCK, 0, 1, 0, &overlapped); err != nil { + _ = file.Close() + return nil, fmt.Errorf("acquire gateway client lock: %w", err) + } + return func() error { + unlockErr := windows.UnlockFileEx(windows.Handle(file.Fd()), 0, 1, 0, &overlapped) + closeErr := file.Close() + if unlockErr != nil { + return fmt.Errorf("unlock gateway client store: %w", unlockErr) + } + if closeErr != nil { + return fmt.Errorf("close gateway client lock: %w", closeErr) + } + return nil + }, nil +} diff --git a/internal/gateway/testproviders_test.go b/internal/gateway/testproviders_test.go new file mode 100644 index 000000000..3b2ed2a22 --- /dev/null +++ b/internal/gateway/testproviders_test.go @@ -0,0 +1,89 @@ +package gateway + +import ( + "context" + "fmt" + "io" + "sync" + + "github.com/samsaffron/term-llm/internal/llm" +) + +type oneEventStream struct { + ctx context.Context + event llm.Event + sent bool + err error +} + +func (s *oneEventStream) Recv() (llm.Event, error) { + if !s.sent { + s.sent = true + return s.event, nil + } + if s.err != nil { + return llm.Event{}, s.err + } + return llm.Event{}, io.EOF +} +func (*oneEventStream) Close() error { return nil } + +type statefulProvider struct { + mu sync.Mutex + imported string + next int +} + +func (*statefulProvider) Name() string { return "stateful" } +func (*statefulProvider) Credential() string { return "mock" } +func (*statefulProvider) Capabilities() llm.Capabilities { return llm.Capabilities{} } +func (p *statefulProvider) ImportProviderState(data []byte) error { + p.mu.Lock() + defer p.mu.Unlock() + p.imported = string(data) + return nil +} +func (p *statefulProvider) ExportProviderState() ([]byte, bool) { + p.mu.Lock() + defer p.mu.Unlock() + return []byte(fmt.Sprintf("state-%d", p.next)), true +} +func (p *statefulProvider) Stream(ctx context.Context, _ llm.Request) (llm.Stream, error) { + p.mu.Lock() + p.next++ + p.mu.Unlock() + return &oneEventStream{ctx: ctx, event: llm.Event{Type: llm.EventTextDelta, Text: "ok"}}, nil +} + +type streamDeathProvider struct{} + +func (*streamDeathProvider) Name() string { return "dead" } +func (*streamDeathProvider) Credential() string { return "mock" } +func (*streamDeathProvider) Capabilities() llm.Capabilities { + return llm.Capabilities{} +} +func (*streamDeathProvider) Stream(ctx context.Context, _ llm.Request) (llm.Stream, error) { + return &oneEventStream{ctx: ctx, event: llm.Event{Type: llm.EventTextDelta, Text: "partial"}, err: fmt.Errorf("secret upstream transport detail")}, nil +} + +type blockingProvider struct{ canceled chan struct{} } + +func (*blockingProvider) Name() string { return "blocking" } +func (*blockingProvider) Credential() string { return "mock" } +func (*blockingProvider) Capabilities() llm.Capabilities { return llm.Capabilities{} } +func (p *blockingProvider) Stream(ctx context.Context, _ llm.Request) (llm.Stream, error) { + return &blockingStream{ctx: ctx, canceled: p.canceled}, nil +} + +type blockingStream struct { + ctx context.Context + canceled chan struct{} + once sync.Once +} + +func (s *blockingStream) Recv() (llm.Event, error) { + <-s.ctx.Done() + s.once.Do(func() { close(s.canceled) }) + return llm.Event{}, s.ctx.Err() +} +func (*blockingStream) Close() error { return nil } diff --git a/internal/gateway/usage.go b/internal/gateway/usage.go new file mode 100644 index 000000000..0ba391f8b --- /dev/null +++ b/internal/gateway/usage.go @@ -0,0 +1,96 @@ +package gateway + +import ( + "bufio" + "encoding/json" + "fmt" + "os" + "path/filepath" + "sync" + "time" + + "github.com/samsaffron/term-llm/internal/llm" + usagepkg "github.com/samsaffron/term-llm/internal/usage" +) + +type UsageRecord struct { + StartedAt time.Time `json:"started_at"` + CompletedAt time.Time `json:"completed_at"` + ClientID string `json:"client_id"` + ClientName string `json:"client_name"` + ProviderKey string `json:"provider_key"` + Model string `json:"model"` + RequestID string `json:"request_id"` + SessionID string `json:"session_id,omitempty"` + InputTokens int `json:"input_tokens"` + OutputTokens int `json:"output_tokens"` + CachedInputTokens int `json:"cached_input_tokens"` + CacheWriteTokens int `json:"cache_write_tokens"` + ReasoningTokens int `json:"reasoning_tokens"` + CostUSD *float64 `json:"cost_usd,omitempty"` + ErrorCode string `json:"error_code,omitempty"` +} + +type UsageRecorder interface{ Record(UsageRecord) error } + +type JSONLUsageRecorder struct { + Path string + mu sync.Mutex +} + +func (r *JSONLUsageRecorder) Record(record UsageRecord) error { + r.mu.Lock() + defer r.mu.Unlock() + record.StartedAt = record.StartedAt.UTC() + record.CompletedAt = record.CompletedAt.UTC() + if err := os.MkdirAll(filepath.Dir(r.Path), 0o700); err != nil { + return err + } + file, err := os.OpenFile(r.Path, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0o600) + if err != nil { + return fmt.Errorf("open gateway usage log: %w", err) + } + defer file.Close() + return json.NewEncoder(file).Encode(record) +} + +func ReadUsageRecords(path string) ([]UsageRecord, error) { + file, err := os.Open(path) + if err != nil { + if os.IsNotExist(err) { + return nil, nil + } + return nil, fmt.Errorf("open gateway usage log: %w", err) + } + defer file.Close() + var records []UsageRecord + scanner := bufio.NewScanner(file) + scanner.Buffer(make([]byte, 64<<10), 4<<20) + for scanner.Scan() { + var record UsageRecord + if err := json.Unmarshal(scanner.Bytes(), &record); err != nil { + return nil, fmt.Errorf("decode gateway usage record: %w", err) + } + records = append(records, record) + } + if err := scanner.Err(); err != nil { + return nil, fmt.Errorf("read gateway usage log: %w", err) + } + return records, nil +} + +func estimateUsageCost(provider, model string, use llm.Usage) *float64 { + if inputPrice, outputPrice, ok := llm.PricingForProviderModel(provider, model); ok { + cost := float64(use.InputTokens+use.CachedInputTokens+use.CacheWriteTokens)*inputPrice/1_000_000 + float64(use.OutputTokens)*outputPrice/1_000_000 + return &cost + } + fetcher := usagepkg.NewPricingFetcher() + cost, err := fetcher.CalculateCostLocal(usagepkg.UsageEntry{ + Model: model, InputTokens: use.InputTokens, OutputTokens: use.OutputTokens, + CacheReadTokens: use.CachedInputTokens, CacheWriteTokens: use.CacheWriteTokens, + }) + if err != nil { + return nil + } + return &cost +} diff --git a/internal/llm/engine.go b/internal/llm/engine.go index 9be7bdcc9..216db7506 100644 --- a/internal/llm/engine.go +++ b/internal/llm/engine.go @@ -254,6 +254,11 @@ func NewEngine(provider Provider, tools *ToolRegistry) *Engine { return e } +func (e *Engine) providerHandlesRetries() bool { + handler, ok := e.provider.(interface{ GatewayHandlesRetries() bool }) + return ok && handler.GatewayHandlesRetries() +} + // TriggerChaosFailure arms a one-shot synthetic replayable stream failure. It is // intentionally tiny and transport-shaped so UI/debug flows exercise the same // recovery paths as a prematurely closed SSE/WebSocket stream. @@ -1581,7 +1586,7 @@ func (e *Engine) Stream(ctx context.Context, req Request) (Stream, error) { stream := newEventStream(ctx, func(ctx context.Context, send eventSender) error { return e.runLoop(ctx, req, send) }) - stream = wrapLoggingStream(stream, e.provider.Name(), req.Model) + stream = wrapLoggingStream(stream, e.provider.Name(), req.Model, providerUsageTrackedExternallyBy(e.provider)) stream = e.wrapDebugLoggingStream(stream) // Wrap with per-turn cleanup for providers that materialize temporary @@ -1604,7 +1609,7 @@ func (e *Engine) Stream(ctx context.Context, req Request) (Stream, error) { stream := newEventStream(ctx, func(ctx context.Context, send eventSender) error { return e.runSimpleScratchpad(ctx, req, send) }) - stream = wrapLoggingStream(stream, e.provider.Name(), req.Model) + stream = wrapLoggingStream(stream, e.provider.Name(), req.Model, providerUsageTrackedExternallyBy(e.provider)) stream = e.wrapDebugLoggingStream(stream) return stream, nil } @@ -1945,7 +1950,7 @@ func (e *Engine) runSimpleScratchpad(ctx context.Context, req Request, send even return err } priorErr = failed - if retry >= defaultUncommittedStreamMaxRetries || !isUncommittedReplayableStreamError(failed) { + if retry >= defaultUncommittedStreamMaxRetries || e.providerHandlesRetries() || !isUncommittedReplayableStreamError(failed) { return failed } attempt := retry + 1 @@ -2701,7 +2706,7 @@ turnLoop: return true, nil } retryUncommittedAttempt := func(cause error) (bool, error) { - if cause == nil || errors.Is(cause, context.Canceled) || errors.Is(cause, context.DeadlineExceeded) || !isUncommittedReplayableStreamError(cause) { + if cause == nil || e.providerHandlesRetries() || errors.Is(cause, context.Canceled) || errors.Is(cause, context.DeadlineExceeded) || !isUncommittedReplayableStreamError(cause) { return false, nil } if recoveredToolWork || len(toolCalls) > 0 || syncToolsExecuted || scratchpadCommitted { @@ -4147,7 +4152,7 @@ func (s *loggingStream) flushLocked() { } s.logged = true _ = s.logger.Log(usage.LogEntry{ - Timestamp: time.Now(), + Timestamp: time.Now().UTC(), Model: s.model, Provider: s.providerName, InputTokens: s.totalInput, @@ -4158,8 +4163,15 @@ func (s *loggingStream) flushLocked() { }) } -// wrapLoggingStream wraps a stream with usage logging -func wrapLoggingStream(inner Stream, providerName, model string) Stream { +func providerUsageTrackedExternallyBy(provider Provider) string { + if tracking, ok := provider.(interface{ UsageTrackedExternallyBy() string }); ok { + return tracking.UsageTrackedExternallyBy() + } + return usage.GetTrackedExternallyBy(provider.Name()) +} + +// wrapLoggingStream wraps a stream with usage logging. +func wrapLoggingStream(inner Stream, providerName, model, trackedExternal string) Stream { // If model is empty, use providerName as the model identifier // This helps identify what was used when providers auto-select models if model == "" { @@ -4170,7 +4182,7 @@ func wrapLoggingStream(inner Stream, providerName, model string) Stream { logger: usage.DefaultLogger(), providerName: providerName, model: model, - trackedExternal: usage.GetTrackedExternallyBy(providerName), + trackedExternal: trackedExternal, } } diff --git a/internal/llm/factory.go b/internal/llm/factory.go index 5a5b08755..4f2efe64a 100644 --- a/internal/llm/factory.go +++ b/internal/llm/factory.go @@ -1,6 +1,7 @@ package llm import ( + "context" "fmt" "os" "strings" @@ -30,7 +31,7 @@ func ParseProviderModel(s string, cfg *config.Config) (string, string, error) { // Check if provider is configured or is a built-in type if cfg != nil { - if _, ok := cfg.Providers[provider]; ok { + if _, ok := cfg.Providers[provider]; ok && (!cfg.Gateway.Enabled() || cfg.IsLocalProvider(provider)) { return provider, model, nil } if len(parts) == 1 { @@ -47,6 +48,22 @@ func ParseProviderModel(s string, cfg *config.Config) (string, string, error) { } } + // A gateway catalog extends ordinary provider:model syntax without adding a + // second addressing scheme. Explicit local routing still won above. + if cfg != nil && cfg.Gateway.Enabled() && !cfg.IsLocalProvider(provider) { + timeout := gatewayDuration(cfg.Gateway.ConnectTimeout, config.DefaultGatewayConnectTimeout) + gatewayDuration(cfg.Gateway.ResponseTimeout, config.DefaultGatewayResponseTimeout) + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + found, err := GatewayCatalogHasProvider(ctx, cfg, provider) + if err != nil { + return "", "", fmt.Errorf("gateway unavailable while resolving provider %q; check gateway URL/network/token: %w", provider, err) + } + if found { + return provider, model, nil + } + return "", "", fmt.Errorf("provider %q is not available through the configured gateway; add it on the gateway or list it in gateway.local_providers", provider) + } + // Also accept built-in provider type names for _, name := range GetBuiltInProviderNames() { if provider == name { @@ -64,6 +81,9 @@ func NewProvider(cfg *config.Config) (Provider, error) { if err != nil { return nil, err } + if _, ok := provider.(interface{ GatewayHandlesRetries() bool }); ok { + return provider, nil + } // Wrap with retry logic (enabled by default) return WrapWithRetry(provider, DefaultRetryConfig()), nil } @@ -73,6 +93,9 @@ func NewProvider(cfg *config.Config) (Provider, error) { // If the provider is a built-in type but not explicitly configured, // it will be created with default settings. func NewProviderByName(cfg *config.Config, name string, model string) (Provider, error) { + if provider, routed, err := newGatewayProviderIfRouted(cfg, name, model); routed || err != nil { + return provider, err + } // Handle hidden debug provider first if name == "debug" { provider := NewDebugProvider(model) @@ -205,7 +228,27 @@ func NewProviderByName(cfg *config.Config, name string, model string) (Provider, return WrapWithRetry(provider, DefaultRetryConfig()), nil } -// NewFastProvider creates a lightweight provider instance for the specified provider key. +func newGatewayProviderIfRouted(cfg *config.Config, name, model string) (Provider, bool, error) { + if cfg == nil || !cfg.Gateway.Enabled() || cfg.IsLocalProvider(name) { + return nil, false, nil + } + timeout := gatewayDuration(cfg.Gateway.ConnectTimeout, config.DefaultGatewayConnectTimeout) + gatewayDuration(cfg.Gateway.ResponseTimeout, config.DefaultGatewayResponseTimeout) + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + found, err := GatewayCatalogHasProvider(ctx, cfg, name) + if err != nil { + return nil, true, fmt.Errorf("gateway unavailable for provider %q; check gateway URL/network/token: %w", name, err) + } + if !found { + return nil, true, fmt.Errorf("provider %q is not available through the configured gateway; add it on the gateway or list it in gateway.local_providers", name) + } + provider, err := NewGatewayProvider(cfg, name, model) + if err != nil { + return nil, true, fmt.Errorf("connect to gateway provider %q: %w", name, err) + } + return provider, true, nil +} + // Resolution order: // 1. providers..fast_provider + fast_model // 2. providers..fast_model on the same provider key @@ -240,9 +283,19 @@ func NewFastProvider(cfg *config.Config, name string) (Provider, error) { // newProviderInternal creates the underlying provider without retry wrapper. func newProviderInternal(cfg *config.Config) (Provider, error) { - // Handle hidden debug provider first + if cfg == nil { + return nil, fmt.Errorf("config is nil") + } + if provider, routed, err := newGatewayProviderIfRouted(cfg, cfg.DefaultProvider, ""); routed || err != nil { + return provider, err + } + // Handle hidden debug provider first. if cfg.DefaultProvider == "debug" { - return NewDebugProvider(""), nil + model := "" + if pc := cfg.GetProviderConfig("debug"); pc != nil { + model = pc.Model + } + return NewDebugProvider(model), nil } providerCfg, ok := cfg.Providers[cfg.DefaultProvider] diff --git a/internal/llm/gateway_catalog_test.go b/internal/llm/gateway_catalog_test.go new file mode 100644 index 000000000..f6b1f12bc --- /dev/null +++ b/internal/llm/gateway_catalog_test.go @@ -0,0 +1,165 @@ +package llm + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "os" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/samsaffron/term-llm/internal/config" + "github.com/samsaffron/term-llm/internal/gateway/protocol" +) + +func resetGatewayCatalogProcessCacheForTest() { + gatewayCatalogProcessCache.Lock() + gatewayCatalogProcessCache.entries = make(map[string]gatewayCatalogCache) + gatewayCatalogProcessCache.Unlock() +} + +func TestGatewayCatalogCacheETagAndStaleOnError(t *testing.T) { + resetGatewayCatalogProcessCacheForTest() + t.Cleanup(resetGatewayCatalogProcessCacheForTest) + t.Setenv("XDG_CACHE_HOME", t.TempDir()) + requests := 0 + fail := false + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requests++ + if r.Header.Get("Authorization") != "Bearer token" { + t.Errorf("authorization = %q", r.Header.Get("Authorization")) + } + if fail { + w.WriteHeader(http.StatusBadGateway) + return + } + if r.Header.Get("If-None-Match") == `"catalog-v1"` { + w.WriteHeader(http.StatusNotModified) + return + } + w.Header().Set("ETag", `"catalog-v1"`) + _ = json.NewEncoder(w).Encode(protocol.Catalog{Version: 1, GeneratedAt: time.Now(), Providers: []protocol.CatalogEntry{{Key: "remote", Models: []protocol.Model{{ID: "m"}}}}}) + })) + defer server.Close() + cfg := &config.Config{Gateway: config.GatewayConfig{URL: server.URL, Token: "token", CatalogTTL: "1ms", ConnectTimeout: "1s", ResponseTimeout: "1s"}} + client, err := NewGatewayProviderForCatalog(cfg) + if err != nil { + t.Fatal(err) + } + catalog, err := client.loadCatalog(t.Context()) + if err != nil || len(catalog.Providers) != 1 { + t.Fatalf("first catalog = %+v, %v", catalog, err) + } + time.Sleep(2 * time.Millisecond) + if _, err := client.loadCatalog(t.Context()); err != nil { + t.Fatalf("ETag refresh: %v", err) + } + if requests < 2 { + t.Fatalf("requests = %d, want ETag revalidation", requests) + } + cachePath := gatewayCatalogCachePath(server.URL, "token") + if cachePath == gatewayCatalogCachePath(server.URL, "other-client-token") { + t.Fatal("gateway catalog cache is not client-scoped") + } + cache, err := readGatewayCatalogCache(cachePath) + if err != nil || cache.ETag != `"catalog-v1"` { + t.Fatalf("cache = %+v, %v", cache, err) + } + cache.FetchedAt = time.Now().Add(-time.Hour) + if err := writeGatewayCatalogCache(cachePath, cache); err != nil { + t.Fatal(err) + } + gatewayCatalogProcessCache.Lock() + gatewayCatalogProcessCache.entries[gatewayCatalogIdentity(server.URL, "token")] = cache + gatewayCatalogProcessCache.Unlock() + fail = true + stale, err := client.loadCatalog(t.Context()) + if err != nil || len(stale.Providers) != 1 || stale.Providers[0].Key != "remote" { + t.Fatalf("stale fallback = %+v, %v", stale, err) + } + if _, err := os.Stat(cachePath); err != nil { + t.Fatal(err) + } +} + +func TestGatewayCatalogProcessMemoizationSingleflightAndCacheOnlyCompletion(t *testing.T) { + resetGatewayCatalogProcessCacheForTest() + t.Cleanup(resetGatewayCatalogProcessCacheForTest) + t.Setenv("XDG_CACHE_HOME", t.TempDir()) + var requests atomic.Int32 + started := make(chan struct{}) + release := make(chan struct{}) + var startedOnce sync.Once + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requests.Add(1) + startedOnce.Do(func() { close(started) }) + <-release + _ = json.NewEncoder(w).Encode(protocol.Catalog{Version: 1, Providers: []protocol.CatalogEntry{{Key: "remote", Models: []protocol.Model{{ID: "model-a"}}}}}) + })) + defer server.Close() + cfg := &config.Config{Gateway: config.GatewayConfig{URL: server.URL, Token: "token", CatalogTTL: "1m", ConnectTimeout: "1s", ResponseTimeout: "1s"}} + provider, err := NewGatewayProviderForCatalog(cfg) + if err != nil { + t.Fatal(err) + } + const callers = 8 + errs := make(chan error, callers) + for range callers { + go func() { + _, err := provider.loadCatalog(t.Context()) + errs <- err + }() + } + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("catalog request did not start") + } + time.Sleep(20 * time.Millisecond) + if got := requests.Load(); got != 1 { + t.Fatalf("concurrent catalog requests = %d, want one", got) + } + close(release) + for range callers { + if err := <-errs; err != nil { + t.Fatal(err) + } + } + before := requests.Load() + completions := GetProviderCompletions("remote:", false, cfg) + if requests.Load() != before { + t.Fatal("shell completion performed gateway network I/O") + } + if len(completions) != 1 || completions[0] != "remote:model-a" { + t.Fatalf("cache-only completions = %v", completions) + } + providers := GetProviderCompletions("", false, cfg) + for _, provider := range providers { + if provider == "openai" || provider == "anthropic" { + t.Fatalf("gateway completion advertised unconfigured built-in %q: %v", provider, providers) + } + } +} + +func TestGatewayCompletionWithHungGatewayIsImmediateAndNetworkFree(t *testing.T) { + resetGatewayCatalogProcessCacheForTest() + t.Cleanup(resetGatewayCatalogProcessCacheForTest) + t.Setenv("XDG_CACHE_HOME", t.TempDir()) + var requests atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + requests.Add(1) + select {} + })) + defer server.Close() + cfg := &config.Config{Gateway: config.GatewayConfig{URL: server.URL, Token: "token", ConnectTimeout: "10s", ResponseTimeout: "30s"}} + started := time.Now() + _ = GetProviderCompletions("remote", false, cfg) + if elapsed := time.Since(started); elapsed > 100*time.Millisecond { + t.Fatalf("cache-only completion took %s", elapsed) + } + if requests.Load() != 0 { + t.Fatalf("cache-only completion contacted hung gateway %d time(s)", requests.Load()) + } +} diff --git a/internal/llm/gateway_provider.go b/internal/llm/gateway_provider.go new file mode 100644 index 000000000..e58a6c11c --- /dev/null +++ b/internal/llm/gateway_provider.go @@ -0,0 +1,701 @@ +package llm + +import ( + "bytes" + "context" + "crypto/rand" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "log/slog" + "net" + "net/http" + "net/url" + "os" + "path/filepath" + "strings" + "sync" + "time" + + "github.com/samsaffron/term-llm/internal/config" + "github.com/samsaffron/term-llm/internal/gateway/protocol" + "golang.org/x/sync/singleflight" +) + +type gatewayCatalogCache struct { + ETag string `json:"etag"` + FetchedAt time.Time `json:"fetched_at"` + Catalog protocol.Catalog `json:"catalog"` +} + +var gatewayCatalogProcessCache = struct { + sync.RWMutex + entries map[string]gatewayCatalogCache +}{entries: make(map[string]gatewayCatalogCache)} + +var gatewayCatalogRefresh singleflight.Group + +// GatewayProvider is the satellite-side remote Provider. The satellite engine +// still owns prompts, sessions, tools, approvals, and the local filesystem. +type GatewayProvider struct { + name string + gateway config.GatewayConfig + token string + baseURL *url.URL + streamHTTP *http.Client + shortHTTP *http.Client + + mu sync.RWMutex + state string + entry protocol.CatalogEntry +} + +func NewGatewayProvider(cfg *config.Config, name, _ string) (*GatewayProvider, error) { + if cfg == nil || !cfg.Gateway.Enabled() { + return nil, fmt.Errorf("gateway is not configured") + } + if err := cfg.Gateway.Validate(); err != nil { + return nil, err + } + token, err := cfg.Gateway.ResolveToken() + if err != nil { + return nil, err + } + base, err := url.Parse(strings.TrimRight(strings.TrimSpace(cfg.Gateway.URL), "/")) + if err != nil { + return nil, fmt.Errorf("parse gateway URL: %w", err) + } + connectTimeout := gatewayDuration(cfg.Gateway.ConnectTimeout, config.DefaultGatewayConnectTimeout) + responseTimeout := gatewayDuration(cfg.Gateway.ResponseTimeout, config.DefaultGatewayResponseTimeout) + transport := &http.Transport{ + Proxy: http.ProxyFromEnvironment, + DialContext: (&net.Dialer{Timeout: connectTimeout, KeepAlive: 30 * time.Second}).DialContext, + ForceAttemptHTTP2: true, + MaxIdleConns: 100, + MaxIdleConnsPerHost: 20, + ResponseHeaderTimeout: responseTimeout, + IdleConnTimeout: 90 * time.Second, + } + p := &GatewayProvider{ + name: name, gateway: cfg.Gateway, token: token, baseURL: base, + streamHTTP: &http.Client{Transport: transport}, + shortHTTP: &http.Client{Transport: transport.Clone(), Timeout: responseTimeout}, + } + catalogCtx, cancel := context.WithTimeout(context.Background(), connectTimeout+responseTimeout) + defer cancel() + catalog, err := p.loadCatalog(catalogCtx) + if err != nil { + return nil, err + } + entry, ok := catalogEntry(catalog, name) + if !ok { + return nil, fmt.Errorf("provider %q is not in gateway catalog", name) + } + p.entry = entry + return p, nil +} + +func gatewayDuration(value, fallback string) time.Duration { + d, err := time.ParseDuration(strings.TrimSpace(value)) + if err == nil && d > 0 { + return d + } + d, _ = time.ParseDuration(fallback) + return d +} + +func (p *GatewayProvider) Name() string { return p.name } +func (p *GatewayProvider) Credential() string { return "gateway" } +func (p *GatewayProvider) Capabilities() Capabilities { + return capabilitiesFromProtocol(p.entry.Capabilities) +} + +// GatewayHandlesRetries marks the provider as a single-attempt transport. The +// central provider owns retry policy, avoiding nested satellite retries. +func (p *GatewayProvider) GatewayHandlesRetries() bool { return true } + +// UsageTrackedExternallyBy keeps satellite-local visibility while marking the +// gateway-attributed copy so aggregate local usage does not double-count it. +func (p *GatewayProvider) UsageTrackedExternallyBy() string { return "gateway" } + +func (p *GatewayProvider) ExportProviderState() ([]byte, bool) { + p.mu.RLock() + defer p.mu.RUnlock() + if p.state == "" { + return nil, false + } + return []byte(p.state), true +} + +func (p *GatewayProvider) ImportProviderState(data []byte) error { + state := strings.TrimSpace(string(data)) + if state == "" { + return fmt.Errorf("gateway provider state is empty") + } + p.mu.Lock() + p.state = state + p.mu.Unlock() + return nil +} + +func (p *GatewayProvider) ListModels(ctx context.Context) ([]ModelInfo, error) { + catalog, err := p.loadCatalog(ctx) + if err != nil { + return nil, err + } + entry, ok := catalogEntry(catalog, p.name) + if !ok { + return nil, fmt.Errorf("provider %q is not in gateway catalog", p.name) + } + models := make([]ModelInfo, 0, len(entry.Models)) + for _, model := range entry.Models { + models = append(models, ModelInfo{ + ID: model.ID, DisplayName: model.DisplayName, Created: model.Created, + OwnedBy: model.OwnedBy, InputLimit: model.InputLimit, + InputPrice: model.InputPrice, OutputPrice: model.OutputPrice, + ReasoningEfforts: model.ReasoningEfforts, + DefaultReasoningEffort: model.DefaultReasoningEffort, + ReasoningModes: model.ReasoningModes, + }) + } + return models, nil +} + +func (p *GatewayProvider) Stream(ctx context.Context, req Request) (Stream, error) { + wireRequest, err := EncodeGatewayRequest(req) + if err != nil { + return nil, fmt.Errorf("encode gateway request: %w", err) + } + requestID, err := newGatewayID("req") + if err != nil { + return nil, err + } + p.mu.RLock() + state := p.state + p.mu.RUnlock() + payload, err := json.Marshal(protocol.InferenceRequest{ + Version: protocol.Version, RequestID: requestID, Provider: p.name, + State: state, Request: wireRequest, + }) + if err != nil { + return nil, fmt.Errorf("encode inference envelope: %w", err) + } + streamCtx, cancel := context.WithCancel(ctx) + httpReq, err := http.NewRequestWithContext(streamCtx, http.MethodPost, p.endpoint("/g1/inference"), bytes.NewReader(payload)) + if err != nil { + cancel() + return nil, fmt.Errorf("create gateway request: %w", err) + } + p.setHeaders(httpReq) + httpReq.Header.Set("Content-Type", "application/json") + httpReq.Header.Set("Accept", "text/event-stream") + resp, err := p.streamHTTP.Do(httpReq) + if err != nil { + cancel() + return nil, fmt.Errorf("gateway inference unavailable; check gateway URL/network and retry: %w", err) + } + if resp.StatusCode != http.StatusOK { + defer resp.Body.Close() + cancel() + return nil, decodeGatewayHTTPError(resp) + } + stream := &gatewayProviderStream{ + provider: p, body: resp.Body, ctx: streamCtx, cancel: cancel, + decoder: newSSEDecoder(resp.Body, sseDecoderOptions{Transport: "gateway SSE"}), + closedSignal: make(chan struct{}), results: make(chan nextGatewaySSE, 1), + } + go stream.readLoop() + return stream, nil +} + +func (p *GatewayProvider) endpoint(path string) string { + base := *p.baseURL + base.Path = strings.TrimRight(base.Path, "/") + path + base.RawQuery = "" + base.Fragment = "" + return base.String() +} + +func (p *GatewayProvider) setHeaders(req *http.Request) { + req.Header.Set("Authorization", "Bearer "+p.token) + req.Header.Set(protocol.VersionHeader, "1") +} + +type nextGatewaySSE struct { + data []byte + err error +} + +type gatewayProviderStream struct { + provider *GatewayProvider + body io.ReadCloser + decoder *sseDecoder + ctx context.Context + cancel context.CancelFunc + results chan nextGatewaySSE + recvMu sync.Mutex + mu sync.Mutex + runID string + closed bool + closedSignal chan struct{} + done bool +} + +func (s *gatewayProviderStream) readLoop() { + defer close(s.results) + for { + _, data, err := s.decoder.Next() + select { + case s.results <- nextGatewaySSE{data: data, err: err}: + case <-s.ctx.Done(): + return + case <-s.closedSignal: + return + } + if err != nil { + return + } + } +} + +func (s *gatewayProviderStream) Recv() (Event, error) { + s.recvMu.Lock() + defer s.recvMu.Unlock() + if s.done { + return Event{}, io.EOF + } + for { + idle := gatewayDuration(s.provider.gateway.IdleTimeout, config.DefaultGatewayIdleTimeout) + timer := time.NewTimer(idle) + var result nextGatewaySSE + var ok bool + select { + case result, ok = <-s.results: + if !timer.Stop() { + select { + case <-timer.C: + default: + } + } + case <-timer.C: + s.terminate() + return Event{}, fmt.Errorf("gateway stream idle timeout after %s; check gateway/provider health and retry", idle) + case <-s.ctx.Done(): + timer.Stop() + if s.isClosed() { + return Event{}, io.EOF + } + err := s.ctx.Err() + s.terminate() + return Event{}, fmt.Errorf("gateway stream ended: %w", err) + case <-s.closedSignal: + timer.Stop() + return Event{}, io.EOF + } + if !ok { + if s.isClosed() { + return Event{}, io.EOF + } + s.terminate() + return Event{}, &StreamIncompleteError{Transport: "gateway SSE", Terminal: "done"} + } + if result.err != nil { + s.terminate() + if errors.Is(result.err, io.EOF) { + return Event{}, &StreamIncompleteError{Transport: "gateway SSE", Terminal: "done"} + } + return Event{}, result.err + } + var record protocol.StreamRecord + dec := json.NewDecoder(bytes.NewReader(result.data)) + dec.DisallowUnknownFields() + if err := dec.Decode(&record); err != nil { + s.terminate() + return Event{}, fmt.Errorf("decode gateway stream record: %w", err) + } + if record.Version != protocol.Version { + s.terminate() + return Event{}, fmt.Errorf("gateway protocol version %d is unsupported", record.Version) + } + s.mu.Lock() + if record.RunID != "" { + s.runID = record.RunID + } + s.mu.Unlock() + switch record.Type { + case "run": + continue + case "event": + event, err := DecodeGatewayEvent(record.Event) + if err != nil { + s.terminate() + return Event{}, err + } + return event, nil + case "tool_callback": + event, err := DecodeGatewayEvent(record.Event) + if err != nil { + s.terminate() + return Event{}, err + } + response := make(chan ToolExecutionResponse, 1) + event.ToolResponse = response + go s.postToolResult(record.CallbackPath, response) + return event, nil + case "state": + s.provider.mu.Lock() + s.provider.state = record.State + s.provider.mu.Unlock() + continue + case "done": + s.done = true + s.terminate() + return Event{Type: EventDone}, nil + case "error": + s.terminate() + if record.Error == nil { + return Event{}, &GatewayError{Code: "gateway_failure", Message: "gateway request failed; retry or check gateway diagnostics"} + } + return Event{}, &GatewayError{Code: record.Error.Code, Message: record.Error.Message, RequestID: record.Error.RequestID} + default: + s.terminate() + return Event{}, fmt.Errorf("unknown gateway stream record %q", record.Type) + } + } +} + +func (s *gatewayProviderStream) postToolResult(callbackPath string, response <-chan ToolExecutionResponse) { + var result ToolExecutionResponse + var ok bool + toolTimeout := gatewayDuration(s.provider.gateway.ToolTimeout, config.DefaultGatewayToolTimeout) + timer := time.NewTimer(toolTimeout) + defer timer.Stop() + select { + case result, ok = <-response: + if !ok { + result.Err = fmt.Errorf("satellite tool response channel closed") + } + case <-timer.C: + result.Err = fmt.Errorf("satellite tool execution timed out after %s", toolTimeout) + case <-s.ctx.Done(): + return + case <-s.closedSignal: + return + } + data, err := EncodeGatewayToolResponse(result) + if err != nil { + slog.Error("encode gateway tool callback", "error", err) + return + } + payload, _ := json.Marshal(protocol.ToolResultRequest{Version: protocol.Version, Result: data}) + callbackURL, err := s.safeCallbackURL(callbackPath) + if err != nil { + slog.Error("validate gateway tool callback", "error", err) + return + } + req, err := http.NewRequestWithContext(s.ctx, http.MethodPost, callbackURL, bytes.NewReader(payload)) + if err != nil { + slog.Error("create gateway tool callback", "error", err) + return + } + s.provider.setHeaders(req) + req.Header.Set("Content-Type", "application/json") + resp, err := s.provider.shortHTTP.Do(req) + if err != nil { + slog.Error("post gateway tool callback", "error", err) + return + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusNoContent { + slog.Error("gateway tool callback rejected", "status", resp.StatusCode) + } +} + +func (s *gatewayProviderStream) safeCallbackURL(path string) (string, error) { + u, err := url.Parse(path) + if err != nil || u.IsAbs() || !strings.HasPrefix(u.Path, "/g1/runs/") { + return "", fmt.Errorf("invalid gateway callback path") + } + return s.provider.endpoint(u.Path), nil +} + +func (s *gatewayProviderStream) isClosed() bool { + s.mu.Lock() + defer s.mu.Unlock() + return s.closed +} + +func (s *gatewayProviderStream) terminate() { + s.mu.Lock() + if s.closed { + s.mu.Unlock() + return + } + s.closed = true + close(s.closedSignal) + s.mu.Unlock() + s.cancel() + _ = s.body.Close() +} + +func (s *gatewayProviderStream) Close() error { + s.mu.Lock() + if s.closed { + s.mu.Unlock() + return nil + } + runID := s.runID + s.mu.Unlock() + s.terminate() + if runID != "" { + timeout := min(gatewayDuration(s.provider.gateway.ResponseTimeout, config.DefaultGatewayResponseTimeout), 500*time.Millisecond) + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + req, err := http.NewRequestWithContext(ctx, http.MethodDelete, s.provider.endpoint("/g1/runs/"+url.PathEscape(runID)), nil) + if err == nil { + s.provider.setHeaders(req) + if resp, doErr := s.provider.shortHTTP.Do(req); doErr == nil { + resp.Body.Close() + } + } + } + return nil +} + +func (p *GatewayProvider) loadCatalog(ctx context.Context) (protocol.Catalog, error) { + identity := gatewayCatalogIdentity(p.baseURL.String(), p.token) + cachePath := gatewayCatalogCachePath(p.baseURL.String(), p.token) + cached, cacheErr := gatewayCatalogCached(identity, cachePath) + ttl := gatewayDuration(p.gateway.CatalogTTL, config.DefaultGatewayCatalogTTL) + if cacheErr == nil && time.Since(cached.FetchedAt) < ttl { + return cached.Catalog, nil + } + if ctx == nil { + ctx = context.Background() + } + result := gatewayCatalogRefresh.DoChan(identity, func() (any, error) { + latest, latestErr := gatewayCatalogCached(identity, cachePath) + if latestErr == nil && time.Since(latest.FetchedAt) < ttl { + return latest.Catalog, nil + } + if latestErr != nil && cacheErr == nil { + latest, latestErr = cached, nil + } + bound := gatewayDuration(p.gateway.ConnectTimeout, config.DefaultGatewayConnectTimeout) + gatewayDuration(p.gateway.ResponseTimeout, config.DefaultGatewayResponseTimeout) + refreshCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), bound) + defer cancel() + return p.fetchCatalog(refreshCtx, identity, cachePath, latest, latestErr == nil) + }) + select { + case <-ctx.Done(): + if cacheErr == nil { + return cached.Catalog, nil + } + return protocol.Catalog{}, fmt.Errorf("gateway catalog unavailable; check gateway URL/network and retry: %w", ctx.Err()) + case loaded := <-result: + if loaded.Err != nil { + if cacheErr == nil { + return cached.Catalog, nil + } + return protocol.Catalog{}, loaded.Err + } + return loaded.Val.(protocol.Catalog), nil + } +} + +func (p *GatewayProvider) fetchCatalog(ctx context.Context, identity, cachePath string, cached gatewayCatalogCache, hasCached bool) (protocol.Catalog, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, p.endpoint("/g1/catalog"), nil) + if err != nil { + return protocol.Catalog{}, err + } + p.setHeaders(req) + if hasCached && cached.ETag != "" { + req.Header.Set("If-None-Match", cached.ETag) + } + resp, err := p.shortHTTP.Do(req) + if err != nil { + if hasCached { + return cached.Catalog, nil + } + return protocol.Catalog{}, fmt.Errorf("fetch gateway catalog: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode == http.StatusNotModified && hasCached { + cached.FetchedAt = time.Now().UTC() + storeGatewayCatalogCache(identity, cachePath, cached) + return cached.Catalog, nil + } + if resp.StatusCode != http.StatusOK { + if hasCached { + return cached.Catalog, nil + } + return protocol.Catalog{}, decodeGatewayHTTPError(resp) + } + var catalog protocol.Catalog + dec := json.NewDecoder(io.LimitReader(resp.Body, 8<<20)) + dec.DisallowUnknownFields() + if err := dec.Decode(&catalog); err != nil { + if hasCached { + return cached.Catalog, nil + } + return protocol.Catalog{}, fmt.Errorf("decode gateway catalog: %w", err) + } + if catalog.Version != protocol.Version { + return protocol.Catalog{}, fmt.Errorf("gateway catalog protocol version %d is unsupported", catalog.Version) + } + storeGatewayCatalogCache(identity, cachePath, gatewayCatalogCache{ETag: resp.Header.Get("ETag"), FetchedAt: time.Now().UTC(), Catalog: catalog}) + return catalog, nil +} + +func (p *GatewayProvider) loadCatalogCacheOnly() (protocol.Catalog, error) { + identity := gatewayCatalogIdentity(p.baseURL.String(), p.token) + cached, err := gatewayCatalogCached(identity, gatewayCatalogCachePath(p.baseURL.String(), p.token)) + if err != nil { + return protocol.Catalog{}, fmt.Errorf("gateway catalog cache is empty; run a normal command while the gateway is available") + } + return cached.Catalog, nil +} + +func gatewayCatalogCached(identity, path string) (gatewayCatalogCache, error) { + gatewayCatalogProcessCache.RLock() + cached, ok := gatewayCatalogProcessCache.entries[identity] + gatewayCatalogProcessCache.RUnlock() + if ok { + return cached, nil + } + cached, err := readGatewayCatalogCache(path) + if err != nil { + return gatewayCatalogCache{}, err + } + gatewayCatalogProcessCache.Lock() + gatewayCatalogProcessCache.entries[identity] = cached + gatewayCatalogProcessCache.Unlock() + return cached, nil +} + +func storeGatewayCatalogCache(identity, path string, cached gatewayCatalogCache) { + gatewayCatalogProcessCache.Lock() + gatewayCatalogProcessCache.entries[identity] = cached + gatewayCatalogProcessCache.Unlock() + _ = writeGatewayCatalogCache(path, cached) +} + +func GatewayCatalogHasProvider(ctx context.Context, cfg *config.Config, name string) (bool, error) { + p, err := NewGatewayProviderForCatalog(cfg) + if err != nil { + return false, err + } + catalog, err := p.loadCatalog(ctx) + if err != nil { + return false, err + } + _, ok := catalogEntry(catalog, name) + return ok, nil +} + +func NewGatewayProviderForCatalog(cfg *config.Config) (*GatewayProvider, error) { + if cfg == nil || !cfg.Gateway.Enabled() { + return nil, fmt.Errorf("gateway is not configured") + } + token, err := cfg.Gateway.ResolveToken() + if err != nil { + return nil, err + } + base, err := url.Parse(strings.TrimRight(strings.TrimSpace(cfg.Gateway.URL), "/")) + if err != nil { + return nil, err + } + connectTimeout := gatewayDuration(cfg.Gateway.ConnectTimeout, config.DefaultGatewayConnectTimeout) + responseTimeout := gatewayDuration(cfg.Gateway.ResponseTimeout, config.DefaultGatewayResponseTimeout) + transport := &http.Transport{Proxy: http.ProxyFromEnvironment, DialContext: (&net.Dialer{Timeout: connectTimeout}).DialContext, ForceAttemptHTTP2: true, ResponseHeaderTimeout: responseTimeout} + return &GatewayProvider{gateway: cfg.Gateway, token: token, baseURL: base, streamHTTP: &http.Client{Transport: transport}, shortHTTP: &http.Client{Transport: transport.Clone(), Timeout: responseTimeout}}, nil +} + +func LoadGatewayCatalogCacheOnly(cfg *config.Config) (protocol.Catalog, error) { + provider, err := NewGatewayProviderForCatalog(cfg) + if err != nil { + return protocol.Catalog{}, err + } + return provider.loadCatalogCacheOnly() +} + +func catalogEntry(catalog protocol.Catalog, name string) (protocol.CatalogEntry, bool) { + for _, entry := range catalog.Providers { + if entry.Key == name { + return entry, true + } + } + return protocol.CatalogEntry{}, false +} + +func gatewayCatalogIdentity(rawURL, token string) string { + sum := sha256.Sum256([]byte(strings.TrimRight(strings.TrimSpace(rawURL), "/") + "\x00" + token)) + return hex.EncodeToString(sum[:]) +} + +func gatewayCatalogCachePath(rawURL, token string) string { + root := strings.TrimSpace(os.Getenv("XDG_CACHE_HOME")) + if root == "" { + if home, err := os.UserHomeDir(); err == nil { + root = filepath.Join(home, ".cache") + } else { + root = os.TempDir() + } + } + identity := gatewayCatalogIdentity(rawURL, token) + return filepath.Join(root, "term-llm", "gateway", identity[:16]+".json") +} + +func readGatewayCatalogCache(path string) (gatewayCatalogCache, error) { + data, err := os.ReadFile(path) + if err != nil { + return gatewayCatalogCache{}, err + } + var cache gatewayCatalogCache + if err := json.Unmarshal(data, &cache); err != nil { + return gatewayCatalogCache{}, err + } + if cache.Catalog.Version != protocol.Version || cache.FetchedAt.IsZero() { + return gatewayCatalogCache{}, fmt.Errorf("invalid gateway catalog cache") + } + return cache, nil +} + +func writeGatewayCatalogCache(path string, cache gatewayCatalogCache) error { + if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil { + return err + } + data, err := json.Marshal(cache) + if err != nil { + return err + } + tmp := path + ".tmp" + if err := os.WriteFile(tmp, data, 0o600); err != nil { + return err + } + return os.Rename(tmp, path) +} + +func decodeGatewayHTTPError(resp *http.Response) error { + var wire protocol.Error + data, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20)) + if err := json.Unmarshal(data, &wire); err == nil && wire.Code != "" { + return &GatewayError{Code: wire.Code, Message: wire.Message, RequestID: wire.RequestID} + } + return &GatewayError{Code: "http_error", Message: fmt.Sprintf("gateway returned HTTP %d; check gateway health and credentials", resp.StatusCode)} +} + +func newGatewayID(prefix string) (string, error) { + var raw [16]byte + if _, err := rand.Read(raw[:]); err != nil { + return "", fmt.Errorf("generate gateway request ID: %w", err) + } + return prefix + "_" + hex.EncodeToString(raw[:]), nil +} + +var _ Provider = (*GatewayProvider)(nil) +var _ ProviderStateExporter = (*GatewayProvider)(nil) +var _ ProviderStateImporter = (*GatewayProvider)(nil) diff --git a/internal/llm/gateway_retry_budget_test.go b/internal/llm/gateway_retry_budget_test.go new file mode 100644 index 000000000..8fa55ec6a --- /dev/null +++ b/internal/llm/gateway_retry_budget_test.go @@ -0,0 +1,105 @@ +package llm + +import ( + "context" + "errors" + "io" + "testing" + "time" +) + +type retryDeadlineProvider struct { + attempts int + block bool +} + +func (p *retryDeadlineProvider) Name() string { return "retry-deadline" } +func (p *retryDeadlineProvider) Credential() string { return "mock" } +func (p *retryDeadlineProvider) Capabilities() Capabilities { return Capabilities{} } +func (p *retryDeadlineProvider) Stream(ctx context.Context, _ Request) (Stream, error) { + p.attempts++ + if p.block { + <-ctx.Done() + return nil, ctx.Err() + } + return nil, errors.New("500 Internal Server Error") +} + +func drainRetryFailure(t *testing.T, provider Provider) error { + t.Helper() + stream, err := provider.Stream(t.Context(), Request{}) + if err != nil { + return err + } + defer stream.Close() + for { + event, recvErr := stream.Recv() + if recvErr != nil { + if recvErr == io.EOF { + return nil + } + return recvErr + } + if event.Type == EventError { + return event.Err + } + } +} + +type gatewayOwnedRetryProvider struct{ retryDeadlineProvider } + +func (*gatewayOwnedRetryProvider) GatewayHandlesRetries() bool { return true } + +func TestEngineDoesNotReplayGatewayOwnedRetryFailures(t *testing.T) { + provider := &gatewayOwnedRetryProvider{} + engine := NewEngine(provider, nil) + stream, err := engine.Stream(t.Context(), Request{Model: "model", Messages: []Message{UserText("hello")}}) + if err != nil { + t.Fatal(err) + } + defer stream.Close() + for { + _, recvErr := stream.Recv() + if recvErr == io.EOF { + break + } + if recvErr != nil { + break + } + } + if provider.attempts != 1 { + t.Fatalf("engine gateway attempts = %d, want 1", provider.attempts) + } +} + +func TestRetryProviderPersistent500HonorsAttemptLimitWithoutNestedRetries(t *testing.T) { + inner := &retryDeadlineProvider{} + first := WrapWithRetry(inner, RetryConfig{MaxAttempts: 20, MaxElapsedTime: time.Second, BaseBackoff: time.Millisecond, MaxBackoff: time.Millisecond}) + provider := WrapWithRetry(first, RetryConfig{MaxAttempts: 3, MaxElapsedTime: time.Second, BaseBackoff: time.Millisecond, MaxBackoff: time.Millisecond}) + started := time.Now() + if err := drainRetryFailure(t, provider); err == nil { + t.Fatal("persistent 500 unexpectedly succeeded") + } + if inner.attempts != 3 { + t.Fatalf("attempts = %d, want exactly 3", inner.attempts) + } + if elapsed := time.Since(started); elapsed > 250*time.Millisecond { + t.Fatalf("attempt-limited retry took %s", elapsed) + } +} + +func TestRetryProviderElapsedBudgetCancelsHungAttempt(t *testing.T) { + inner := &retryDeadlineProvider{block: true} + provider := WrapWithRetry(inner, RetryConfig{MaxAttempts: 5, MaxElapsedTime: 40 * time.Millisecond, BaseBackoff: time.Millisecond, MaxBackoff: time.Millisecond}) + started := time.Now() + err := drainRetryFailure(t, provider) + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("error = %v, want deadline exceeded", err) + } + if inner.attempts != 1 { + t.Fatalf("hung attempt count = %d, want 1", inner.attempts) + } + if elapsed := time.Since(started); elapsed > 250*time.Millisecond { + t.Fatalf("elapsed retry budget took %s", elapsed) + } +} diff --git a/internal/llm/gateway_routing_test.go b/internal/llm/gateway_routing_test.go new file mode 100644 index 000000000..94916dc65 --- /dev/null +++ b/internal/llm/gateway_routing_test.go @@ -0,0 +1,131 @@ +package llm + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/samsaffron/term-llm/internal/config" + "github.com/samsaffron/term-llm/internal/gateway/protocol" +) + +func gatewayCatalogServer(t *testing.T) *httptest.Server { + t.Helper() + return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/g1/catalog" { + http.NotFound(w, r) + return + } + _ = json.NewEncoder(w).Encode(protocol.Catalog{Version: protocol.Version, GeneratedAt: time.Now(), Providers: []protocol.CatalogEntry{{Key: "remote", Models: []protocol.Model{{ID: "model-a"}}}}}) + })) +} + +func TestGatewayRoutingPrecedence(t *testing.T) { + t.Setenv("XDG_CACHE_HOME", t.TempDir()) + server := gatewayCatalogServer(t) + defer server.Close() + gatewayCfg := config.GatewayConfig{URL: server.URL, Token: "token", CatalogTTL: "1ms", ConnectTimeout: "1s", ResponseTimeout: "1s"} + + remoteCfg := &config.Config{Gateway: gatewayCfg, Providers: map[string]config.ProviderConfig{}} + provider, err := NewProviderByName(remoteCfg, "remote", "model-a") + if err != nil { + t.Fatal(err) + } + if _, ok := provider.(*GatewayProvider); !ok { + t.Fatalf("catalog provider routed to %T, want GatewayProvider", provider) + } + if got, _, err := ParseProviderModel("remote:model-a", remoteCfg); err != nil || got != "remote" { + t.Fatalf("ParseProviderModel remote = %q, %v", got, err) + } + + explicit := &config.Config{Gateway: gatewayCfg, Providers: map[string]config.ProviderConfig{"remote": {Type: config.ProviderTypeZen, Model: "model-a"}}} + provider, err = NewProviderByName(explicit, "remote", "model-a") + if err != nil { + t.Fatal(err) + } + if _, ok := provider.(*GatewayProvider); ok { + t.Fatalf("explicit local provider was routed remotely") + } + + localList := &config.Config{Gateway: gatewayCfg, Providers: map[string]config.ProviderConfig{"remote": {Type: config.ProviderTypeZen, Model: "model-a"}}} + localList.Gateway.LocalProviders = []string{"remote"} + provider, err = NewProviderByName(localList, "remote", "model-a") + if err != nil { + t.Fatal(err) + } + if _, ok := provider.(*GatewayProvider); ok { + t.Fatalf("gateway.local_providers did not win") + } +} + +func TestGatewayAdvertisedDebugRoutesRemoteUnlessExplicitlyLocal(t *testing.T) { + t.Setenv("XDG_CACHE_HOME", t.TempDir()) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _ = json.NewEncoder(w).Encode(protocol.Catalog{Version: protocol.Version, Providers: []protocol.CatalogEntry{{Key: "debug", Type: string(config.ProviderTypeDebug), Models: []protocol.Model{{ID: "fast"}}}}}) + })) + defer server.Close() + gatewayCfg := config.GatewayConfig{URL: server.URL, Token: "token", CatalogTTL: "1ms", ConnectTimeout: "1s", ResponseTimeout: "1s"} + + remoteCfg := &config.Config{Gateway: gatewayCfg, Providers: map[string]config.ProviderConfig{}} + provider, err := NewProviderByName(remoteCfg, "debug", "fast") + if err != nil { + t.Fatal(err) + } + if _, ok := provider.(*GatewayProvider); !ok { + t.Fatalf("gateway-advertised debug routed to %T, want GatewayProvider", provider) + } + + explicitCfg := &config.Config{Gateway: gatewayCfg, Providers: map[string]config.ProviderConfig{"debug": {Model: "fast", Models: []string{"fast"}}}} + provider, err = NewProviderByName(explicitCfg, "debug", "fast") + if err != nil { + t.Fatal(err) + } + if _, ok := provider.(*GatewayProvider); ok { + t.Fatal("explicitly configured local debug was routed remotely") + } + if provider.Name() != "debug:fast" { + t.Fatalf("explicit local debug model = %q, want debug:fast", provider.Name()) + } + + listedCfg := &config.Config{Gateway: gatewayCfg, Providers: map[string]config.ProviderConfig{}} + listedCfg.Gateway.LocalProviders = []string{"debug"} + provider, err = NewProviderByName(listedCfg, "debug", "fast") + if err != nil { + t.Fatal(err) + } + if _, ok := provider.(*GatewayProvider); ok { + t.Fatal("gateway.local_providers debug exception was routed remotely") + } +} + +func TestGatewayFailsClosedByDefaultAndNoGatewayCompatibility(t *testing.T) { + cfg := &config.Config{Gateway: config.GatewayConfig{URL: "http://127.0.0.1:1", Token: "token", ConnectTimeout: "10ms", ResponseTimeout: "10ms"}, Providers: map[string]config.ProviderConfig{}} + if _, err := NewProviderByName(cfg, "remote", "model"); err == nil || !strings.Contains(err.Error(), "gateway unavailable") || !strings.Contains(err.Error(), "URL/network/token") { + t.Fatalf("default gateway outage error = %v", err) + } + localOverride := &config.Config{ + Gateway: config.GatewayConfig{URL: "http://127.0.0.1:1", Token: "token", LocalProviders: []string{"zen"}, ConnectTimeout: "10ms", ResponseTimeout: "10ms"}, + Providers: map[string]config.ProviderConfig{"zen": {Type: config.ProviderTypeZen, Model: "minimax-m2.5-free"}}, + } + provider, err := NewProviderByName(localOverride, "zen", "minimax-m2.5-free") + if err != nil { + t.Fatalf("explicit local provider failed during gateway outage: %v", err) + } + if _, ok := provider.(*GatewayProvider); ok { + t.Fatal("gateway outage overrode explicit local provider") + } + local := &config.Config{DefaultProvider: "debug", Providers: map[string]config.ProviderConfig{"debug": {Model: "fast"}}} + provider, err = NewProvider(local) + if err != nil { + t.Fatalf("no-gateway behavior failed: %v", err) + } + if _, ok := provider.(*GatewayProvider); ok { + t.Fatalf("no-gateway config unexpectedly routed remotely") + } + if provider.Name() != "debug:fast" { + t.Fatalf("local debug model override = %q, want debug:fast", provider.Name()) + } +} diff --git a/internal/llm/gateway_stream_test.go b/internal/llm/gateway_stream_test.go new file mode 100644 index 000000000..da7eda530 --- /dev/null +++ b/internal/llm/gateway_stream_test.go @@ -0,0 +1,65 @@ +package llm + +import ( + "context" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/samsaffron/term-llm/internal/config" + "github.com/samsaffron/term-llm/internal/gateway/protocol" +) + +func TestGatewayStreamUsesSingleReaderAndIdleTimeoutTerminatesIt(t *testing.T) { + resetGatewayCatalogProcessCacheForTest() + t.Cleanup(resetGatewayCatalogProcessCacheForTest) + t.Setenv("XDG_CACHE_HOME", t.TempDir()) + streamCanceled := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/g1/catalog": + _ = json.NewEncoder(w).Encode(protocol.Catalog{Version: protocol.Version, Providers: []protocol.CatalogEntry{{Key: "remote", Models: []protocol.Model{{ID: "model-a"}}}}}) + case "/g1/inference": + w.Header().Set("Content-Type", "text/event-stream") + flusher := w.(http.Flusher) + event, _ := EncodeGatewayEvent(Event{Type: EventTextDelta, Text: "first"}, nil) + record, _ := json.Marshal(protocol.StreamRecord{Version: protocol.Version, Type: "event", Event: event}) + fmt.Fprintf(w, "event: gateway\ndata: %s\n\n", record) + flusher.Flush() + <-r.Context().Done() + close(streamCanceled) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + cfg := &config.Config{Gateway: config.GatewayConfig{URL: server.URL, Token: "token", CatalogTTL: "1m", ConnectTimeout: "1s", ResponseTimeout: "1s", IdleTimeout: "25ms"}} + provider, err := NewGatewayProvider(cfg, "remote", "model-a") + if err != nil { + t.Fatal(err) + } + stream, err := provider.Stream(context.Background(), Request{Model: "model-a", Messages: []Message{UserText("hi")}}) + if err != nil { + t.Fatal(err) + } + first, err := stream.Recv() + if err != nil || first.Text != "first" { + t.Fatalf("first event = %+v, %v", first, err) + } + _, err = stream.Recv() + if err == nil || !strings.Contains(err.Error(), "idle timeout") || !strings.Contains(err.Error(), "gateway") { + t.Fatalf("idle error = %v", err) + } + select { + case <-streamCanceled: + case <-time.After(time.Second): + t.Fatal("idle timeout left gateway SSE reader/request running") + } + if err := stream.Close(); err != nil { + t.Fatal(err) + } +} diff --git a/internal/llm/gateway_wire.go b/internal/llm/gateway_wire.go new file mode 100644 index 000000000..58fab7733 --- /dev/null +++ b/internal/llm/gateway_wire.go @@ -0,0 +1,263 @@ +package llm + +import ( + "bytes" + "encoding/json" + "fmt" + + "github.com/samsaffron/term-llm/internal/gateway/protocol" +) + +// gatewayRequest is deliberately distinct from Request. In particular it has +// no WorkingDir, approval transcript, execution filters, or debug fields. This +// makes local filesystem/process state structurally impossible on the wire. +type gatewayRequest struct { + Model string `json:"model"` + SessionID string `json:"session_id,omitempty"` + Ephemeral bool `json:"ephemeral,omitempty"` + IncludeDeveloperInContinuation bool `json:"include_developer_in_continuation,omitempty"` + Messages []Message `json:"messages"` + Tools []ToolSpec `json:"tools,omitempty"` + ToolChoice ToolChoice `json:"tool_choice"` + LastTurnToolChoice *ToolChoice `json:"last_turn_tool_choice,omitempty"` + ParallelToolCalls bool `json:"parallel_tool_calls,omitempty"` + Search bool `json:"search,omitempty"` + ForceExternalSearch bool `json:"force_external_search,omitempty"` + DisableExternalWebFetch bool `json:"disable_external_web_fetch,omitempty"` + ReasoningEffort string `json:"reasoning_effort,omitempty"` + Responses *ResponsesOptions `json:"responses,omitempty"` + MaxOutputTokens int `json:"max_output_tokens,omitempty"` + Temperature float32 `json:"temperature,omitempty"` + TemperatureSet bool `json:"temperature_set,omitempty"` + TopP float32 `json:"top_p,omitempty"` + TopPSet bool `json:"top_p_set,omitempty"` + ServiceTier string `json:"service_tier,omitempty"` + ServiceTierSet bool `json:"service_tier_set,omitempty"` + MaxTurns int `json:"max_turns,omitempty"` + ToolMap map[string]string `json:"tool_map,omitempty"` +} + +func sanitizeGatewayMessages(messages []Message) []Message { + out := make([]Message, len(messages)) + for i, message := range messages { + out[i] = message + out[i].Parts = make([]Part, 0, len(message.Parts)) + for _, part := range message.Parts { + if part.Type == PartSkillActivation || part.Type == PartToolActivity { + continue + } + part.ImagePath = "" + part.FilePath = "" + if part.ToolResult != nil { + result := *part.ToolResult + result.Diffs = nil + result.Images = nil + part.ToolResult = &result + } + out[i].Parts = append(out[i].Parts, part) + } + } + return out +} + +// EncodeGatewayRequest serializes provider-neutral request data while omitting +// satellite-only fields by construction. +func EncodeGatewayRequest(req Request) (json.RawMessage, error) { + wire := gatewayRequest{ + Model: req.Model, SessionID: req.SessionID, Ephemeral: req.Ephemeral, + IncludeDeveloperInContinuation: req.IncludeDeveloperInContinuation, + Messages: sanitizeGatewayMessages(req.Messages), Tools: req.Tools, + ToolChoice: req.ToolChoice, LastTurnToolChoice: req.LastTurnToolChoice, + ParallelToolCalls: req.ParallelToolCalls, Search: req.Search, + ForceExternalSearch: req.ForceExternalSearch, DisableExternalWebFetch: req.DisableExternalWebFetch, + ReasoningEffort: req.ReasoningEffort, Responses: req.Responses, + MaxOutputTokens: req.MaxOutputTokens, Temperature: req.Temperature, + TemperatureSet: req.TemperatureSet, TopP: req.TopP, TopPSet: req.TopPSet, + ServiceTier: req.ServiceTier, ServiceTierSet: req.ServiceTierSet, + MaxTurns: req.MaxTurns, ToolMap: req.ToolMap, + } + data, err := json.Marshal(wire) + return data, err +} + +// DecodeGatewayRequest strictly decodes the current top-level request schema. +// Nested provider-neutral structs retain Go's forward-compatible unknown-field +// behavior so additive message metadata can cross mixed patch versions. +func DecodeGatewayRequest(data []byte) (Request, error) { + var wire gatewayRequest + dec := json.NewDecoder(bytes.NewReader(data)) + dec.DisallowUnknownFields() + if err := dec.Decode(&wire); err != nil { + return Request{}, fmt.Errorf("decode gateway request: %w", err) + } + return Request{ + Model: wire.Model, SessionID: wire.SessionID, Ephemeral: wire.Ephemeral, + IncludeDeveloperInContinuation: wire.IncludeDeveloperInContinuation, + Messages: wire.Messages, Tools: wire.Tools, ToolChoice: wire.ToolChoice, + LastTurnToolChoice: wire.LastTurnToolChoice, ParallelToolCalls: wire.ParallelToolCalls, + Search: wire.Search, ForceExternalSearch: wire.ForceExternalSearch, + DisableExternalWebFetch: wire.DisableExternalWebFetch, ReasoningEffort: wire.ReasoningEffort, + Responses: wire.Responses, MaxOutputTokens: wire.MaxOutputTokens, + Temperature: wire.Temperature, TemperatureSet: wire.TemperatureSet, + TopP: wire.TopP, TopPSet: wire.TopPSet, ServiceTier: wire.ServiceTier, + ServiceTierSet: wire.ServiceTierSet, MaxTurns: wire.MaxTurns, ToolMap: wire.ToolMap, + }, nil +} + +type gatewayEvent struct { + Type EventType `json:"type"` + Text string `json:"text,omitempty"` + Model string `json:"model,omitempty"` + ReasoningEffort string `json:"reasoning_effort,omitempty"` + InterjectionID string `json:"interjection_id,omitempty"` + InterjectionStatus InterjectionStatus `json:"interjection_status,omitempty"` + Message *Message `json:"message,omitempty"` + ReasoningItemID string `json:"reasoning_item_id,omitempty"` + ReasoningEncryptedContent string `json:"reasoning_encrypted_content,omitempty"` + ReasoningKind ReasoningKind `json:"reasoning_kind,omitempty"` + ReasoningSummaryParts []string `json:"reasoning_summary_parts,omitempty"` + ReasoningIndex int `json:"reasoning_index,omitempty"` + ReasoningFinal bool `json:"reasoning_final,omitempty"` + Tool *ToolCall `json:"tool,omitempty"` + ToolCallID string `json:"tool_call_id,omitempty"` + ToolName string `json:"tool_name,omitempty"` + ToolInfo string `json:"tool_info,omitempty"` + ToolArgs json.RawMessage `json:"tool_args,omitempty"` + ToolSuccess bool `json:"tool_success,omitempty"` + ToolOutput string `json:"tool_output,omitempty"` + ToolDiffs []DiffData `json:"tool_diffs,omitempty"` + ToolFileChanges []FileChange `json:"tool_file_changes,omitempty"` + ToolImages []string `json:"tool_images,omitempty"` + Use *Usage `json:"usage,omitempty"` + Error *protocol.Error `json:"error,omitempty"` + RetryAttempt int `json:"retry_attempt,omitempty"` + RetryMaxAttempts int `json:"retry_max_attempts,omitempty"` + RetryWaitSecs float64 `json:"retry_wait_secs,omitempty"` + ToolActivity *ToolActivity `json:"tool_activity,omitempty"` + ProviderReplay *ProviderReplayItem `json:"provider_replay,omitempty"` + ImageData []byte `json:"image_data,omitempty"` + ImageMimeType string `json:"image_mime_type,omitempty"` + RevisedPrompt string `json:"revised_prompt,omitempty"` +} + +// EncodeGatewayEvent serializes every provider-neutral event field. Callback +// channels and concrete error values never cross the boundary. +func EncodeGatewayEvent(event Event, publicError *protocol.Error) (json.RawMessage, error) { + var message *Message + if event.Message.Role != "" || len(event.Message.Parts) > 0 || event.Message.CacheAnchor || event.Message.ApprovalRole != "" || event.Message.ClientMessageID != "" || event.Message.ResponseID != "" || event.Message.AssistantSegmentOrdinal != 0 || event.Message.SegmentStartSequence != 0 || event.Message.SegmentEndSequence != 0 { + copy := event.Message + message = © + } + wire := gatewayEvent{ + Type: event.Type, Text: event.Text, Model: event.Model, ReasoningEffort: event.ReasoningEffort, + InterjectionID: event.InterjectionID, InterjectionStatus: event.InterjectionStatus, Message: message, + ReasoningItemID: event.ReasoningItemID, ReasoningEncryptedContent: event.ReasoningEncryptedContent, + ReasoningKind: event.ReasoningKind, ReasoningSummaryParts: event.ReasoningSummaryParts, + ReasoningIndex: event.ReasoningIndex, ReasoningFinal: event.ReasoningFinal, + Tool: event.Tool, ToolCallID: event.ToolCallID, ToolName: event.ToolName, ToolInfo: event.ToolInfo, + ToolArgs: event.ToolArgs, ToolSuccess: event.ToolSuccess, ToolOutput: event.ToolOutput, + ToolDiffs: event.ToolDiffs, ToolFileChanges: event.ToolFileChanges, ToolImages: event.ToolImages, + Use: event.Use, Error: publicError, RetryAttempt: event.RetryAttempt, + RetryMaxAttempts: event.RetryMaxAttempts, RetryWaitSecs: event.RetryWaitSecs, + ToolActivity: event.ToolActivity, ProviderReplay: event.ProviderReplay, + ImageData: event.ImageData, ImageMimeType: event.ImageMimeType, RevisedPrompt: event.RevisedPrompt, + } + return json.Marshal(wire) +} + +func DecodeGatewayEvent(data []byte) (Event, error) { + var wire gatewayEvent + dec := json.NewDecoder(bytes.NewReader(data)) + dec.DisallowUnknownFields() + if err := dec.Decode(&wire); err != nil { + return Event{}, fmt.Errorf("decode gateway event: %w", err) + } + event := Event{ + Type: wire.Type, Text: wire.Text, Model: wire.Model, ReasoningEffort: wire.ReasoningEffort, + InterjectionID: wire.InterjectionID, InterjectionStatus: wire.InterjectionStatus, + ReasoningItemID: wire.ReasoningItemID, ReasoningEncryptedContent: wire.ReasoningEncryptedContent, + ReasoningKind: wire.ReasoningKind, ReasoningSummaryParts: wire.ReasoningSummaryParts, + ReasoningIndex: wire.ReasoningIndex, ReasoningFinal: wire.ReasoningFinal, + Tool: wire.Tool, ToolCallID: wire.ToolCallID, ToolName: wire.ToolName, ToolInfo: wire.ToolInfo, + ToolArgs: wire.ToolArgs, ToolSuccess: wire.ToolSuccess, ToolOutput: wire.ToolOutput, + ToolDiffs: wire.ToolDiffs, ToolFileChanges: wire.ToolFileChanges, ToolImages: wire.ToolImages, + Use: wire.Use, RetryAttempt: wire.RetryAttempt, RetryMaxAttempts: wire.RetryMaxAttempts, + RetryWaitSecs: wire.RetryWaitSecs, ToolActivity: wire.ToolActivity, + ProviderReplay: wire.ProviderReplay, ImageData: wire.ImageData, + ImageMimeType: wire.ImageMimeType, RevisedPrompt: wire.RevisedPrompt, + } + if wire.Message != nil { + event.Message = *wire.Message + } + if wire.Error != nil { + event.Err = &GatewayError{Code: wire.Error.Code, Message: wire.Error.Message, RequestID: wire.Error.RequestID} + } + return event, nil +} + +// GatewayError is a safe, structured gateway failure. Its message is prepared +// by the gateway and never contains raw upstream bodies, secrets, or paths. +type GatewayError struct { + Code string + Message string + RequestID string +} + +func (e *GatewayError) Error() string { + if e == nil { + return "gateway request failed" + } + if e.Code == "" { + return "gateway: " + e.Message + } + return fmt.Sprintf("gateway %s: %s", e.Code, e.Message) +} + +func EncodeGatewayToolResponse(response ToolExecutionResponse) (json.RawMessage, error) { + response.Result.Diffs = nil + response.Result.Images = nil + response.Result.FileChanges = nil + payload := struct { + Result ToolOutput `json:"result"` + Error string `json:"error,omitempty"` + }{Result: response.Result} + if response.Err != nil { + payload.Error = response.Err.Error() + } + return json.Marshal(payload) +} + +func DecodeGatewayToolResponse(data []byte) (ToolExecutionResponse, error) { + var payload struct { + Result ToolOutput `json:"result"` + Error string `json:"error,omitempty"` + } + dec := json.NewDecoder(bytes.NewReader(data)) + dec.DisallowUnknownFields() + if err := dec.Decode(&payload); err != nil { + return ToolExecutionResponse{}, fmt.Errorf("decode gateway tool response: %w", err) + } + response := ToolExecutionResponse{Result: payload.Result} + if payload.Error != "" { + response.Err = fmt.Errorf("%s", payload.Error) + } + return response, nil +} + +func CapabilitiesToGatewayProtocol(c Capabilities) protocol.Capabilities { + return protocol.Capabilities{ + NativeWebSearch: c.NativeWebSearch, NativeWebFetch: c.NativeWebFetch, + ToolCalls: c.ToolCalls, SupportsToolChoice: c.SupportsToolChoice, + ManagesOwnContext: c.ManagesOwnContext, InlineToolLoop: c.InlineToolLoop, + OrderedInlineToolEvents: c.OrderedInlineToolEvents, + } +} + +func capabilitiesFromProtocol(c protocol.Capabilities) Capabilities { + return Capabilities{ + NativeWebSearch: c.NativeWebSearch, NativeWebFetch: c.NativeWebFetch, + ToolCalls: c.ToolCalls, SupportsToolChoice: c.SupportsToolChoice, + ManagesOwnContext: c.ManagesOwnContext, InlineToolLoop: c.InlineToolLoop, + OrderedInlineToolEvents: c.OrderedInlineToolEvents, + } +} diff --git a/internal/llm/gateway_wire_test.go b/internal/llm/gateway_wire_test.go new file mode 100644 index 000000000..584f9d3e1 --- /dev/null +++ b/internal/llm/gateway_wire_test.go @@ -0,0 +1,100 @@ +package llm + +import ( + "encoding/json" + "errors" + "strings" + "testing" + + "github.com/samsaffron/term-llm/internal/gateway/protocol" +) + +func TestGatewayWireRequestRoundTripAndOmitsLocalState(t *testing.T) { + req := Request{ + Model: "vision-model", SessionID: "session-1", WorkingDir: "/satellite/private", + ApprovalTranscriptPrefix: []Message{UserText("private approval")}, + AllowedTools: []string{"read_file"}, AllowedToolsPresent: true, Debug: true, DebugRaw: true, + Messages: []Message{ + {Role: RoleUser, Parts: []Part{ + {Type: PartText, Text: "hello", ReasoningContent: "summary", ReasoningKind: ReasoningKindSummary}, + {Type: PartImage, ImagePath: "/private/image.png", ImageData: &ToolImageData{MediaType: "image/png", Base64: "aW1hZ2U="}}, + {Type: PartFile, FilePath: "/private/file.pdf", FileData: &ToolFileData{MediaType: "application/pdf", Base64: "ZmlsZQ==", Filename: "file.pdf"}}, + }}, + {Role: RoleAssistant, Parts: []Part{ + {Type: PartToolCall, ToolCall: &ToolCall{ID: "c1", Name: "view", Arguments: json.RawMessage(`{"x":1}`), Caller: "programmatic", ThoughtSig: []byte("sig")}}, + {Type: PartProviderReplay, ProviderReplay: &ProviderReplayItem{Raw: json.RawMessage(`{"type":"opaque"}`)}}, + }}, + {Role: RoleTool, Parts: []Part{{Type: PartToolResult, ToolResult: &ToolResult{ID: "c1", Name: "view", Content: "ok", ContentParts: []ToolContentPart{{Type: ToolContentPartText, Text: "ok"}, {Type: ToolContentPartImageData, ImageData: &ToolImageData{MediaType: "image/png", Base64: "aQ=="}}}, Images: []string{"/private/result.png"}}}}}, + }, + Tools: []ToolSpec{{Name: "view", Description: "view", Schema: map[string]any{"type": "object"}, Strict: true, AllowedCallers: []string{"programmatic"}, OutputSchema: map[string]any{"type": "string"}}}, + ToolChoice: ToolChoice{Mode: ToolChoiceName, Name: "view"}, ParallelToolCalls: true, + Search: true, ForceExternalSearch: true, ReasoningEffort: "high", MaxOutputTokens: 42, + Temperature: 0, TemperatureSet: true, TopP: .5, TopPSet: true, + Responses: &ResponsesOptions{ReasoningMode: "summary", MultiAgent: MultiAgentOptions{Enabled: true, EnabledSet: true, MaxConcurrentSubagents: 2}}, + } + wire, err := EncodeGatewayRequest(req) + if err != nil { + t.Fatal(err) + } + text := string(wire) + for _, forbidden := range []string{"WorkingDir", "working_dir", "/satellite/private", "/private/image.png", "/private/file.pdf", "/private/result.png", "private approval", "AllowedTools", "DebugRaw"} { + if strings.Contains(text, forbidden) { + t.Fatalf("wire request leaked forbidden %q: %s", forbidden, text) + } + } + got, err := DecodeGatewayRequest(wire) + if err != nil { + t.Fatal(err) + } + if got.Model != req.Model || got.SessionID != req.SessionID || len(got.Messages) != 3 || !got.Search || !got.TemperatureSet || got.WorkingDir != "" { + t.Fatalf("round trip mismatch: %+v", got) + } + if got.Messages[0].Parts[1].ImageData == nil || got.Messages[0].Parts[2].FileData == nil { + t.Fatalf("vision/file data lost: %+v", got.Messages[0].Parts) + } + if got.Messages[2].Parts[0].ToolResult.Images != nil { + t.Fatalf("local result paths crossed wire: %+v", got.Messages[2].Parts[0].ToolResult) + } +} + +func TestGatewayWireEventRoundTripAllFields(t *testing.T) { + event := Event{ + Type: EventReasoningDelta, Text: "reason", Model: "m", ReasoningEffort: "high", + ReasoningItemID: "r1", ReasoningEncryptedContent: "sealed", ReasoningKind: ReasoningKindSummary, + ReasoningSummaryParts: []string{"a", "b"}, ReasoningIndex: 2, ReasoningFinal: true, + Tool: &ToolCall{ID: "c", Name: "tool", Arguments: json.RawMessage(`{"a":1}`), ThoughtSig: []byte("sig")}, + ToolCallID: "c", ToolName: "tool", ToolInfo: "info", ToolArgs: json.RawMessage(`{"a":1}`), + ToolSuccess: true, ToolOutput: "out", ToolDiffs: []DiffData{{File: "f", Old: "a", New: "b", Line: 1}}, + ToolFileChanges: []FileChange{{Path: "f", Kind: "modify", Adds: 1}}, ToolImages: []string{"image"}, + Use: &Usage{InputTokens: 1, OutputTokens: 2, CachedInputTokens: 3, CacheWriteTokens: 4, ReasoningTokens: 5}, + RetryAttempt: 1, RetryMaxAttempts: 3, RetryWaitSecs: .25, + ToolActivity: &ToolActivity{ID: "a", Name: "search", Status: ToolActivityCompleted}, + ProviderReplay: &ProviderReplayItem{Raw: json.RawMessage(`{"opaque":true}`)}, + ImageData: []byte("image"), ImageMimeType: "image/png", RevisedPrompt: "better", + } + wire, err := EncodeGatewayEvent(event, &protocol.Error{Code: "provider_upstream_failure", Message: "safe error"}) + if err != nil { + t.Fatal(err) + } + if strings.Contains(string(wire), `"message":{}`) { + t.Fatalf("empty event message was serialized: %s", wire) + } + got, err := DecodeGatewayEvent(wire) + if err != nil { + t.Fatal(err) + } + if got.Type != event.Type || got.Text != event.Text || got.Tool == nil || got.Tool.Name != "tool" || got.Use == nil || got.Use.ReasoningTokens != 5 || got.ProviderReplay == nil || got.Err == nil || !strings.Contains(got.Err.Error(), "safe error") || string(got.ImageData) != "image" { + t.Fatalf("event round trip mismatch: %+v", got) + } + var gatewayErr *GatewayError + if !errors.As(got.Err, &gatewayErr) || gatewayErr.Code != "provider_upstream_failure" { + t.Fatalf("structured gateway error = %#v", got.Err) + } +} + +func TestDecodeGatewayRequestRejectsUnknownTopLevelField(t *testing.T) { + _, err := DecodeGatewayRequest([]byte(`{"model":"m","working_dir":"/forged"}`)) + if err == nil || !strings.Contains(err.Error(), "unknown field") { + t.Fatalf("got %v, want controlled unknown-field error", err) + } +} diff --git a/internal/llm/models.go b/internal/llm/models.go index b0b3d23e1..7bb54c22a 100644 --- a/internal/llm/models.go +++ b/internal/llm/models.go @@ -797,12 +797,30 @@ func GetImageProviderNames() []string { func GetProviderCompletions(toComplete string, isImage bool, cfg *config.Config) []string { var providerNames []string var getModelIDs func(string) []string + remoteModels := make(map[string][]string) + if !isImage && cfg != nil && cfg.Gateway.Enabled() { + if catalog, err := LoadGatewayCatalogCacheOnly(cfg); err == nil { + for _, entry := range catalog.Providers { + providerNames = append(providerNames, entry.Key) + for _, model := range entry.Models { + remoteModels[entry.Key] = append(remoteModels[entry.Key], model.ID) + } + } + } + } if isImage { providerNames = GetImageProviderNames() getModelIDs = func(p string) []string { return ImageProviderModels[p] } } else { - providerNames = GetProviderNames(cfg) + if cfg != nil && cfg.Gateway.Enabled() { + providerNames = append(providerNames, "debug") + providerNames = append(providerNames, cfg.ExplicitProviderNames()...) + providerNames = append(providerNames, cfg.Gateway.LocalProviders...) + } else { + providerNames = append(providerNames, GetProviderNames(cfg)...) + } + providerNames = dedupeCompletionStrings(providerNames) getModelIDs = ProviderModelIDs } @@ -814,6 +832,12 @@ func GetProviderCompletions(toComplete string, isImage bool, cfg *config.Config) // Get models for completion var models []string + if cfg != nil && !cfg.IsLocalProvider(provider) { + models = append(models, remoteModels[provider]...) + if cfg.Gateway.Enabled() && len(models) == 0 { + return nil + } + } // Check if config has a models list for this provider var configModels []string @@ -825,7 +849,7 @@ func GetProviderCompletions(toComplete string, isImage bool, cfg *config.Config) } } - if len(configModels) > 0 { + if len(models) == 0 && len(configModels) > 0 { // Use config-defined models list, plus configured model (deduped) seen := make(map[string]bool) if configModel != "" { @@ -838,7 +862,8 @@ func GetProviderCompletions(toComplete string, isImage bool, cfg *config.Config) seen[m] = true } } - } else { + } + if len(models) == 0 { // Resolve provider type, including custom aliases (e.g., "acme" → "venice") providerType := resolveProviderType(provider) @@ -893,3 +918,17 @@ func GetProviderCompletions(toComplete string, isImage bool, cfg *config.Config) } return completions } + +func dedupeCompletionStrings(values []string) []string { + seen := make(map[string]bool, len(values)) + out := make([]string, 0, len(values)) + for _, value := range values { + if value == "" || seen[value] { + continue + } + seen[value] = true + out = append(out, value) + } + sort.Strings(out) + return out +} diff --git a/internal/llm/retry.go b/internal/llm/retry.go index 304e70284..91120bc08 100644 --- a/internal/llm/retry.go +++ b/internal/llm/retry.go @@ -43,8 +43,12 @@ type RetryProvider struct { config RetryConfig } -// WrapWithRetry wraps a provider with retry logic. +// WrapWithRetry wraps a provider with retry logic. Reconfiguring an existing +// RetryProvider replaces its policy instead of nesting retry loops. func WrapWithRetry(p Provider, config RetryConfig) Provider { + if existing, ok := p.(*RetryProvider); ok { + return &RetryProvider{inner: existing.inner, config: normalizeRetryConfig(config)} + } return &RetryProvider{inner: p, config: normalizeRetryConfig(config)} } @@ -200,12 +204,18 @@ func (r *RetryProvider) ListModels(ctx context.Context) ([]ModelInfo, error) { func (r *RetryProvider) Stream(ctx context.Context, req Request) (Stream, error) { config := normalizeRetryConfig(r.config) return newEventStream(ctx, func(ctx context.Context, send eventSender) error { - _, err := retryCall(ctx, config, func() (struct{}, error) { - stream, err := r.inner.Stream(ctx, req) + retryCtx := ctx + cancel := func() {} + if config.MaxElapsedTime > 0 { + retryCtx, cancel = context.WithTimeout(ctx, config.MaxElapsedTime) + } + defer cancel() + _, err := retryCall(retryCtx, config, func() (struct{}, error) { + stream, err := r.inner.Stream(retryCtx, req) if err != nil { return struct{}{}, err } - return struct{}{}, r.forwardAttempt(ctx, stream, send) + return struct{}{}, r.forwardAttempt(retryCtx, stream, send) }, func(info retryInfo) error { // Emit retry event so UI can show progress. RetryMaxAttempts==0 means // time-budgeted retry with no fixed attempt ceiling. diff --git a/internal/search/factory.go b/internal/search/factory.go index 05e94e6f2..a1c72b75a 100644 --- a/internal/search/factory.go +++ b/internal/search/factory.go @@ -9,6 +9,9 @@ import ( // NewSearcher creates a Searcher based on the config. // Returns Exa MCP as the default if no provider is specified. func NewSearcher(cfg *config.Config) (Searcher, error) { + if cfg != nil && cfg.Gateway.Enabled() && cfg.Gateway.RouteSearch() { + return NewGatewayClient(cfg.Gateway) + } provider := cfg.Search.Provider if provider == "" { provider = config.DefaultSearchProvider diff --git a/internal/search/gateway.go b/internal/search/gateway.go new file mode 100644 index 000000000..9f79469da --- /dev/null +++ b/internal/search/gateway.go @@ -0,0 +1,109 @@ +package search + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net" + "net/http" + "strings" + "time" + + "github.com/samsaffron/term-llm/internal/config" + "github.com/samsaffron/term-llm/internal/gateway/protocol" +) + +type GatewayClient struct { + baseURL string + token string + client *http.Client +} + +func NewGatewayClient(cfg config.GatewayConfig) (*GatewayClient, error) { + if !cfg.Enabled() { + return nil, fmt.Errorf("gateway is not configured") + } + token, err := cfg.ResolveToken() + if err != nil { + return nil, err + } + timeout := config.DefaultGatewayResponseTimeout + if strings.TrimSpace(cfg.ResponseTimeout) != "" { + timeout = cfg.ResponseTimeout + } + d, err := time.ParseDuration(timeout) + if err != nil || d <= 0 { + d = 30 * time.Second + } + connectTimeout := config.DefaultGatewayConnectTimeout + if strings.TrimSpace(cfg.ConnectTimeout) != "" { + connectTimeout = cfg.ConnectTimeout + } + connect, err := time.ParseDuration(connectTimeout) + if err != nil || connect <= 0 { + connect = 2 * time.Second + } + transport := &http.Transport{ + Proxy: http.ProxyFromEnvironment, + DialContext: (&net.Dialer{Timeout: connect, KeepAlive: 30 * time.Second}).DialContext, + ResponseHeaderTimeout: d, + } + return &GatewayClient{baseURL: strings.TrimRight(cfg.URL, "/"), token: token, client: &http.Client{Transport: transport, Timeout: d}}, nil +} + +func (c *GatewayClient) Search(ctx context.Context, query string, maxResults int) ([]Result, error) { + var response protocol.SearchResponse + if err := c.post(ctx, "/g1/search", protocol.SearchRequest{Version: protocol.Version, Query: query, MaxResults: maxResults}, &response); err != nil { + return nil, err + } + results := make([]Result, 0, len(response.Results)) + for _, result := range response.Results { + results = append(results, Result{Title: result.Title, URL: result.URL, Snippet: result.Snippet}) + } + return results, nil +} + +func (c *GatewayClient) FetchURL(ctx context.Context, rawURL string) (string, error) { + var response protocol.FetchResponse + if err := c.post(ctx, "/g1/fetch", protocol.FetchRequest{Version: protocol.Version, URL: rawURL}, &response); err != nil { + return "", err + } + return response.Content, nil +} + +func (c *GatewayClient) post(ctx context.Context, endpoint string, payload, target any) error { + data, err := json.Marshal(payload) + if err != nil { + return err + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.baseURL+endpoint, bytes.NewReader(data)) + if err != nil { + return err + } + req.Header.Set("Authorization", "Bearer "+c.token) + req.Header.Set(protocol.VersionHeader, "1") + req.Header.Set("Content-Type", "application/json") + resp, err := c.client.Do(req) + if err != nil { + return fmt.Errorf("gateway %s unavailable; check gateway URL/network and retry: %w", strings.TrimPrefix(endpoint, "/g1/"), err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + data, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20)) + var wire protocol.Error + if json.Unmarshal(data, &wire) == nil && wire.Message != "" { + return fmt.Errorf("gateway %s: %s", wire.Code, wire.Message) + } + return fmt.Errorf("gateway HTTP %d", resp.StatusCode) + } + dec := json.NewDecoder(io.LimitReader(resp.Body, config.DefaultGatewayMaxResponseBytes)) + dec.DisallowUnknownFields() + if err := dec.Decode(target); err != nil { + return fmt.Errorf("decode gateway response: %w", err) + } + return nil +} + +var _ Searcher = (*GatewayClient)(nil) diff --git a/internal/search/gateway_test.go b/internal/search/gateway_test.go new file mode 100644 index 000000000..d180673d3 --- /dev/null +++ b/internal/search/gateway_test.go @@ -0,0 +1,22 @@ +package search + +import ( + "testing" + + "github.com/samsaffron/term-llm/internal/config" +) + +func TestGatewayExplicitFalsePreservesLocalSearch(t *testing.T) { + route := false + cfg := &config.Config{ + Gateway: config.GatewayConfig{URL: "https://gateway.invalid", Search: &route}, + Search: config.SearchConfig{Provider: "duckduckgo"}, + } + searcher, err := NewSearcher(cfg) + if err != nil { + t.Fatal(err) + } + if _, ok := searcher.(*DuckDuckGoLite); !ok { + t.Fatalf("explicit gateway.search false created %T, want local DuckDuckGo", searcher) + } +} diff --git a/internal/tools/config.go b/internal/tools/config.go index b2018344d..e815fe72f 100644 --- a/internal/tools/config.go +++ b/internal/tools/config.go @@ -389,13 +389,17 @@ func optionalToolConfig(configs []*ToolConfig) *ToolConfig { } // ParseToolsFlag parses a comma-separated list of tool names. -// Special values: "all" or "*" expand to all available tools. +// Special values: "all" or "*" expand to all available tools; "none" disables +// every tool, including tools enabled in configuration. func ParseToolsFlag(value string) []string { if value == "" { return nil } - // Handle "all" or "*" to enable all tools trimmed := strings.TrimSpace(value) + if strings.EqualFold(trimmed, "none") { + return []string{} + } + // Handle "all" or "*" to enable all tools if trimmed == "all" || trimmed == "*" { return StandardToolNames() } diff --git a/internal/tools/gateway_tools_test.go b/internal/tools/gateway_tools_test.go new file mode 100644 index 000000000..706770c0f --- /dev/null +++ b/internal/tools/gateway_tools_test.go @@ -0,0 +1,12 @@ +package tools + +import "testing" + +func TestParseToolsFlagNoneDisablesAllTools(t *testing.T) { + if got := ParseToolsFlag("none"); got == nil || len(got) != 0 { + t.Fatalf("ParseToolsFlag(none) = %#v, want non-nil empty list", got) + } + if got := ParseToolsFlag(" NONE "); got == nil || len(got) != 0 { + t.Fatalf("ParseToolsFlag(NONE) = %#v, want non-nil empty list", got) + } +} diff --git a/internal/usage/gateway_test.go b/internal/usage/gateway_test.go new file mode 100644 index 000000000..901e2e8a6 --- /dev/null +++ b/internal/usage/gateway_test.go @@ -0,0 +1,49 @@ +package usage + +import ( + "encoding/json" + "os" + "path/filepath" + "strings" + "testing" + "time" +) + +func TestGatewayTrackedSatelliteUsageVisibleWithoutAggregateDoubleCount(t *testing.T) { + entry := UsageEntry{Timestamp: time.Now().UTC(), Provider: ProviderTermLLM, Model: "remote-model", InputTokens: 10, TrackedExternallyBy: ProviderGateway} + result := LoadResult{Entries: []UsageEntry{entry}} + if got := result.Filter(FilterOptions{}); len(got) != 0 { + t.Fatalf("aggregate included gateway-tracked satellite copy: %+v", got) + } + visible := result.Filter(FilterOptions{Provider: ProviderTermLLM, IncludeExternal: true}) + if len(visible) != 1 || visible[0].TrackedExternallyBy != ProviderGateway { + t.Fatalf("explicit local usage view hid gateway-tracked copy: %+v", visible) + } + allVisible := result.Filter(FilterOptions{IncludeExternal: true}) + if len(allVisible) != 1 || allVisible[0].TrackedExternallyBy != ProviderGateway { + t.Fatalf("include-external did not affect aggregate view: %+v", allVisible) + } +} + +func TestUsageLoggerNormalizesTimestampToUTC(t *testing.T) { + logger := &Logger{baseDir: t.TempDir()} + local := time.Date(2026, 8, 1, 12, 0, 0, 0, time.FixedZone("local", 10*60*60)) + if err := logger.Log(LogEntry{Timestamp: local, Model: "m", Provider: "p", InputTokens: 1}); err != nil { + t.Fatal(err) + } + files, err := filepath.Glob(filepath.Join(logger.baseDir, "*.jsonl")) + if err != nil || len(files) != 1 { + t.Fatalf("usage files = %v, %v", files, err) + } + data, err := os.ReadFile(files[0]) + if err != nil { + t.Fatal(err) + } + var entry LogEntry + if err := json.Unmarshal([]byte(strings.TrimSpace(string(data))), &entry); err != nil { + t.Fatal(err) + } + if entry.Timestamp.Location() != time.UTC || entry.Timestamp.Hour() != 2 { + t.Fatalf("usage timestamp = %s (%s), want UTC", entry.Timestamp, entry.Timestamp.Location()) + } +} diff --git a/internal/usage/logger.go b/internal/usage/logger.go index 71367b94b..7a4651b08 100644 --- a/internal/usage/logger.go +++ b/internal/usage/logger.go @@ -53,6 +53,7 @@ func NewLogger() *Logger { func (l *Logger) Log(entry LogEntry) error { l.mu.Lock() defer l.mu.Unlock() + entry.Timestamp = entry.Timestamp.UTC() // Ensure directory exists if err := os.MkdirAll(l.baseDir, 0755); err != nil { diff --git a/internal/usage/types.go b/internal/usage/types.go index 9558a68ef..1dd81dab8 100644 --- a/internal/usage/types.go +++ b/internal/usage/types.go @@ -7,6 +7,7 @@ const ( ProviderClaudeCode = "claude-code" ProviderGeminiCLI = "gemini-cli" ProviderTermLLM = "term-llm" + ProviderGateway = "gateway" ) // UsageEntry represents a single usage event from any provider @@ -95,17 +96,11 @@ func (r LoadResult) Filter(opts FilterOptions) []UsageEntry { if !opts.Until.IsZero() && e.Timestamp.After(opts.Until) { continue } - // Handle term-llm externally-tracked entries - if e.Provider == ProviderTermLLM && e.TrackedExternallyBy != "" { - // When showing all providers (no filter), exclude externally-tracked term-llm entries - // to avoid double-counting with the external provider's data - if opts.Provider == "" { - continue - } - // When filtering to term-llm specifically, only include if IncludeExternal is set - if opts.Provider == ProviderTermLLM && !opts.IncludeExternal { - continue - } + // Externally tracked term-llm entries are hidden by default to avoid + // aggregate double-counting. --include-external means include them in both + // the all-provider view and an explicit term-llm view. + if e.Provider == ProviderTermLLM && e.TrackedExternallyBy != "" && !opts.IncludeExternal { + continue } result = append(result, e) } diff --git a/ops/gateway-compose.yaml b/ops/gateway-compose.yaml new file mode 100644 index 000000000..8d7738da1 --- /dev/null +++ b/ops/gateway-compose.yaml @@ -0,0 +1,66 @@ +# Private-network inference gateway example. +# Supply or build a `term-llm:latest` image, copy the adjacent example configs, +# and place provider/client secrets in ./secrets before starting. +services: + gateway: + image: term-llm:latest + command: + - gateway + - serve + - --listen + - 0.0.0.0:8787 + - --state-dir + - /var/lib/term-llm-gateway + volumes: + - gateway-state:/var/lib/term-llm-gateway + - ./gateway-config.yaml:/home/agent/.config/term-llm/config.yaml:ro + secrets: + - anthropic_api_key + networks: [provider-plane] + restart: unless-stopped + healthcheck: + test: ["CMD", "term-llm", "gateway", "health", "http://127.0.0.1:8787/g1/health"] + interval: 10s + timeout: 5s + retries: 5 + start_period: 5s + # Deliberately no ports: the gateway is reachable only on provider-plane. + + satellite: + image: term-llm:latest + command: [serve, web, --host, 0.0.0.0, --port, "8080"] + volumes: + - satellite-state:/home/agent + - ./satellite-config.yaml:/home/agent/.config/term-llm/config.yaml:ro + secrets: + - gateway_client_token + networks: + - provider-plane + - ui-plane + ports: + - "127.0.0.1:8080:8080" + depends_on: + gateway: + condition: service_healthy + restart: unless-stopped + healthcheck: + test: ["CMD", "term-llm", "gateway", "health", "http://127.0.0.1:8080/healthz"] + interval: 10s + timeout: 5s + retries: 5 + start_period: 10s + +networks: + provider-plane: + internal: true + ui-plane: {} + +volumes: + gateway-state: {} + satellite-state: {} + +secrets: + anthropic_api_key: + file: ./secrets/anthropic_api_key + gateway_client_token: + file: ./secrets/gateway_client_token diff --git a/ops/gateway-config.yaml b/ops/gateway-config.yaml new file mode 100644 index 000000000..c33e5fb7c --- /dev/null +++ b/ops/gateway-config.yaml @@ -0,0 +1,10 @@ +# Central gateway provider configuration. Provider credentials stay here. +default_provider: anthropic +providers: + anthropic: + model: claude-sonnet-4-6 + # Lazy resolution happens only on the gateway. + api_key: "$(cat /run/secrets/anthropic_api_key)" +search: + provider: exa_mcp + fetch_provider: jina diff --git a/ops/gateway_compose_test.go b/ops/gateway_compose_test.go new file mode 100644 index 000000000..ad5c0e407 --- /dev/null +++ b/ops/gateway_compose_test.go @@ -0,0 +1,46 @@ +package ops + +import ( + "bytes" + "os" + "testing" + + "gopkg.in/yaml.v3" +) + +func TestGatewayComposeExampleParsesAndHasHealthGating(t *testing.T) { + data, err := os.ReadFile("gateway-compose.yaml") + if err != nil { + t.Fatal(err) + } + var document struct { + Services map[string]struct { + Image string `yaml:"image"` + Healthcheck struct { + Test []string `yaml:"test"` + Interval string `yaml:"interval"` + Timeout string `yaml:"timeout"` + Retries int `yaml:"retries"` + } `yaml:"healthcheck"` + DependsOn map[string]struct { + Condition string `yaml:"condition"` + } `yaml:"depends_on"` + } `yaml:"services"` + } + dec := yaml.NewDecoder(bytes.NewReader(data)) + if err := dec.Decode(&document); err != nil { + t.Fatalf("compose YAML is invalid: %v", err) + } + for _, name := range []string{"gateway", "satellite"} { + service, ok := document.Services[name] + if !ok || service.Image == "" { + t.Fatalf("missing %s service/image", name) + } + if len(service.Healthcheck.Test) < 4 || service.Healthcheck.Test[0] != "CMD" || service.Healthcheck.Interval == "" || service.Healthcheck.Timeout == "" || service.Healthcheck.Retries <= 0 { + t.Fatalf("%s healthcheck is incomplete: %+v", name, service.Healthcheck) + } + } + if got := document.Services["satellite"].DependsOn["gateway"].Condition; got != "service_healthy" { + t.Fatalf("satellite gateway dependency condition = %q", got) + } +} diff --git a/ops/satellite-config.yaml b/ops/satellite-config.yaml new file mode 100644 index 000000000..7d9b141c5 --- /dev/null +++ b/ops/satellite-config.yaml @@ -0,0 +1,11 @@ +# Satellite configuration: no provider or search credentials. +default_provider: anthropic +gateway: + url: http://gateway:8787 + token_file: /run/secrets/gateway_client_token + required: true + search: true + fetch: true + catalog_ttl: 15m + connect_timeout: 2s + response_timeout: 5s From 5b5390108390c49fda9619ff77e78ffc9fce9531 Mon Sep 17 00:00:00 2001 From: Jarvis Date: Sat, 1 Aug 2026 19:06:25 +1000 Subject: [PATCH 2/3] Harden gateway session resumption equivalence --- docs-site/content/guides/inference-gateway.md | 16 +- internal/gateway/equivalence_test.go | 459 ++++++++++++++++++ internal/gateway/execution.go | 4 +- internal/gateway/integration_test.go | 7 +- internal/gateway/seal.go | 31 +- internal/gateway/security_test.go | 23 +- internal/gateway/server.go | 12 +- internal/llm/chatgpt.go | 15 +- internal/llm/chatgpt_test.go | 212 ++++++++ internal/llm/gateway_provider.go | 14 +- internal/llm/gateway_wire.go | 24 +- internal/llm/gateway_wire_test.go | 12 +- 12 files changed, 788 insertions(+), 41 deletions(-) create mode 100644 internal/gateway/equivalence_test.go diff --git a/docs-site/content/guides/inference-gateway.md b/docs-site/content/guides/inference-gateway.md index 0fefaf4bd..a8a6b0630 100644 --- a/docs-site/content/guides/inference-gateway.md +++ b/docs-site/content/guides/inference-gateway.md @@ -30,10 +30,24 @@ Important invariants: 1. Provider API keys, OAuth tokens, CLI homes, and provider configuration are never returned by catalog or inference endpoints. 2. A satellite `WorkingDir` is not a wire field. Each gateway request gets a new empty directory under the gateway state-owned `runs/` root. Finished runs are removed, and stale gateway-prefixed directories are scavenged safely at startup. 3. The gateway has no satellite or external-client tool registry. Normal `/g1` tool calls return to the satellite engine. Inline CLI-provider calls on `/g1` use an authenticated callback POST and block the gateway provider until the satellite result arrives. `/v1/responses` accepts only client-defined function tools, returns function calls to the client, and never executes them; tool-bearing requests to incompatible inline-loop CLI providers are rejected before provider startup. -4. Provider resume state is opaque to satellites and authenticated with an AES-GCM gateway key. It is bound to both client ID and provider key; tampered or cross-client state is rejected. +4. Provider resume state is opaque to satellites and authenticated with an AES-GCM gateway key. It is bound to protocol version, authenticated client ID, provider key, and the satellite `SessionID`; tampered, cross-client, cross-provider, or cross-session state is rejected. Requests with an empty `SessionID` are deliberately stateless and neither import nor export sealed provider state. 5. Each request is recorded centrally with client, provider key, model, request/session IDs, token counters, outcome, and locally calculable cost. 6. The gateway performs provider retries. By default it makes at most three upstream attempts within 20 seconds; `--upstream-retry-attempts` and `--upstream-retry-elapsed` tighten or extend those bounds. Request cancellation and stream idle deadlines still win. The satellite gateway transport and engine do not add a second retry loop. +### Session continuity and direct-provider equivalence + +For an ordinary satellite session, the satellite remains the source of truth for `SessionID` and the complete transcript. It stores provider state under that session plus provider key, includes the same `SessionID` and provider-facing history on later turns, and can export/import the `GatewayProvider`'s opaque sealed blob when reconstructing a runtime. The gateway creates a fresh central provider for each request. Providers that implement provider-state export/import round-trip that state through the session-bound sealed blob; providers that do not are reconstructed from the complete transcript. + +After intentional boundary normalization, the central provider receives the same model-semantic `llm.Request` as a direct provider: model and session identity; system/developer/user/assistant/tool history; tool definitions and choices; tool-call and tool-result continuation data; cache anchors; reasoning summaries, encrypted reasoning, and opaque provider replay; search controls; Responses options; sampling, output-token, service-tier, and ephemeral/continuation controls. This equivalence has deliberately narrow exceptions: + +- `WorkingDir` is replaced with a fresh empty gateway-owned run directory. +- Satellite image/file paths are removed while inline image/file data is retained. Tool-result diff/image paths, legacy display strings, formatted `ToolInfo`, provider tool-activity display records, skill-activation provenance, and persisted UI identity/segment fields are omitted. Expanded developer instructions, actual tool arguments/results, multimodal result data, caller provenance, thought signatures, reasoning, and provider replay remain. +- Approval-only transcripts/roles, request-scoped execution filters (`AllowedTools`), satellite engine controls (`MaxTurns` and `ToolMap`), and debug flags are local and do not cross the provider boundary. The satellite engine consumes these before or around provider execution; omitting them cannot change the upstream model payload. + +ChatGPT transport behavior needs one qualification. ChatGPT's HTTP/SSE path sends full history on every turn in both direct and gateway modes, so second-turn upstream request payloads are equivalent. A long-lived **direct WebSocket** provider may instead reuse its connection-local `previous_response_id` and send only the new continuation suffix. The gateway recreates the central provider (and therefore its WebSocket connection) per request, so it sends the full transcript without `previous_response_id`. Both forms are semantically equivalent, and rejection of a direct WebSocket parent falls back to the same full-history request. They are not strictly transport-, cache-, or latency-equivalent. + +Gateway restart and satellite session restoration remain correct when the client database and `state.key` are retained: the satellite replays its full transcript and restores any exportable provider state from the sealed blob. A restart does **not** retain ChatGPT's connection-local WebSocket optimization, so the next request uses full history and may have different cache/latency characteristics even though conversation semantics are preserved. + Run the gateway only on a trusted private network. For traffic that is not already protected by a private overlay, service mesh, or TLS reverse proxy, configure the built-in TLS listener with both `--tls-cert` and `--tls-key`. Bearer credentials authenticate access but plaintext HTTP does not encrypt them. ## Start a gateway diff --git a/internal/gateway/equivalence_test.go b/internal/gateway/equivalence_test.go new file mode 100644 index 000000000..5f45455be --- /dev/null +++ b/internal/gateway/equivalence_test.go @@ -0,0 +1,459 @@ +package gateway + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http/httptest" + "path/filepath" + "reflect" + "strings" + "sync" + "testing" + "time" + + "github.com/samsaffron/term-llm/internal/config" + "github.com/samsaffron/term-llm/internal/llm" +) + +type equivalenceProviderRecord struct { + Request llm.Request + Imported string +} + +type equivalenceProviderRecorder struct { + mu sync.Mutex + records []equivalenceProviderRecord +} + +func (r *equivalenceProviderRecorder) add(req llm.Request, imported string) { + r.mu.Lock() + defer r.mu.Unlock() + r.records = append(r.records, equivalenceProviderRecord{Request: req, Imported: imported}) +} + +func (r *equivalenceProviderRecorder) snapshot() []equivalenceProviderRecord { + r.mu.Lock() + defer r.mu.Unlock() + return append([]equivalenceProviderRecord(nil), r.records...) +} + +type equivalenceStateProvider struct { + recorder *equivalenceProviderRecorder + + mu sync.Mutex + imported string + exported string +} + +func (*equivalenceStateProvider) Name() string { return "equivalence-state" } +func (*equivalenceStateProvider) Credential() string { return "mock" } +func (*equivalenceStateProvider) Capabilities() llm.Capabilities { + return llm.Capabilities{ToolCalls: true} +} + +func (p *equivalenceStateProvider) ImportProviderState(data []byte) error { + p.mu.Lock() + defer p.mu.Unlock() + p.imported = string(data) + return nil +} + +func (p *equivalenceStateProvider) ExportProviderState() ([]byte, bool) { + p.mu.Lock() + defer p.mu.Unlock() + if p.exported == "" { + return nil, false + } + return []byte(p.exported), true +} + +func (p *equivalenceStateProvider) Stream(ctx context.Context, req llm.Request) (llm.Stream, error) { + p.mu.Lock() + imported := p.imported + wantState := "state:" + req.SessionID + if imported != "" && imported != wantState { + p.mu.Unlock() + return nil, fmt.Errorf("provider state %q does not match session %q", imported, req.SessionID) + } + p.exported = wantState + p.mu.Unlock() + p.recorder.add(req, imported) + return &oneEventStream{ctx: ctx, event: llm.Event{Type: llm.EventTextDelta, Text: "ok"}}, nil +} + +func installReconstructedProviderFactory(fixture *gatewayFixture, recorder *equivalenceProviderRecorder) { + fixture.gateway.cfg.ProviderFactory = func(*config.Config, string, string) (llm.Provider, error) { + return &equivalenceStateProvider{recorder: recorder}, nil + } +} + +func TestGatewayMultiTurnTranscriptToolReasoningAndPersistedState(t *testing.T) { + recorder := &equivalenceProviderRecorder{} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, &equivalenceStateProvider{recorder: recorder}, time.Second) + installReconstructedProviderFactory(fixture, recorder) + + firstSatellite, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + firstMessages := []llm.Message{llm.UserText("first question")} + stream, err := firstSatellite.Stream(t.Context(), llm.Request{ + Model: "model-a", SessionID: "session-a", Messages: firstMessages, + }) + if err != nil { + t.Fatal(err) + } + collectStream(t, stream) + + persisted, ok := firstSatellite.ExportProviderState() + if !ok || strings.Contains(string(persisted), "state:session-a") { + t.Fatalf("exported gateway state = %q, %t; want opaque sealed state", persisted, ok) + } + + // Reconstruct the satellite provider, as a resumed runtime does after loading + // ProviderState from its session database. + resumedSatellite, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + if err := resumedSatellite.ImportProviderState(persisted); err != nil { + t.Fatal(err) + } + + replayRaw := json.RawMessage(`{"type":"reasoning","id":"rs_1","encrypted_content":"sealed-reasoning"}`) + secondMessages := []llm.Message{ + llm.UserText("first question"), + {Role: llm.RoleAssistant, Parts: []llm.Part{ + {Type: llm.PartText, Text: "I need a tool", ReasoningContent: "summary", ReasoningSummaryParts: []string{"summary"}, ReasoningItemID: "rs_1", ReasoningEncryptedContent: "sealed-reasoning", ReasoningKind: llm.ReasoningKindSummary}, + {Type: llm.PartProviderReplay, ProviderReplay: &llm.ProviderReplayItem{Raw: replayRaw}}, + {Type: llm.PartToolCall, ToolCall: &llm.ToolCall{ID: "call_1", Name: "lookup", Arguments: json.RawMessage(`{"q":"term-llm"}`), Caller: "programmatic", ToolInfo: "(/satellite/private)", ThoughtSig: []byte("thought")}}, + }}, + llm.ToolResultMessageFromOutput("call_1", "lookup", llm.ToolOutput{Content: "tool answer", ContentParts: []llm.ToolContentPart{{Type: llm.ToolContentPartText, Text: "tool answer"}}, Diffs: []llm.DiffData{{File: "/satellite/private", New: "secret"}}, Images: []string{"/satellite/result.png"}}, nil), + llm.UserText("second question"), + } + stream, err = resumedSatellite.Stream(t.Context(), llm.Request{ + Model: "model-a", SessionID: "session-a", Messages: secondMessages, + }) + if err != nil { + t.Fatal(err) + } + collectStream(t, stream) + + records := recorder.snapshot() + if len(records) != 2 { + t.Fatalf("central provider requests = %d, want 2", len(records)) + } + if records[0].Request.SessionID != "session-a" || !reflect.DeepEqual(records[0].Request.Messages, firstMessages) { + t.Fatalf("first central request lost session/transcript: %#v", records[0].Request) + } + if records[1].Request.SessionID != "session-a" || len(records[1].Request.Messages) != len(secondMessages) { + t.Fatalf("second central request lost session/transcript: %#v", records[1].Request) + } + if records[1].Imported != "state:session-a" { + t.Fatalf("reconstructed central provider imported %q, want session state", records[1].Imported) + } + assistant := records[1].Request.Messages[1] + if len(assistant.Parts) != 3 || assistant.Parts[1].ProviderReplay == nil || string(assistant.Parts[1].ProviderReplay.Raw) != string(replayRaw) { + t.Fatalf("reasoning/provider replay was not preserved: %#v", assistant.Parts) + } + if assistant.Parts[0].ReasoningEncryptedContent != "sealed-reasoning" || assistant.Parts[0].ReasoningContent != "summary" { + t.Fatalf("reasoning fields were not preserved: %#v", assistant.Parts[0]) + } + if assistant.Parts[2].ToolCall == nil || assistant.Parts[2].ToolCall.ID != "call_1" || assistant.Parts[2].ToolCall.ToolInfo != "" { + t.Fatalf("tool call semantics/display sanitization mismatch: %#v", assistant.Parts[2]) + } + result := records[1].Request.Messages[2].Parts[0].ToolResult + if result == nil || result.ID != "call_1" || result.Content != "tool answer" || len(result.ContentParts) != 1 || len(result.Diffs) != 0 || len(result.Images) != 0 { + t.Fatalf("tool-result continuation mismatch: %#v", result) + } +} + +func TestGatewayServerRestartRestoresSealedProviderStateWithSameKey(t *testing.T) { + recorder := &equivalenceProviderRecorder{} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, &equivalenceStateProvider{recorder: recorder}, time.Second) + installReconstructedProviderFactory(fixture, recorder) + + provider, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + stream, err := provider.Stream(t.Context(), llm.Request{Model: "model-a", SessionID: "restart-session", Messages: []llm.Message{llm.UserText("before restart")}}) + if err != nil { + t.Fatal(err) + } + collectStream(t, stream) + persisted, ok := provider.ExportProviderState() + if !ok { + t.Fatal("gateway provider did not export state before restart") + } + + fixture.server.Close() + sealer, err := OpenStateSealer(filepath.Join(fixture.stateDir, "state.key")) + if err != nil { + t.Fatal(err) + } + clients, err := OpenClientStore(filepath.Join(fixture.stateDir, "clients.json")) + if err != nil { + t.Fatal(err) + } + restarted, err := NewServer(ServerConfig{ + Config: fixture.central, Clients: clients, Sealer: sealer, Usage: fixture.usage, + ProviderFactory: func(*config.Config, string, string) (llm.Provider, error) { + return &equivalenceStateProvider{recorder: recorder}, nil + }, + Policy: Policy{AllowCLI: true, AllowSearch: true, AllowFetch: true}, + }) + if err != nil { + t.Fatal(err) + } + restartedHTTP := httptest.NewServer(restarted.Handler()) + defer restartedHTTP.Close() + satelliteConfig := fixture.satelliteConfig() + satelliteConfig.Gateway.URL = restartedHTTP.URL + + resumed, err := llm.NewGatewayProvider(satelliteConfig, "remote", "model-a") + if err != nil { + t.Fatal(err) + } + if err := resumed.ImportProviderState(persisted); err != nil { + t.Fatal(err) + } + stream, err = resumed.Stream(t.Context(), llm.Request{Model: "model-a", SessionID: "restart-session", Messages: []llm.Message{llm.UserText("before restart"), llm.AssistantText("ok"), llm.UserText("after restart")}}) + if err != nil { + t.Fatal(err) + } + collectStream(t, stream) + + records := recorder.snapshot() + if len(records) != 2 || records[1].Imported != "state:restart-session" || len(records[1].Request.Messages) != 3 { + t.Fatalf("restart records = %#v, want imported state plus full transcript", records) + } +} + +func TestGatewaySimultaneousSessionsDoNotCrossBindState(t *testing.T) { + recorder := &equivalenceProviderRecorder{} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, &equivalenceStateProvider{recorder: recorder}, time.Second) + installReconstructedProviderFactory(fixture, recorder) + + providers := make(map[string]*llm.GatewayProvider) + for _, sessionID := range []string{"session-a", "session-b"} { + provider, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + providers[sessionID] = provider + stream, err := provider.Stream(t.Context(), llm.Request{Model: "model-a", SessionID: sessionID, Messages: []llm.Message{llm.UserText("first " + sessionID)}}) + if err != nil { + t.Fatal(err) + } + collectStream(t, stream) + } + + var wg sync.WaitGroup + errs := make(chan error, 2) + for sessionID, provider := range providers { + sessionID, provider := sessionID, provider + wg.Add(1) + go func() { + defer wg.Done() + stream, err := provider.Stream(context.Background(), llm.Request{Model: "model-a", SessionID: sessionID, Messages: []llm.Message{llm.UserText("first " + sessionID), llm.AssistantText("ok"), llm.UserText("second " + sessionID)}}) + if err != nil { + errs <- err + return + } + defer stream.Close() + for { + _, err = stream.Recv() + if err != nil { + if !errors.Is(err, io.EOF) { + errs <- err + } + return + } + } + }() + } + wg.Wait() + close(errs) + for err := range errs { + t.Fatal(err) + } + + secondTurns := map[string]bool{} + for _, record := range recorder.snapshot() { + if len(record.Request.Messages) == 3 { + want := "state:" + record.Request.SessionID + if record.Imported != want { + t.Fatalf("session %q imported %q, want %q", record.Request.SessionID, record.Imported, want) + } + secondTurns[record.Request.SessionID] = true + } + } + if !secondTurns["session-a"] || !secondTurns["session-b"] { + t.Fatalf("simultaneous second turns = %#v", secondTurns) + } +} + +func TestGatewayRejectsSealedStateAcrossSessions(t *testing.T) { + recorder := &equivalenceProviderRecorder{} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, &equivalenceStateProvider{recorder: recorder}, time.Second) + installReconstructedProviderFactory(fixture, recorder) + + sessionA, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + stream, err := sessionA.Stream(t.Context(), llm.Request{Model: "model-a", SessionID: "session-a", Messages: []llm.Message{llm.UserText("a")}}) + if err != nil { + t.Fatal(err) + } + collectStream(t, stream) + sealed, ok := sessionA.ExportProviderState() + if !ok { + t.Fatal("session A did not export sealed state") + } + + sessionB, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + if err := sessionB.ImportProviderState(sealed); err != nil { + t.Fatal(err) + } + if _, err := sessionB.Stream(t.Context(), llm.Request{Model: "model-a", SessionID: "session-b", Messages: []llm.Message{llm.UserText("b")}}); err == nil || !strings.Contains(err.Error(), "invalid_state") { + t.Fatalf("cross-session state error = %v, want invalid_state", err) + } + if records := recorder.snapshot(); len(records) != 1 { + t.Fatalf("cross-session state reached provider Stream: %#v", records) + } +} + +func TestGatewayStatelessRequestDoesNotRoundTripProviderState(t *testing.T) { + recorder := &equivalenceProviderRecorder{} + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, &equivalenceStateProvider{recorder: recorder}, time.Second) + installReconstructedProviderFactory(fixture, recorder) + provider, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + stream, err := provider.Stream(t.Context(), llm.Request{Model: "model-a", Messages: []llm.Message{llm.UserText("one shot")}}) + if err != nil { + t.Fatal(err) + } + collectStream(t, stream) + if state, ok := provider.ExportProviderState(); ok || len(state) != 0 { + t.Fatalf("stateless request exported provider state %q, %t", state, ok) + } +} + +func cloneRequestForComparison(t *testing.T, req llm.Request) llm.Request { + t.Helper() + data, err := json.Marshal(req) + if err != nil { + t.Fatal(err) + } + var cloned llm.Request + if err := json.Unmarshal(data, &cloned); err != nil { + t.Fatal(err) + } + return cloned +} + +// normalizeDocumentedGatewayDifferences removes only fields that are deliberately +// satellite-local or display-only. Every remaining field can affect provider +// semantics and must compare exactly. +func normalizeDocumentedGatewayDifferences(req llm.Request) llm.Request { + req.WorkingDir = "" + req.ApprovalTranscriptPrefix = nil + req.AllowedTools = nil + req.AllowedToolsPresent = false + req.MaxTurns = 0 + req.ToolMap = nil + req.Debug = false + req.DebugRaw = false + for i := range req.Messages { + message := &req.Messages[i] + message.ApprovalRole = "" + message.ClientMessageID = "" + message.ResponseID = "" + message.AssistantSegmentOrdinal = 0 + message.SegmentStartSequence = 0 + message.SegmentEndSequence = 0 + parts := message.Parts[:0] + for _, part := range message.Parts { + if part.Type == llm.PartSkillActivation || part.Type == llm.PartToolActivity { + continue + } + part.ImagePath = "" + part.FilePath = "" + if part.ToolCall != nil { + part.ToolCall.ToolInfo = "" + } + if part.ToolResult != nil { + part.ToolResult.Display = "" + part.ToolResult.Diffs = nil + part.ToolResult.Images = nil + } + parts = append(parts, part) + } + message.Parts = parts + } + return req +} + +func TestGatewayProviderRequestMatchesDirectAfterDocumentedNormalization(t *testing.T) { + mock := llm.NewMockProvider("central").AddTextResponse("ok") + fixture := newGatewayFixture(t, config.ProviderTypeOpenAI, mock, time.Second) + provider, err := llm.NewGatewayProvider(fixture.satelliteConfig(), "remote", "model-a") + if err != nil { + t.Fatal(err) + } + lastChoice := llm.ToolChoice{Mode: llm.ToolChoiceName, Name: "lookup"} + req := llm.Request{ + Model: "model-a", SessionID: "differential-session", WorkingDir: "/satellite/private", Ephemeral: true, + IncludeDeveloperInContinuation: true, + Messages: []llm.Message{{ + Role: llm.RoleUser, CacheAnchor: true, ApprovalRole: "reviewer", ClientMessageID: "client-message", ResponseID: "response", AssistantSegmentOrdinal: 2, SegmentStartSequence: 3, SegmentEndSequence: 4, + Parts: []llm.Part{ + {Type: llm.PartText, Text: "hello", ReasoningContent: "reason", ReasoningSummaryParts: []string{"one", "two"}, ReasoningItemID: "rs", ReasoningEncryptedContent: "encrypted", ReasoningKind: llm.ReasoningKindSummary}, + {Type: llm.PartImage, ImagePath: "/satellite/image.png", ImageData: &llm.ToolImageData{MediaType: "image/png", Base64: "aQ==", Detail: "high"}}, + {Type: llm.PartFile, FilePath: "/satellite/file.pdf", FileData: &llm.ToolFileData{MediaType: "application/pdf", Base64: "Zg==", Filename: "file.pdf", SizeBytes: 1}}, + {Type: llm.PartToolCall, ToolCall: &llm.ToolCall{ID: "call", Name: "lookup", Arguments: json.RawMessage(`{"x":1}`), Caller: "programmatic", ToolInfo: "(private)", ThoughtSig: []byte("sig")}}, + {Type: llm.PartProviderReplay, ProviderReplay: &llm.ProviderReplayItem{Raw: json.RawMessage(`{"type":"reasoning","id":"rs"}`)}}, + {Type: llm.PartToolActivity, ToolActivity: &llm.ToolActivity{Name: "search", Info: "private display", Status: llm.ToolActivityCompleted}}, + {Type: llm.PartSkillActivation, SkillActivation: &llm.SkillActivationProvenance{Name: "private-skill", SourcePath: "/satellite/skill"}}, + }, + }, { + Role: llm.RoleTool, Parts: []llm.Part{{Type: llm.PartToolResult, ToolResult: &llm.ToolResult{ID: "call", Name: "lookup", Content: "result", ContentParts: []llm.ToolContentPart{{Type: llm.ToolContentPartText, Text: "result"}}, Display: "private display", Diffs: []llm.DiffData{{File: "/satellite/private", New: "x"}}, Images: []string{"/satellite/result.png"}, IsError: true, Caller: "programmatic", ThoughtSig: []byte("sig")}}}, + }}, + ApprovalTranscriptPrefix: []llm.Message{llm.UserText("approval-only")}, + Tools: []llm.ToolSpec{{Name: "lookup", Description: "look up", Schema: map[string]any{"type": "object", "additionalProperties": false}, Strict: true, AllowedCallers: []string{"programmatic"}, OutputSchema: map[string]any{"type": "string"}}}, + ToolChoice: llm.ToolChoice{Mode: llm.ToolChoiceRequired}, LastTurnToolChoice: &lastChoice, ParallelToolCalls: true, + AllowedTools: []string{"lookup"}, AllowedToolsPresent: true, + Search: true, ForceExternalSearch: true, DisableExternalWebFetch: true, + ReasoningEffort: "high", Responses: &llm.ResponsesOptions{ReasoningMode: "summary", ReasoningContext: "preserve", MultiAgent: llm.MultiAgentOptions{Enabled: true, EnabledSet: true, MaxConcurrentSubagents: 2}, ProgrammaticToolCalling: llm.ProgrammaticToolCallingOptions{Enabled: true, EnabledSet: true, Tools: []string{"lookup"}}, PromptCache: llm.PromptCacheOptions{Mode: "memory", TTL: "1h"}}, + MaxOutputTokens: 321, Temperature: 0, TemperatureSet: true, TopP: .7, TopPSet: true, + ServiceTier: "priority", ServiceTierSet: true, MaxTurns: 7, ToolMap: map[string]string{"lookup": "local_lookup"}, Debug: true, DebugRaw: true, + } + + stream, err := provider.Stream(t.Context(), req) + if err != nil { + t.Fatal(err) + } + collectStream(t, stream) + recorded := mock.RecordedRequests() + if len(recorded) != 1 { + t.Fatalf("recorded requests = %d, want 1", len(recorded)) + } + + want := normalizeDocumentedGatewayDifferences(cloneRequestForComparison(t, req)) + got := normalizeDocumentedGatewayDifferences(cloneRequestForComparison(t, recorded[0])) + if !reflect.DeepEqual(got, want) { + wantJSON, _ := json.MarshalIndent(want, "", " ") + gotJSON, _ := json.MarshalIndent(got, "", " ") + t.Fatalf("direct/gateway provider request semantic mismatch\nwant: %s\n got: %s", wantJSON, gotJSON) + } +} diff --git a/internal/gateway/execution.go b/internal/gateway/execution.go index ada212d6a..24e300fd7 100644 --- a/internal/gateway/execution.go +++ b/internal/gateway/execution.go @@ -83,9 +83,9 @@ func (s *Server) startInference(parent context.Context, client Client, envelope MaxBackoff: 5 * time.Second, }) if envelope.State != "" { - plain, openErr := s.cfg.Sealer.Open(envelope.State, client.ID, envelope.Provider) + plain, openErr := s.cfg.Sealer.Open(envelope.State, client.ID, envelope.Provider, providerReq.SessionID) if openErr != nil { - return failStarted(http.StatusBadRequest, "invalid_state", "provider state is invalid or does not belong to this client/provider") + return failStarted(http.StatusBadRequest, "invalid_state", "provider state is invalid or does not belong to this client/provider/session") } importer, ok := provider.(llm.ProviderStateImporter) if !ok { diff --git a/internal/gateway/integration_test.go b/internal/gateway/integration_test.go index 386279571..51e2edff5 100644 --- a/internal/gateway/integration_test.go +++ b/internal/gateway/integration_test.go @@ -134,6 +134,7 @@ func (p *setupFailureProvider) attemptCount() int { type gatewayFixture struct { server *httptest.Server gateway *Server + stateDir string central *config.Config clients *ClientStore client Client @@ -178,7 +179,7 @@ func newGatewayFixture(t *testing.T, providerType config.ProviderType, provider } ts := httptest.NewServer(server.Handler()) t.Cleanup(ts.Close) - return &gatewayFixture{server: ts, gateway: server, central: central, clients: clients, client: client, token: token, usage: usage, provider: provider} + return &gatewayFixture{server: ts, gateway: server, stateDir: dir, central: central, clients: clients, client: client, token: token, usage: usage, provider: provider} } func (f *gatewayFixture) satelliteConfig() *config.Config { @@ -467,7 +468,7 @@ func TestGatewayCrossClientStateAndRunAccessDenied(t *testing.T) { if err := provider.ImportProviderState([]byte("forged-state")); err != nil { t.Fatal(err) } - if _, err := provider.Stream(context.Background(), llm.Request{Model: "model-a", Messages: []llm.Message{llm.UserText("x")}}); err == nil || !strings.Contains(err.Error(), "invalid_state") { + if _, err := provider.Stream(context.Background(), llm.Request{Model: "model-a", SessionID: "forged-session", Messages: []llm.Message{llm.UserText("x")}}); err == nil || !strings.Contains(err.Error(), "invalid_state") { t.Fatalf("tampered state error = %v", err) } other, otherToken, err := fixture.clients.Add("satellite-b", Policy{AllowSearch: true, AllowFetch: true}) @@ -498,7 +499,7 @@ func TestGatewayProviderStateRoundTripsSealed(t *testing.T) { t.Fatal(err) } for i := 0; i < 2; i++ { - stream, err := provider.Stream(context.Background(), llm.Request{Model: "model-a", Messages: []llm.Message{llm.UserText("turn")}}) + stream, err := provider.Stream(context.Background(), llm.Request{Model: "model-a", SessionID: "session-state", Messages: []llm.Message{llm.UserText("turn")}}) if err != nil { t.Fatal(err) } diff --git a/internal/gateway/seal.go b/internal/gateway/seal.go index 476ad8791..9e42d3f17 100644 --- a/internal/gateway/seal.go +++ b/internal/gateway/seal.go @@ -10,15 +10,17 @@ import ( "io" "os" "path/filepath" + "strings" "github.com/samsaffron/term-llm/internal/gateway/protocol" ) type sealedProviderState struct { - Version int `json:"version"` - ClientID string `json:"client_id"` - Provider string `json:"provider"` - State []byte `json:"state"` + Version int `json:"version"` + ClientID string `json:"client_id"` + Provider string `json:"provider"` + SessionID string `json:"session_id"` + State []byte `json:"state"` } type StateSealer struct{ aead cipher.AEAD } @@ -54,8 +56,14 @@ func OpenStateSealer(path string) (*StateSealer, error) { return &StateSealer{aead: aead}, nil } -func (s *StateSealer) Seal(clientID, provider string, state []byte) (string, error) { - plain, err := json.Marshal(sealedProviderState{Version: protocol.Version, ClientID: clientID, Provider: provider, State: state}) +func (s *StateSealer) Seal(clientID, provider, sessionID string, state []byte) (string, error) { + if strings.TrimSpace(sessionID) == "" { + return "", fmt.Errorf("cannot seal provider state without a session ID") + } + plain, err := json.Marshal(sealedProviderState{ + Version: protocol.Version, ClientID: clientID, Provider: provider, + SessionID: sessionID, State: state, + }) if err != nil { return "", err } @@ -67,7 +75,10 @@ func (s *StateSealer) Seal(clientID, provider string, state []byte) (string, err return base64.RawURLEncoding.EncodeToString(append(nonce, ciphertext...)), nil } -func (s *StateSealer) Open(blob, clientID, provider string) ([]byte, error) { +func (s *StateSealer) Open(blob, clientID, provider, sessionID string) ([]byte, error) { + if strings.TrimSpace(sessionID) == "" { + return nil, fmt.Errorf("cannot open provider state without a session ID") + } raw, err := base64.RawURLEncoding.DecodeString(blob) if err != nil || len(raw) < s.aead.NonceSize() { return nil, fmt.Errorf("invalid sealed provider state") @@ -78,8 +89,10 @@ func (s *StateSealer) Open(blob, clientID, provider string) ([]byte, error) { return nil, fmt.Errorf("invalid sealed provider state") } var state sealedProviderState - if err := json.Unmarshal(plain, &state); err != nil || state.Version != protocol.Version || state.ClientID != clientID || state.Provider != provider { - return nil, fmt.Errorf("sealed provider state does not belong to this client/provider") + if err := json.Unmarshal(plain, &state); err != nil || + state.Version != protocol.Version || state.ClientID != clientID || + state.Provider != provider || state.SessionID != sessionID { + return nil, fmt.Errorf("sealed provider state does not belong to this client/provider/session") } return append([]byte(nil), state.State...), nil } diff --git a/internal/gateway/security_test.go b/internal/gateway/security_test.go index e75d9cfcf..1c3196431 100644 --- a/internal/gateway/security_test.go +++ b/internal/gateway/security_test.go @@ -248,16 +248,16 @@ func TestGatewayRunTempRootScavengesOnlyOwnedPrefixDirectories(t *testing.T) { } } -func TestStateSealerRoundTripTamperAndCrossClient(t *testing.T) { +func TestStateSealerRoundTripTamperAndBinding(t *testing.T) { sealer, err := OpenStateSealer(filepath.Join(t.TempDir(), "state.key")) if err != nil { t.Fatal(err) } - blob, err := sealer.Seal("client-a", "claude-bin", []byte("gateway-local-state")) + blob, err := sealer.Seal("client-a", "claude-bin", "session-a", []byte("gateway-local-state")) if err != nil { t.Fatal(err) } - plain, err := sealer.Open(blob, "client-a", "claude-bin") + plain, err := sealer.Open(blob, "client-a", "claude-bin", "session-a") if err != nil || string(plain) != "gateway-local-state" { t.Fatalf("round trip = %q, %v", plain, err) } @@ -267,15 +267,22 @@ func TestStateSealerRoundTripTamperAndCrossClient(t *testing.T) { } else { tampered[len(tampered)/2] = 'A' } - for _, tc := range []struct{ blob, client, provider string }{ - {string(tampered), "client-a", "claude-bin"}, - {blob, "client-b", "claude-bin"}, - {blob, "client-a", "grok-bin"}, + for _, tc := range []struct{ blob, client, provider, session string }{ + {string(tampered), "client-a", "claude-bin", "session-a"}, + {blob, "client-b", "claude-bin", "session-a"}, + {blob, "client-a", "grok-bin", "session-a"}, + {blob, "client-a", "claude-bin", "session-b"}, } { - if _, err := sealer.Open(tc.blob, tc.client, tc.provider); err == nil { + if _, err := sealer.Open(tc.blob, tc.client, tc.provider, tc.session); err == nil { t.Fatalf("accepted tampered/foreign state: %+v", tc) } } + if _, err := sealer.Seal("client-a", "claude-bin", "", []byte("state")); err == nil { + t.Fatal("sealed state for an empty stateless session") + } + if _, err := sealer.Open(blob, "client-a", "claude-bin", ""); err == nil { + t.Fatal("opened state for an empty stateless session") + } } func TestEnrollmentCreatesUniqueAuthenticatedClient(t *testing.T) { diff --git a/internal/gateway/server.go b/internal/gateway/server.go index a55bec33e..33a4a7b80 100644 --- a/internal/gateway/server.go +++ b/internal/gateway/server.go @@ -473,11 +473,13 @@ func (s *Server) handleInference(w http.ResponseWriter, r *http.Request, client if errorCode == "" && errors.Is(ctx.Err(), context.Canceled) { errorCode = "canceled" } - if exporter, ok := provider.(llm.ProviderStateExporter); ok { - if plain, valid := exporter.ExportProviderState(); valid { - if sealed, sealErr := s.cfg.Sealer.Seal(client.ID, envelope.Provider, plain); sealErr == nil { - _ = writeSSE(w, protocol.StreamRecord{Version: protocol.Version, Type: "state", RequestID: envelope.RequestID, RunID: runID, State: sealed}) - flusher.Flush() + if strings.TrimSpace(providerReq.SessionID) != "" { + if exporter, ok := provider.(llm.ProviderStateExporter); ok { + if plain, valid := exporter.ExportProviderState(); valid { + if sealed, sealErr := s.cfg.Sealer.Seal(client.ID, envelope.Provider, providerReq.SessionID, plain); sealErr == nil { + _ = writeSSE(w, protocol.StreamRecord{Version: protocol.Version, Type: "state", RequestID: envelope.RequestID, RunID: runID, State: sealed}) + flusher.Flush() + } } } } diff --git a/internal/llm/chatgpt.go b/internal/llm/chatgpt.go index 594df1f92..206d86da5 100644 --- a/internal/llm/chatgpt.go +++ b/internal/llm/chatgpt.go @@ -231,6 +231,17 @@ func (p *ChatGPTProvider) Stream(ctx context.Context, req Request) (Stream, erro } } + responsesReq, err := p.buildResponsesRequest(req) + if err != nil { + return nil, err + } + return p.responsesClient.Stream(ctx, responsesReq, req.DebugRaw) +} + +// buildResponsesRequest constructs the provider-facing request independently of +// transport state. ResponsesClient decides whether HTTP/SSE sends full history or +// a reused WebSocket sends previous_response_id plus a continuation suffix. +func (p *ChatGPTProvider) buildResponsesRequest(req Request) (ResponsesRequest, error) { // Effort precedence: req.ReasoningEffort wins over model suffix, which wins over provider-level effort. reqModel, reqEffort := parseModelEffortForProvider("chatgpt", req.Model) model := chooseModel(reqModel, p.model) @@ -244,7 +255,7 @@ func (p *ChatGPTProvider) Stream(ctx context.Context, req Request) (Stream, erro responsesOptions := mergeResponsesOptions(p.responsesOptions, req.Responses, req.Ephemeral) if _, err := validateResponsesOptions("chatgpt", model, &responsesOptions, req.Tools); err != nil { - return nil, err + return ResponsesRequest{}, err } // Build tools. Public-API Pro and advanced Responses controls are not @@ -296,7 +307,7 @@ func (p *ChatGPTProvider) Stream(ctx context.Context, req Request) (Stream, erro responsesReq.Reasoning.Effort = effort } - return p.responsesClient.Stream(ctx, responsesReq, req.DebugRaw) + return responsesReq, nil } // ResetConversation clears server state for the Responses API client. diff --git a/internal/llm/chatgpt_test.go b/internal/llm/chatgpt_test.go index 4a569bda1..5f8eaa75b 100644 --- a/internal/llm/chatgpt_test.go +++ b/internal/llm/chatgpt_test.go @@ -3,12 +3,17 @@ package llm import ( "context" "encoding/json" + "fmt" "io" "net/http" + "net/http/httptest" + "reflect" "strings" + "sync" "testing" "time" + "github.com/gorilla/websocket" "github.com/samsaffron/term-llm/internal/credentials" ) @@ -443,3 +448,210 @@ func TestChatGPTStream_ReasoningSummaryByOutputIndex(t *testing.T) { t.Fatalf("reasoning tokens = %d, want 1", usageEvent.Use.ReasoningTokens) } } + +func chatGPTTestCredentials() *credentials.ChatGPTCredentials { + return &credentials.ChatGPTCredentials{ + AccessToken: "test-token", + AccountID: "test-account", + ExpiresAt: time.Now().Add(time.Hour).Unix(), + } +} + +func chatGPTCompletedSSE(responseID string) string { + return fmt.Sprintf("event: response.completed\ndata: {\"type\":\"response.completed\",\"response\":{\"id\":%q,\"usage\":{\"input_tokens\":1,\"output_tokens\":1,\"total_tokens\":2}}}\n\ndata: [DONE]\n\n", responseID) +} + +func TestChatGPTHTTPSecondTurnMatchesReconstructedGatewayProvider(t *testing.T) { + origClient := chatGPTHTTPClient + defer func() { chatGPTHTTPClient = origClient }() + + var captured []map[string]any + chatGPTHTTPClient = &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + var payload map[string]any + if err := json.NewDecoder(req.Body).Decode(&payload); err != nil { + return nil, err + } + captured = append(captured, payload) + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(strings.NewReader(chatGPTCompletedSSE(fmt.Sprintf("resp_%d", len(captured))))), + Header: make(http.Header), + }, nil + })} + + first := Request{Model: "gpt-5.6-sol", SessionID: "session-equivalence", Messages: []Message{UserText("first")}} + second := Request{Model: "gpt-5.6-sol", SessionID: "session-equivalence", Messages: []Message{UserText("first"), AssistantText("answer"), UserText("second")}} + + // Direct HTTP/SSE keeps one provider/client alive across turns. + direct := NewChatGPTProviderWithCreds(chatGPTTestCredentials(), "gpt-5.6-sol") + stream, err := direct.Stream(t.Context(), first) + if err != nil { + t.Fatal(err) + } + drainStreamToDone(t, stream) + _ = stream.Close() + stream, err = direct.Stream(t.Context(), second) + if err != nil { + t.Fatal(err) + } + drainStreamToDone(t, stream) + _ = stream.Close() + + // The gateway creates a fresh central provider for this second request. + reconstructed := NewChatGPTProviderWithCreds(chatGPTTestCredentials(), "gpt-5.6-sol") + stream, err = reconstructed.Stream(t.Context(), second) + if err != nil { + t.Fatal(err) + } + drainStreamToDone(t, stream) + _ = stream.Close() + + if len(captured) != 3 { + t.Fatalf("captured payloads = %d, want 3", len(captured)) + } + if !reflect.DeepEqual(captured[1], captured[2]) { + directJSON, _ := json.Marshal(captured[1]) + gatewayJSON, _ := json.Marshal(captured[2]) + t.Fatalf("direct/gateway HTTP second-turn payloads differ\ndirect: %s\ngateway: %s", directJSON, gatewayJSON) + } + for index, payload := range captured[1:] { + if _, ok := payload["previous_response_id"]; ok { + t.Fatalf("second-turn HTTP payload %d used previous_response_id: %#v", index, payload) + } + input, ok := payload["input"].([]any) + if !ok || len(input) != 3 { + t.Fatalf("second-turn HTTP payload %d input = %#v, want full transcript", index, payload["input"]) + } + } +} + +func TestChatGPTWebSocketDirectContinuationGatewayDifferenceAndFullHistoryFallback(t *testing.T) { + var mu sync.Mutex + connections := 0 + var captured []map[string]any + capture := func(data []byte) map[string]any { + var payload map[string]any + if err := json.Unmarshal(data, &payload); err != nil { + t.Errorf("decode WebSocket request: %v", err) + } + mu.Lock() + captured = append(captured, payload) + mu.Unlock() + return payload + } + + upgrader := websocket.Upgrader{} + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := upgrader.Upgrade(w, r, nil) + if err != nil { + t.Errorf("upgrade: %v", err) + return + } + defer conn.Close() + mu.Lock() + connections++ + connection := connections + mu.Unlock() + + if connection == 1 { + // Direct mode: first full turn, optimized second turn, then semantic + // full-history retry after the connection-local parent is rejected. + _, data, err := conn.ReadMessage() + if err != nil { + t.Errorf("read direct first request: %v", err) + return + } + capture(data) + _ = conn.WriteJSON(map[string]any{"type": "response.completed", "response": map[string]any{"id": "resp_direct_1"}}) + + _, data, err = conn.ReadMessage() + if err != nil { + t.Errorf("read direct continuation: %v", err) + return + } + capture(data) + _ = conn.WriteJSON(map[string]any{ + "type": "response.failed", "status": 400, + "response": map[string]any{"error": map[string]any{"code": "previous_response_not_found", "message": "Previous response not found", "param": "previous_response_id"}}, + }) + + _, data, err = conn.ReadMessage() + if err != nil { + t.Errorf("read direct full-history retry: %v", err) + return + } + capture(data) + _ = conn.WriteJSON(map[string]any{"type": "response.completed", "response": map[string]any{"id": "resp_direct_2"}}) + return + } + + // Gateway-equivalent reconstruction has no connection-local response ID, + // so its second-turn transcript is complete on its first frame. + _, data, err := conn.ReadMessage() + if err != nil { + t.Errorf("read reconstructed request: %v", err) + return + } + capture(data) + _ = conn.WriteJSON(map[string]any{"type": "response.completed", "response": map[string]any{"id": "resp_gateway"}}) + })) + defer server.Close() + + newWSProvider := func() *ChatGPTProvider { + provider := NewChatGPTProviderWithCredsAndOptions(chatGPTTestCredentials(), "gpt-5.6-sol", ChatGPTProviderOptions{UseWebSocket: true}) + provider.responsesClient = &ResponsesClient{ + BaseURL: server.URL, HTTPClient: server.Client(), UseWebSocket: true, + WebSocketServerState: true, DisableServerState: true, + } + return provider + } + first := Request{Model: "gpt-5.6-sol", SessionID: "session-ws", Messages: []Message{UserText("first")}} + second := Request{Model: "gpt-5.6-sol", SessionID: "session-ws", Messages: []Message{UserText("first"), AssistantText("answer"), UserText("second")}} + + direct := newWSProvider() + stream, err := direct.Stream(t.Context(), first) + if err != nil { + t.Fatal(err) + } + drainStreamToDone(t, stream) + _ = stream.Close() + stream, err = direct.Stream(t.Context(), second) + if err != nil { + t.Fatal(err) + } + drainStreamToDone(t, stream) + _ = stream.Close() + + reconstructed := newWSProvider() + stream, err = reconstructed.Stream(t.Context(), second) + if err != nil { + t.Fatal(err) + } + drainStreamToDone(t, stream) + _ = stream.Close() + + mu.Lock() + requests := append([]map[string]any(nil), captured...) + mu.Unlock() + if len(requests) != 4 { + t.Fatalf("captured WebSocket requests = %d, want 4", len(requests)) + } + directContinuation := requests[1] + if directContinuation["previous_response_id"] != "resp_direct_1" { + t.Fatalf("direct continuation previous_response_id = %#v", directContinuation["previous_response_id"]) + } + if input, ok := directContinuation["input"].([]any); !ok || len(input) != 1 { + t.Fatalf("direct continuation input = %#v, want only new suffix", directContinuation["input"]) + } + for name, request := range map[string]map[string]any{"direct fallback": requests[2], "gateway reconstruction": requests[3]} { + if _, ok := request["previous_response_id"]; ok { + t.Fatalf("%s retained previous_response_id: %#v", name, request) + } + if input, ok := request["input"].([]any); !ok || len(input) != 3 { + t.Fatalf("%s input = %#v, want full transcript", name, request["input"]) + } + } + if !reflect.DeepEqual(requests[2], requests[3]) { + t.Fatalf("full-history WS fallback and gateway reconstruction differ\nfallback: %#v\ngateway: %#v", requests[2], requests[3]) + } +} diff --git a/internal/llm/gateway_provider.go b/internal/llm/gateway_provider.go index e58a6c11c..c29504c08 100644 --- a/internal/llm/gateway_provider.go +++ b/internal/llm/gateway_provider.go @@ -176,6 +176,13 @@ func (p *GatewayProvider) Stream(ctx context.Context, req Request) (Stream, erro p.mu.RLock() state := p.state p.mu.RUnlock() + stateSessionID := req.SessionID + if strings.TrimSpace(stateSessionID) == "" { + // Empty SessionID requests are deliberately stateless: never attach opaque + // continuation state from another one-shot call. + state = "" + stateSessionID = "" + } payload, err := json.Marshal(protocol.InferenceRequest{ Version: protocol.Version, RequestID: requestID, Provider: p.name, State: state, Request: wireRequest, @@ -203,7 +210,7 @@ func (p *GatewayProvider) Stream(ctx context.Context, req Request) (Stream, erro return nil, decodeGatewayHTTPError(resp) } stream := &gatewayProviderStream{ - provider: p, body: resp.Body, ctx: streamCtx, cancel: cancel, + provider: p, body: resp.Body, ctx: streamCtx, cancel: cancel, sessionID: stateSessionID, decoder: newSSEDecoder(resp.Body, sseDecoderOptions{Transport: "gateway SSE"}), closedSignal: make(chan struct{}), results: make(chan nextGatewaySSE, 1), } @@ -235,6 +242,7 @@ type gatewayProviderStream struct { decoder *sseDecoder ctx context.Context cancel context.CancelFunc + sessionID string results chan nextGatewaySSE recvMu sync.Mutex mu sync.Mutex @@ -346,6 +354,10 @@ func (s *gatewayProviderStream) Recv() (Event, error) { go s.postToolResult(record.CallbackPath, response) return event, nil case "state": + if s.sessionID == "" { + s.terminate() + return Event{}, fmt.Errorf("gateway returned provider state for a stateless request") + } s.provider.mu.Lock() s.provider.state = record.State s.provider.mu.Unlock() diff --git a/internal/llm/gateway_wire.go b/internal/llm/gateway_wire.go index 58fab7733..723f5a582 100644 --- a/internal/llm/gateway_wire.go +++ b/internal/llm/gateway_wire.go @@ -9,8 +9,9 @@ import ( ) // gatewayRequest is deliberately distinct from Request. In particular it has -// no WorkingDir, approval transcript, execution filters, or debug fields. This -// makes local filesystem/process state structurally impossible on the wire. +// no WorkingDir, approval transcript/metadata, execution filters/budgets/maps, +// debug fields, local paths, or display-only tool metadata. This makes local +// filesystem/process state structurally impossible on the wire. type gatewayRequest struct { Model string `json:"model"` SessionID string `json:"session_id,omitempty"` @@ -33,14 +34,20 @@ type gatewayRequest struct { TopPSet bool `json:"top_p_set,omitempty"` ServiceTier string `json:"service_tier,omitempty"` ServiceTierSet bool `json:"service_tier_set,omitempty"` - MaxTurns int `json:"max_turns,omitempty"` - ToolMap map[string]string `json:"tool_map,omitempty"` } func sanitizeGatewayMessages(messages []Message) []Message { out := make([]Message, len(messages)) for i, message := range messages { out[i] = message + // Approval and persisted/UI projection metadata are satellite-local. CacheAnchor + // remains because it changes provider prompt-cache semantics. + out[i].ApprovalRole = "" + out[i].ClientMessageID = "" + out[i].ResponseID = "" + out[i].AssistantSegmentOrdinal = 0 + out[i].SegmentStartSequence = 0 + out[i].SegmentEndSequence = 0 out[i].Parts = make([]Part, 0, len(message.Parts)) for _, part := range message.Parts { if part.Type == PartSkillActivation || part.Type == PartToolActivity { @@ -48,8 +55,14 @@ func sanitizeGatewayMessages(messages []Message) []Message { } part.ImagePath = "" part.FilePath = "" + if part.ToolCall != nil { + call := *part.ToolCall + call.ToolInfo = "" + part.ToolCall = &call + } if part.ToolResult != nil { result := *part.ToolResult + result.Display = "" result.Diffs = nil result.Images = nil part.ToolResult = &result @@ -74,7 +87,6 @@ func EncodeGatewayRequest(req Request) (json.RawMessage, error) { MaxOutputTokens: req.MaxOutputTokens, Temperature: req.Temperature, TemperatureSet: req.TemperatureSet, TopP: req.TopP, TopPSet: req.TopPSet, ServiceTier: req.ServiceTier, ServiceTierSet: req.ServiceTierSet, - MaxTurns: req.MaxTurns, ToolMap: req.ToolMap, } data, err := json.Marshal(wire) return data, err @@ -100,7 +112,7 @@ func DecodeGatewayRequest(data []byte) (Request, error) { Responses: wire.Responses, MaxOutputTokens: wire.MaxOutputTokens, Temperature: wire.Temperature, TemperatureSet: wire.TemperatureSet, TopP: wire.TopP, TopPSet: wire.TopPSet, ServiceTier: wire.ServiceTier, - ServiceTierSet: wire.ServiceTierSet, MaxTurns: wire.MaxTurns, ToolMap: wire.ToolMap, + ServiceTierSet: wire.ServiceTierSet, }, nil } diff --git a/internal/llm/gateway_wire_test.go b/internal/llm/gateway_wire_test.go index 584f9d3e1..17cddfc24 100644 --- a/internal/llm/gateway_wire_test.go +++ b/internal/llm/gateway_wire_test.go @@ -15,29 +15,30 @@ func TestGatewayWireRequestRoundTripAndOmitsLocalState(t *testing.T) { ApprovalTranscriptPrefix: []Message{UserText("private approval")}, AllowedTools: []string{"read_file"}, AllowedToolsPresent: true, Debug: true, DebugRaw: true, Messages: []Message{ - {Role: RoleUser, Parts: []Part{ + {Role: RoleUser, ApprovalRole: "private-approval-role", ClientMessageID: "private-client-id", ResponseID: "private-response-id", AssistantSegmentOrdinal: 2, SegmentStartSequence: 3, SegmentEndSequence: 4, Parts: []Part{ {Type: PartText, Text: "hello", ReasoningContent: "summary", ReasoningKind: ReasoningKindSummary}, {Type: PartImage, ImagePath: "/private/image.png", ImageData: &ToolImageData{MediaType: "image/png", Base64: "aW1hZ2U="}}, {Type: PartFile, FilePath: "/private/file.pdf", FileData: &ToolFileData{MediaType: "application/pdf", Base64: "ZmlsZQ==", Filename: "file.pdf"}}, }}, {Role: RoleAssistant, Parts: []Part{ - {Type: PartToolCall, ToolCall: &ToolCall{ID: "c1", Name: "view", Arguments: json.RawMessage(`{"x":1}`), Caller: "programmatic", ThoughtSig: []byte("sig")}}, + {Type: PartToolCall, ToolCall: &ToolCall{ID: "c1", Name: "view", Arguments: json.RawMessage(`{"x":1}`), Caller: "programmatic", ToolInfo: "private-tool-info", ThoughtSig: []byte("sig")}}, {Type: PartProviderReplay, ProviderReplay: &ProviderReplayItem{Raw: json.RawMessage(`{"type":"opaque"}`)}}, }}, - {Role: RoleTool, Parts: []Part{{Type: PartToolResult, ToolResult: &ToolResult{ID: "c1", Name: "view", Content: "ok", ContentParts: []ToolContentPart{{Type: ToolContentPartText, Text: "ok"}, {Type: ToolContentPartImageData, ImageData: &ToolImageData{MediaType: "image/png", Base64: "aQ=="}}}, Images: []string{"/private/result.png"}}}}}, + {Role: RoleTool, Parts: []Part{{Type: PartToolResult, ToolResult: &ToolResult{ID: "c1", Name: "view", Content: "ok", ContentParts: []ToolContentPart{{Type: ToolContentPartText, Text: "ok"}, {Type: ToolContentPartImageData, ImageData: &ToolImageData{MediaType: "image/png", Base64: "aQ=="}}}, Display: "private-result-display", Images: []string{"/private/result.png"}}}}}, }, Tools: []ToolSpec{{Name: "view", Description: "view", Schema: map[string]any{"type": "object"}, Strict: true, AllowedCallers: []string{"programmatic"}, OutputSchema: map[string]any{"type": "string"}}}, ToolChoice: ToolChoice{Mode: ToolChoiceName, Name: "view"}, ParallelToolCalls: true, Search: true, ForceExternalSearch: true, ReasoningEffort: "high", MaxOutputTokens: 42, Temperature: 0, TemperatureSet: true, TopP: .5, TopPSet: true, Responses: &ResponsesOptions{ReasoningMode: "summary", MultiAgent: MultiAgentOptions{Enabled: true, EnabledSet: true, MaxConcurrentSubagents: 2}}, + MaxTurns: 9, ToolMap: map[string]string{"view": "local_view"}, } wire, err := EncodeGatewayRequest(req) if err != nil { t.Fatal(err) } text := string(wire) - for _, forbidden := range []string{"WorkingDir", "working_dir", "/satellite/private", "/private/image.png", "/private/file.pdf", "/private/result.png", "private approval", "AllowedTools", "DebugRaw"} { + for _, forbidden := range []string{"WorkingDir", "working_dir", "/satellite/private", "/private/image.png", "/private/file.pdf", "/private/result.png", "private approval", "private-approval-role", "private-client-id", "private-response-id", "private-tool-info", "private-result-display", "AllowedTools", "max_turns", "tool_map", "DebugRaw"} { if strings.Contains(text, forbidden) { t.Fatalf("wire request leaked forbidden %q: %s", forbidden, text) } @@ -52,6 +53,9 @@ func TestGatewayWireRequestRoundTripAndOmitsLocalState(t *testing.T) { if got.Messages[0].Parts[1].ImageData == nil || got.Messages[0].Parts[2].FileData == nil { t.Fatalf("vision/file data lost: %+v", got.Messages[0].Parts) } + if got.Messages[1].Parts[0].ToolCall.ToolInfo != "" || got.Messages[2].Parts[0].ToolResult.Display != "" || got.Messages[0].ApprovalRole != "" || got.Messages[0].ClientMessageID != "" || got.Messages[0].ResponseID != "" { + t.Fatalf("local/display metadata crossed wire: %+v", got.Messages) + } if got.Messages[2].Parts[0].ToolResult.Images != nil { t.Fatalf("local result paths crossed wire: %+v", got.Messages[2].Parts[0].ToolResult) } From 437d2fbeb9e2ec31279c8cbb2457c24f4c1d2cd0 Mon Sep 17 00:00:00 2001 From: Jarvis Date: Sat, 1 Aug 2026 19:35:08 +1000 Subject: [PATCH 3/3] Reuse gateway WebSocket provider sessions --- cmd/gateway.go | 82 +- cmd/gateway_session_test.go | 19 + docs-site/content/guides/inference-gateway.md | 16 +- internal/gateway/execution.go | 103 ++- internal/gateway/provider_sessions.go | 266 ++++++ internal/gateway/provider_sessions_test.go | 795 ++++++++++++++++++ internal/gateway/responses.go | 6 +- internal/gateway/server.go | 165 +++- internal/llm/chatgpt_test.go | 57 +- 9 files changed, 1394 insertions(+), 115 deletions(-) create mode 100644 cmd/gateway_session_test.go create mode 100644 internal/gateway/provider_sessions.go create mode 100644 internal/gateway/provider_sessions_test.go diff --git a/cmd/gateway.go b/cmd/gateway.go index 256f34f9a..f415da9e8 100644 --- a/cmd/gateway.go +++ b/cmd/gateway.go @@ -21,42 +21,43 @@ import ( ) var ( - gatewayStateDir string - gatewayListen string - gatewayTLSCert string - gatewayTLSKey string - gatewayAllowProviders []string - gatewayDenyProviders []string - gatewayAllowModels []string - gatewayDenyModels []string - gatewayAllowCLI bool - gatewayNoSearch bool - gatewayNoFetch bool - gatewayIdleTimeout time.Duration - gatewayToolTimeout time.Duration - gatewayCatalogTTL time.Duration - gatewayRetryAttempts int - gatewayRetryElapsed time.Duration - gatewayClientAllowCLI bool - gatewayClientAllow []string - gatewayClientDeny []string - gatewayClientModels []string - gatewayClientDenyModel []string - gatewayClientSearch bool - gatewayClientFetch bool - gatewayClientInference int - gatewayClientSearchRPM int - gatewayClientSearchMax int - gatewayClientFetchRPM int - gatewayClientFetchMax int - gatewayClientEnroll bool - gatewayEnrollmentTTL time.Duration - gatewayEnrollName string - gatewayEnrollWrite bool - gatewayEnrollTokenFile string - gatewayEnrollPrintOnly bool - gatewayUsageClient string - gatewayUsageJSON bool + gatewayStateDir string + gatewayListen string + gatewayTLSCert string + gatewayTLSKey string + gatewayAllowProviders []string + gatewayDenyProviders []string + gatewayAllowModels []string + gatewayDenyModels []string + gatewayAllowCLI bool + gatewayNoSearch bool + gatewayNoFetch bool + gatewayIdleTimeout time.Duration + gatewayToolTimeout time.Duration + gatewayCatalogTTL time.Duration + gatewayRetryAttempts int + gatewayRetryElapsed time.Duration + gatewayProviderSessionTimeout time.Duration + gatewayClientAllowCLI bool + gatewayClientAllow []string + gatewayClientDeny []string + gatewayClientModels []string + gatewayClientDenyModel []string + gatewayClientSearch bool + gatewayClientFetch bool + gatewayClientInference int + gatewayClientSearchRPM int + gatewayClientSearchMax int + gatewayClientFetchRPM int + gatewayClientFetchMax int + gatewayClientEnroll bool + gatewayEnrollmentTTL time.Duration + gatewayEnrollName string + gatewayEnrollWrite bool + gatewayEnrollTokenFile string + gatewayEnrollPrintOnly bool + gatewayUsageClient string + gatewayUsageJSON bool ) var gatewayCmd = &cobra.Command{Use: "gateway", Short: "Serve and manage the private inference gateway"} @@ -96,6 +97,7 @@ func init() { gatewayServeCmd.Flags().DurationVar(&gatewayCatalogTTL, "catalog-ttl", 5*time.Minute, "Refresh provider config and live model catalogs after this interval") gatewayServeCmd.Flags().IntVar(&gatewayRetryAttempts, "upstream-retry-attempts", gateway.DefaultUpstreamRetryAttempts, "Maximum upstream attempts per gateway inference request") gatewayServeCmd.Flags().DurationVar(&gatewayRetryElapsed, "upstream-retry-elapsed", gateway.DefaultUpstreamRetryElapsed, "Maximum elapsed time across upstream attempts") + gatewayServeCmd.Flags().DurationVar(&gatewayProviderSessionTimeout, "provider-session-idle-timeout", gateway.DefaultProviderSessionIdleTimeout, "Keep successful satellite WebSocket provider sessions warm for this idle duration (0 disables)") gatewayClientAddCmd.Flags().BoolVar(&gatewayClientAllowCLI, "allow-cli", false, "Allow this client to use CLI providers") gatewayClientAddCmd.Flags().StringSliceVar(&gatewayClientAllow, "allow-provider", nil, "Allowed provider key/prefix") @@ -133,6 +135,9 @@ func runGatewayServe(cmd *cobra.Command, _ []string) error { if gatewayRetryElapsed <= 0 { return fmt.Errorf("--upstream-retry-elapsed must be positive") } + if gatewayProviderSessionTimeout < 0 { + return fmt.Errorf("--provider-session-idle-timeout cannot be negative; use 0 to disable session reuse") + } stateDir, err := resolveGatewayStateDir() if err != nil { return err @@ -173,7 +178,9 @@ func runGatewayServe(cmd *cobra.Command, _ []string) error { Searcher: searcher, FetchTool: fetchTool, IdleTimeout: gatewayIdleTimeout, ToolTimeout: gatewayToolTimeout, CatalogTTL: gatewayCatalogTTL, UpstreamRetryAttempts: gatewayRetryAttempts, UpstreamRetryMaxElapsed: gatewayRetryElapsed, - RunTempRoot: filepath.Join(stateDir, "runs"), + ProviderSessionIdleTimeout: gatewayProviderSessionTimeout, + DisableProviderSessionReuse: gatewayProviderSessionTimeout == 0, + RunTempRoot: filepath.Join(stateDir, "runs"), Policy: gateway.Policy{ AllowProviders: gatewayAllowProviders, DenyProviders: gatewayDenyProviders, AllowModels: gatewayAllowModels, DenyModels: gatewayDenyModels, @@ -183,6 +190,7 @@ func runGatewayServe(cmd *cobra.Command, _ []string) error { if err != nil { return err } + defer server.Close() httpServer := &http.Server{Addr: gatewayListen, Handler: server.Handler(), ReadHeaderTimeout: 10 * time.Second, IdleTimeout: 2 * time.Minute} go func() { <-cmd.Context().Done() diff --git a/cmd/gateway_session_test.go b/cmd/gateway_session_test.go new file mode 100644 index 000000000..d33dacbff --- /dev/null +++ b/cmd/gateway_session_test.go @@ -0,0 +1,19 @@ +package cmd + +import ( + "strings" + "testing" +) + +func TestGatewayProviderSessionIdleTimeoutFlag(t *testing.T) { + flag := gatewayServeCmd.Flags().Lookup("provider-session-idle-timeout") + if flag == nil { + t.Fatal("gateway serve is missing --provider-session-idle-timeout") + } + if flag.DefValue != "30s" { + t.Fatalf("provider session idle timeout default = %q, want 30s", flag.DefValue) + } + if !strings.Contains(flag.Usage, "0 disables") || !strings.Contains(flag.Usage, "WebSocket") { + t.Fatalf("provider session idle timeout help is incomplete: %q", flag.Usage) + } +} diff --git a/docs-site/content/guides/inference-gateway.md b/docs-site/content/guides/inference-gateway.md index a8a6b0630..00c77be12 100644 --- a/docs-site/content/guides/inference-gateway.md +++ b/docs-site/content/guides/inference-gateway.md @@ -36,7 +36,15 @@ Important invariants: ### Session continuity and direct-provider equivalence -For an ordinary satellite session, the satellite remains the source of truth for `SessionID` and the complete transcript. It stores provider state under that session plus provider key, includes the same `SessionID` and provider-facing history on later turns, and can export/import the `GatewayProvider`'s opaque sealed blob when reconstructing a runtime. The gateway creates a fresh central provider for each request. Providers that implement provider-state export/import round-trip that state through the session-bound sealed blob; providers that do not are reconstructed from the complete transcript. +For an ordinary satellite session, the satellite remains the source of truth for `SessionID` and the complete transcript. It stores provider state under that session plus provider key, includes the same `SessionID` and provider-facing history on later turns, and can export/import the `GatewayProvider`'s opaque sealed blob when reconstructing a runtime. Providers that implement provider-state export/import continue to round-trip that state through the session-bound sealed blob; providers that do not can always reconstruct from the complete transcript. + +For sessionful satellite `/g1` inference, the gateway keeps a successfully completed central provider warm for **30 seconds** only when the request has a non-empty `SessionID`, is not ephemeral, and that explicit central OpenAI/ChatGPT provider has `use_websocket: true`. A lease is isolated by authenticated gateway client ID, provider key, and satellite `SessionID`, and turns for one lease are serialized. Other clients, providers, and sessions remain independent and concurrent. The model is deliberately not part of the lease key: this matches a direct provider instance and permits model changes, while `ResponsesClient` rejects an incompatible `previous_response_id` continuation and starts a full-history chain on the same connection when model or other non-input controls change. + +A warm follow-up uses the same provider and WebSocket, so ChatGPT sends the same `previous_response_id` plus continuation-only suffix as direct term-llm. The public `/v1/responses` edge does not opt into these private session semantics and remains stateless; unrelated Discourse requests are never attached to a retained satellite provider. WebSocket providers currently export no sealed provider state, so a warm provider is used directly rather than importing the satellite's prior blob a second time. Exportable-state providers and ordinary non-WebSocket providers retain the sealed per-request behavior above and are not added to this cache. + +Idle expiry is refreshed only after a fully successful turn. Cancellation, provider/stream failure, invalid continuation state, an incompatible config/policy decision, credential/config generation change, or gateway shutdown evicts the lease and calls provider conversation reset/cleanup through the retry wrapper. Idle eviction and shutdown close the WebSocket. The cache uses one bounded reaper rather than one timer or goroutine per session. + +After idle expiry or gateway restart, the next request creates a provider and sends the satellite's complete transcript without `previous_response_id`. This cold form is semantically correct and matches the direct provider's full-history fallback, but incurs a new WebSocket handshake and may have different first-request cache/latency characteristics. While the lease is warm, direct and gateway ChatGPT WebSocket upstream payloads are transport-equivalent. After intentional boundary normalization, the central provider receives the same model-semantic `llm.Request` as a direct provider: model and session identity; system/developer/user/assistant/tool history; tool definitions and choices; tool-call and tool-result continuation data; cache anchors; reasoning summaries, encrypted reasoning, and opaque provider replay; search controls; Responses options; sampling, output-token, service-tier, and ephemeral/continuation controls. This equivalence has deliberately narrow exceptions: @@ -44,10 +52,6 @@ After intentional boundary normalization, the central provider receives the same - Satellite image/file paths are removed while inline image/file data is retained. Tool-result diff/image paths, legacy display strings, formatted `ToolInfo`, provider tool-activity display records, skill-activation provenance, and persisted UI identity/segment fields are omitted. Expanded developer instructions, actual tool arguments/results, multimodal result data, caller provenance, thought signatures, reasoning, and provider replay remain. - Approval-only transcripts/roles, request-scoped execution filters (`AllowedTools`), satellite engine controls (`MaxTurns` and `ToolMap`), and debug flags are local and do not cross the provider boundary. The satellite engine consumes these before or around provider execution; omitting them cannot change the upstream model payload. -ChatGPT transport behavior needs one qualification. ChatGPT's HTTP/SSE path sends full history on every turn in both direct and gateway modes, so second-turn upstream request payloads are equivalent. A long-lived **direct WebSocket** provider may instead reuse its connection-local `previous_response_id` and send only the new continuation suffix. The gateway recreates the central provider (and therefore its WebSocket connection) per request, so it sends the full transcript without `previous_response_id`. Both forms are semantically equivalent, and rejection of a direct WebSocket parent falls back to the same full-history request. They are not strictly transport-, cache-, or latency-equivalent. - -Gateway restart and satellite session restoration remain correct when the client database and `state.key` are retained: the satellite replays its full transcript and restores any exportable provider state from the sealed blob. A restart does **not** retain ChatGPT's connection-local WebSocket optimization, so the next request uses full history and may have different cache/latency characteristics even though conversation semantics are preserved. - Run the gateway only on a trusted private network. For traffic that is not already protected by a private overlay, service mesh, or TLS reverse proxy, configure the built-in TLS listener with both `--tls-cert` and `--tls-key`. Bearer credentials authenticate access but plaintext HTTP does not encrypt them. ## Start a gateway @@ -72,6 +76,8 @@ term-llm gateway serve \ --state-dir /var/lib/term-llm-gateway ``` +Successful sessionful satellite WebSocket providers are retained for 30 seconds by default. Tune this with `--provider-session-idle-timeout DURATION`; `0` disables reuse, while negative durations are rejected. This setting does not make `/v1/responses` stateful. + Gateway state contains: ```text diff --git a/internal/gateway/execution.go b/internal/gateway/execution.go index 24e300fd7..dbbd6a8a7 100644 --- a/internal/gateway/execution.go +++ b/internal/gateway/execution.go @@ -6,8 +6,10 @@ import ( "log/slog" "net/http" "os" + "strings" "time" + "github.com/samsaffron/term-llm/internal/config" "github.com/samsaffron/term-llm/internal/gateway/protocol" "github.com/samsaffron/term-llm/internal/llm" ) @@ -16,7 +18,9 @@ import ( // the private /g1 transport and the public OpenAI-compatible Responses edge. // It applies catalog, policy, concurrency, credential, retry, state, filesystem // isolation, and cancellation rules without instantiating an agent runtime. -func (s *Server) startInference(parent context.Context, client Client, envelope protocol.InferenceRequest, providerReq llm.Request, rejectInlineToolLoop bool) (*inferenceExecution, *inferenceRequestError) { +// Provider sessions may be retained only when allowProviderSession is true; +// callers for the public Responses edge always pass false and remain stateless. +func (s *Server) startInference(parent context.Context, client Client, envelope protocol.InferenceRequest, providerReq llm.Request, rejectInlineToolLoop, allowProviderSession bool) (*inferenceExecution, *inferenceRequestError) { if providerReq.Model == "" { if pc := s.currentConfig().GetProviderConfig(envelope.Provider); pc != nil { providerReq.Model = pc.Model @@ -41,47 +45,92 @@ func (s *Server) startInference(parent context.Context, client Client, envelope entry, found, catalogErr := s.currentCatalogProvider(parent, envelope.Provider) if catalogErr != nil { slog.Error("refresh gateway provider catalog for inference", "request_id", envelope.RequestID, "provider", envelope.Provider, "error", catalogErr) + s.evictIncompatibleProviderSession(parent, client, envelope, providerReq, allowProviderSession) return fail(http.StatusServiceUnavailable, "catalog_unavailable", "gateway provider catalog is temporarily unavailable; retry or contact the gateway operator") } if !found { s.recordFailure(client, envelope, providerReq, "unknown_provider", time.Now().UTC()) + s.evictIncompatibleProviderSession(parent, client, envelope, providerReq, allowProviderSession) return fail(http.StatusNotFound, "unknown_provider", "provider is not available") } if !catalogEntryAllowsModel(entry, envelope.Provider, providerReq.Model) { s.recordFailure(client, envelope, providerReq, "unknown_model", time.Now().UTC()) + s.evictIncompatibleProviderSession(parent, client, envelope, providerReq, allowProviderSession) return fail(http.StatusNotFound, "unknown_model", "model is not available for this gateway provider; choose a catalog model or ask the gateway operator to allow unlisted models") } if !s.cfg.Policy.Allows(envelope.Provider, providerReq.Model, entry.CLI) || !client.Policy.Allows(envelope.Provider, providerReq.Model, entry.CLI) { s.recordFailure(client, envelope, providerReq, "policy_denied", time.Now().UTC()) + s.evictIncompatibleProviderSession(parent, client, envelope, providerReq, allowProviderSession) return fail(http.StatusForbidden, "policy_denied", "provider/model is denied by gateway policy; choose an allowed model or contact the gateway operator") } if rejectInlineToolLoop && len(providerReq.Tools) > 0 && entry.CLI && entry.Capabilities.InlineToolLoop { s.recordFailure(client, envelope, providerReq, "incompatible_tool_request", time.Now().UTC()) + s.evictIncompatibleProviderSession(parent, client, envelope, providerReq, allowProviderSession) return fail(http.StatusBadRequest, "incompatible_tool_request", "this CLI provider requires an inline tool loop; Responses function tools are rejected because the gateway never executes client tools") } started := time.Now().UTC() + ctx, cancel := context.WithCancel(parent) + contextStop := context.AfterFunc(s.lifecycleCtx, cancel) + var providerSession *providerSessionLease + var provider llm.Provider + cleanupProvider := false failStarted := func(status int, code, message string) (*inferenceExecution, *inferenceRequestError) { + if providerSession != nil { + providerSession.release(false) + } else if cleanupProvider && provider != nil { + cleanupProviderSession(provider) + } + contextStop() + cancel() s.recordUsage(client, envelope, providerReq, llm.Usage{}, code, started) return fail(status, code, message) } if err := nonInteractiveAuthReady(entry.Type); err != nil { status, code := classifyProviderError(err, entry.Type) slog.Error("gateway provider authentication unavailable", "request_id", envelope.RequestID, "provider", envelope.Provider, "error", err) + s.evictIncompatibleProviderSession(ctx, client, envelope, providerReq, allowProviderSession) return failStarted(status, code, safeProviderErrorMessage(code, envelope.Provider)) } - provider, err := s.cfg.ProviderFactory(s.centralConfig(), envelope.Provider, providerReq.Model) - if err != nil { - status, code := classifyProviderError(err, entry.Type) - slog.Error("create gateway provider", "request_id", envelope.RequestID, "provider", envelope.Provider, "error", err) - return failStarted(status, code, safeProviderErrorMessage(code, envelope.Provider)) + + sessionKey, fingerprint, sessionEligible := s.providerSessionSpec(client, envelope, providerReq, allowProviderSession) + if sessionKey != (providerSessionKey{}) && !sessionEligible { + if err := s.providerSessions.evict(ctx, sessionKey); err != nil { + return failStarted(499, "canceled", "gateway request was canceled") + } + } + if sessionEligible { + var err error + providerSession, err = s.providerSessions.acquire(ctx, sessionKey, fingerprint) + if err != nil { + return failStarted(499, "canceled", "gateway request was canceled") + } + provider = providerSession.provider + if providerSession.reused && envelope.State != "" { + return failStarted(http.StatusBadRequest, "invalid_state", "provider state must not be re-imported into a live gateway WebSocket session") + } + } + + if provider == nil { + var err error + provider, err = s.cfg.ProviderFactory(s.centralConfig(), envelope.Provider, providerReq.Model) + if err != nil { + status, code := classifyProviderError(err, entry.Type) + slog.Error("create gateway provider", "request_id", envelope.RequestID, "provider", envelope.Provider, "error", err) + return failStarted(status, code, safeProviderErrorMessage(code, envelope.Provider)) + } + provider = llm.WrapWithRetry(provider, llm.RetryConfig{ + MaxAttempts: s.cfg.UpstreamRetryAttempts, + MaxElapsedTime: s.cfg.UpstreamRetryMaxElapsed, + BaseBackoff: time.Second, + MaxBackoff: 5 * time.Second, + }) + if providerSession != nil { + providerSession.setProvider(provider, fingerprint) + } else { + cleanupProvider = true + } } - provider = llm.WrapWithRetry(provider, llm.RetryConfig{ - MaxAttempts: s.cfg.UpstreamRetryAttempts, - MaxElapsedTime: s.cfg.UpstreamRetryMaxElapsed, - BaseBackoff: time.Second, - MaxBackoff: 5 * time.Second, - }) if envelope.State != "" { plain, openErr := s.cfg.Sealer.Open(envelope.State, client.ID, envelope.Provider, providerReq.SessionID) if openErr != nil { @@ -101,10 +150,8 @@ func (s *Server) startInference(parent context.Context, client Client, envelope return failStarted(http.StatusInternalServerError, "internal", "could not create isolated provider directory") } providerReq.WorkingDir = tempDir - ctx, cancel := context.WithCancel(parent) stream, err := provider.Stream(ctx, providerReq) if err != nil { - cancel() _ = os.RemoveAll(tempDir) status, code := classifyProviderError(err, entry.Type) slog.Error("gateway provider request failed", "request_id", envelope.RequestID, "provider", envelope.Provider, "error", err) @@ -113,11 +160,37 @@ func (s *Server) startInference(parent context.Context, client Client, envelope return &inferenceExecution{ server: s, client: client, envelope: envelope, request: providerReq, entry: entry, - provider: provider, stream: stream, ctx: ctx, cancel: cancel, release: releaseInference, + provider: provider, stream: stream, ctx: ctx, cancel: cancel, contextStop: contextStop, + providerSession: providerSession, cleanupProvider: cleanupProvider, release: releaseInference, tempDir: tempDir, started: started, }, nil } +func (s *Server) providerSessionSpec(client Client, envelope protocol.InferenceRequest, req llm.Request, allow bool) (providerSessionKey, [32]byte, bool) { + var zero [32]byte + if s.providerSessions == nil || !allow || strings.TrimSpace(req.SessionID) == "" { + return providerSessionKey{}, zero, false + } + key := providerSessionKey{clientID: client.ID, provider: envelope.Provider, sessionID: req.SessionID} + cfg, generation := s.currentConfigAndGeneration() + providerConfig, ok := cfg.Providers[envelope.Provider] + if !ok || !cfg.IsExplicitProvider(envelope.Provider) || req.Ephemeral || !providerConfig.UseWebSocket { + return key, zero, false + } + providerType := config.InferProviderType(envelope.Provider, providerConfig.Type) + if providerType != config.ProviderTypeOpenAI && providerType != config.ProviderTypeChatGPT { + return key, zero, false + } + return key, providerSessionFingerprint(envelope.Provider, providerType, providerConfig, generation), true +} + +func (s *Server) evictIncompatibleProviderSession(ctx context.Context, client Client, envelope protocol.InferenceRequest, req llm.Request, allow bool) { + key, _, _ := s.providerSessionSpec(client, envelope, req, allow) + if key != (providerSessionKey{}) { + _ = s.providerSessions.evict(ctx, key) + } +} + func (s *Server) logProviderStreamError(execution *inferenceExecution, err error) (int, string) { status, code := classifyProviderError(err, execution.entry.Type) if errors.Is(err, context.Canceled) || errors.Is(execution.ctx.Err(), context.Canceled) { diff --git a/internal/gateway/provider_sessions.go b/internal/gateway/provider_sessions.go new file mode 100644 index 000000000..522c85f10 --- /dev/null +++ b/internal/gateway/provider_sessions.go @@ -0,0 +1,266 @@ +package gateway + +import ( + "context" + "crypto/sha256" + "encoding/json" + "fmt" + "sync" + "time" + + "github.com/samsaffron/term-llm/internal/config" + "github.com/samsaffron/term-llm/internal/llm" +) + +type providerSessionKey struct { + clientID string + provider string + sessionID string +} + +type providerSessionEntry struct { + token chan struct{} + provider llm.Provider + fingerprint [sha256.Size]byte + expiresAt time.Time +} + +type providerSessionLease struct { + cache *providerSessionCache + key providerSessionKey + entry *providerSessionEntry + provider llm.Provider + reused bool + once sync.Once +} + +type providerSessionCache struct { + idleTimeout time.Duration + + mu sync.Mutex + entries map[providerSessionKey]*providerSessionEntry + closed bool + reaperStarted bool + stop chan struct{} + done chan struct{} + close sync.Once +} + +func newProviderSessionCache(idleTimeout time.Duration) *providerSessionCache { + cache := &providerSessionCache{ + idleTimeout: idleTimeout, + entries: make(map[providerSessionKey]*providerSessionEntry), + stop: make(chan struct{}), + done: make(chan struct{}), + } + return cache +} + +func (c *providerSessionCache) acquire(ctx context.Context, key providerSessionKey, fingerprint [sha256.Size]byte) (*providerSessionLease, error) { + for { + c.mu.Lock() + if c.closed { + c.mu.Unlock() + return nil, fmt.Errorf("gateway provider session cache is closed") + } + entry := c.entries[key] + if entry == nil { + entry = &providerSessionEntry{token: make(chan struct{}, 1)} + entry.token <- struct{}{} + c.entries[key] = entry + } + c.mu.Unlock() + + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-entry.token: + } + + c.mu.Lock() + current := !c.closed && c.entries[key] == entry + c.mu.Unlock() + if !current { + entry.token <- struct{}{} + continue + } + + reused := entry.provider != nil && entry.fingerprint == fingerprint && time.Now().Before(entry.expiresAt) + if entry.provider != nil && !reused { + cleanupProviderSession(entry.provider) + entry.provider = nil + entry.expiresAt = time.Time{} + } + return &providerSessionLease{cache: c, key: key, entry: entry, provider: entry.provider, reused: reused}, nil + } +} + +func (l *providerSessionLease) setProvider(provider llm.Provider, fingerprint [sha256.Size]byte) { + l.provider = provider + l.entry.provider = provider + l.entry.fingerprint = fingerprint +} + +func (l *providerSessionLease) release(success bool) { + if l == nil { + return + } + l.once.Do(func() { + l.cache.mu.Lock() + current := !l.cache.closed && l.cache.entries[l.key] == l.entry + if !success || !current { + if l.cache.entries[l.key] == l.entry { + delete(l.cache.entries, l.key) + } + } + l.cache.mu.Unlock() + + if success && current { + l.entry.expiresAt = time.Now().Add(l.cache.idleTimeout) + l.cache.startReaper() + } else if l.entry.provider != nil { + cleanupProviderSession(l.entry.provider) + l.entry.provider = nil + l.entry.expiresAt = time.Time{} + } + l.entry.token <- struct{}{} + }) +} + +func (c *providerSessionCache) evict(ctx context.Context, key providerSessionKey) error { + c.mu.Lock() + entry := c.entries[key] + c.mu.Unlock() + if entry == nil { + return nil + } + + select { + case <-ctx.Done(): + return ctx.Err() + case <-entry.token: + } + c.mu.Lock() + if c.entries[key] == entry { + delete(c.entries, key) + } + c.mu.Unlock() + if entry.provider != nil { + cleanupProviderSession(entry.provider) + entry.provider = nil + entry.expiresAt = time.Time{} + } + entry.token <- struct{}{} + return nil +} + +func (c *providerSessionCache) startReaper() { + c.mu.Lock() + defer c.mu.Unlock() + if c.closed || c.reaperStarted { + return + } + c.reaperStarted = true + go c.reap() +} + +func (c *providerSessionCache) reap() { + defer close(c.done) + interval := c.idleTimeout / 2 + if interval > time.Second { + interval = time.Second + } + if interval < time.Millisecond { + interval = time.Millisecond + } + ticker := time.NewTicker(interval) + defer ticker.Stop() + for { + select { + case <-c.stop: + return + case <-ticker.C: + c.reapExpired(time.Now()) + } + } +} + +func (c *providerSessionCache) reapExpired(now time.Time) { + c.mu.Lock() + entries := make(map[providerSessionKey]*providerSessionEntry, len(c.entries)) + for key, entry := range c.entries { + entries[key] = entry + } + c.mu.Unlock() + + for key, entry := range entries { + select { + case <-entry.token: + c.mu.Lock() + expired := c.entries[key] == entry && entry.provider != nil && !now.Before(entry.expiresAt) + if expired { + delete(c.entries, key) + } + c.mu.Unlock() + if expired { + cleanupProviderSession(entry.provider) + entry.provider = nil + entry.expiresAt = time.Time{} + } + entry.token <- struct{}{} + default: + } + } +} + +func (c *providerSessionCache) Close() { + if c == nil { + return + } + c.close.Do(func() { + c.mu.Lock() + c.closed = true + reaperStarted := c.reaperStarted + entries := make([]*providerSessionEntry, 0, len(c.entries)) + for key, entry := range c.entries { + entries = append(entries, entry) + delete(c.entries, key) + } + c.mu.Unlock() + + if reaperStarted { + close(c.stop) + <-c.done + } + + for _, entry := range entries { + <-entry.token + if entry.provider != nil { + cleanupProviderSession(entry.provider) + entry.provider = nil + } + entry.token <- struct{}{} + } + }) +} + +func cleanupProviderSession(provider llm.Provider) { + if resetter, ok := provider.(interface{ ResetConversation() }); ok { + resetter.ResetConversation() + } + if cleaner, ok := provider.(llm.ProviderCleaner); ok { + cleaner.CleanupMCP() + } +} + +func providerSessionFingerprint(providerKey string, providerType config.ProviderType, providerConfig config.ProviderConfig, generation uint64) [sha256.Size]byte { + data, _ := json.Marshal(struct { + ProviderKey string + ProviderType config.ProviderType + Config config.ProviderConfig + Generation uint64 + }{ + ProviderKey: providerKey, ProviderType: providerType, Config: providerConfig, Generation: generation, + }) + return sha256.Sum256(data) +} diff --git a/internal/gateway/provider_sessions_test.go b/internal/gateway/provider_sessions_test.go new file mode 100644 index 000000000..9c338f073 --- /dev/null +++ b/internal/gateway/provider_sessions_test.go @@ -0,0 +1,795 @@ +package gateway + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptest" + "reflect" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/gorilla/websocket" + "github.com/samsaffron/term-llm/internal/config" + "github.com/samsaffron/term-llm/internal/gateway/protocol" + "github.com/samsaffron/term-llm/internal/llm" +) + +type leaseWebSocketRecorder struct { + server *httptest.Server + + mu sync.Mutex + requests []map[string]any + connections int + active int + closed int + handshakeDelay time.Duration + block <-chan struct{} + requestSeen chan struct{} + seenOnce sync.Once +} + +func newLeaseWebSocketRecorder(t *testing.T, handshakeDelay time.Duration) *leaseWebSocketRecorder { + t.Helper() + recorder := &leaseWebSocketRecorder{handshakeDelay: handshakeDelay, requestSeen: make(chan struct{})} + upgrader := websocket.Upgrader{} + recorder.server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if recorder.handshakeDelay > 0 { + timer := time.NewTimer(recorder.handshakeDelay) + select { + case <-r.Context().Done(): + timer.Stop() + return + case <-timer.C: + } + } + conn, err := upgrader.Upgrade(w, r, nil) + if err != nil { + return + } + recorder.mu.Lock() + recorder.connections++ + recorder.active++ + recorder.mu.Unlock() + defer func() { + _ = conn.Close() + recorder.mu.Lock() + recorder.active-- + recorder.closed++ + recorder.mu.Unlock() + }() + for { + _, data, err := conn.ReadMessage() + if err != nil { + return + } + var payload map[string]any + if json.Unmarshal(data, &payload) != nil { + return + } + recorder.mu.Lock() + recorder.requests = append(recorder.requests, payload) + block := recorder.block + recorder.mu.Unlock() + recorder.seenOnce.Do(func() { close(recorder.requestSeen) }) + if block != nil { + <-block + } + if strings.Contains(string(data), `"fail"`) { + _ = conn.WriteJSON(map[string]any{"type": "response.output_text.delta", "delta": "partial"}) + _ = conn.WriteJSON(map[string]any{ + "type": "response.failed", "status": http.StatusBadGateway, + "response": map[string]any{"error": map[string]any{"code": "upstream_failed", "message": "failed"}}, + }) + continue + } + _ = conn.WriteJSON(map[string]any{"type": "response.output_text.delta", "delta": "ok"}) + _ = conn.WriteJSON(map[string]any{ + "type": "response.completed", + "response": map[string]any{ + "id": "resp_parent", + "usage": map[string]any{"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}, + }, + }) + } + })) + t.Cleanup(func() { recorder.server.Close() }) + return recorder +} + +func (r *leaseWebSocketRecorder) snapshot() (requests []map[string]any, connections, active, closed int) { + r.mu.Lock() + defer r.mu.Unlock() + requests = make([]map[string]any, len(r.requests)) + for i, request := range r.requests { + copyRequest := make(map[string]any, len(request)) + for key, value := range request { + copyRequest[key] = value + } + requests[i] = copyRequest + } + return requests, r.connections, r.active, r.closed +} + +type leaseWebSocketProvider struct { + client *llm.ResponsesClient + cleanups *atomic.Int32 +} + +func newLeaseWebSocketProvider(recorder *leaseWebSocketRecorder, cleanups *atomic.Int32) *leaseWebSocketProvider { + return &leaseWebSocketProvider{ + client: &llm.ResponsesClient{ + BaseURL: recorder.server.URL, HTTPClient: recorder.server.Client(), UseWebSocket: true, + WebSocketServerState: true, DisableServerState: true, + }, + cleanups: cleanups, + } +} + +func (*leaseWebSocketProvider) Name() string { return "lease-websocket" } +func (*leaseWebSocketProvider) Credential() string { return "mock" } +func (*leaseWebSocketProvider) Capabilities() llm.Capabilities { return llm.Capabilities{} } +func (p *leaseWebSocketProvider) Stream(ctx context.Context, req llm.Request) (llm.Stream, error) { + return p.client.Stream(ctx, llm.ResponsesRequest{ + Model: req.Model, Messages: req.Messages, Stream: true, SessionID: req.SessionID, + }, false) +} +func (p *leaseWebSocketProvider) ResetConversation() { + p.client.ResetConversation() + if p.cleanups != nil { + p.cleanups.Add(1) + } +} + +type leaseGatewayFixture struct { + server *Server + httpServer *httptest.Server + config *config.Config + clients *ClientStore + client Client + token string + factory ProviderFactory + factoryRuns atomic.Int32 +} + +func newLeaseGatewayFixture(t *testing.T, idle time.Duration, providerKeys []string, factory ProviderFactory) *leaseGatewayFixture { + t.Helper() + dir := t.TempDir() + clients, err := OpenClientStore(dir + "/clients.json") + if err != nil { + t.Fatal(err) + } + client, token, err := clients.Add("satellite-a", Policy{AllowProviders: providerKeys, MaxConcurrentInference: 8}) + if err != nil { + t.Fatal(err) + } + sealer, err := OpenStateSealer(dir + "/state.key") + if err != nil { + t.Fatal(err) + } + providers := make(map[string]config.ProviderConfig, len(providerKeys)) + for _, key := range providerKeys { + providers[key] = config.ProviderConfig{ + Type: config.ProviderTypeOpenAI, Model: "model-a", Models: []string{"model-a", "model-b"}, + APIKey: "test-key", UseWebSocket: true, + } + } + central := &config.Config{DefaultProvider: providerKeys[0], Providers: providers} + fixture := &leaseGatewayFixture{config: central, clients: clients, client: client, token: token, factory: factory} + server, err := NewServer(ServerConfig{ + Config: central, Clients: clients, Sealer: sealer, ProviderSessionIdleTimeout: idle, + DisableProviderSessionReuse: idle == 0, + ProviderFactory: func(cfg *config.Config, provider, model string) (llm.Provider, error) { + fixture.factoryRuns.Add(1) + return factory(cfg, provider, model) + }, + Policy: Policy{AllowProviders: providerKeys}, UpstreamRetryAttempts: 1, + }) + if err != nil { + t.Fatal(err) + } + fixture.server = server + for _, key := range providerKeys { + server.storeCatalogProvider(protocol.CatalogEntry{ + Key: key, Type: string(config.ProviderTypeOpenAI), AllowUnlistedModels: false, + Models: []protocol.Model{{ID: "model-a"}, {ID: "model-b"}}, + }) + } + fixture.httpServer = httptest.NewServer(server.Handler()) + t.Cleanup(func() { + fixture.httpServer.Close() + server.Close() + }) + return fixture +} + +func (f *leaseGatewayFixture) satellite(t *testing.T, provider, model string) llm.Provider { + t.Helper() + providerClient, err := llm.NewGatewayProvider(&config.Config{Gateway: config.GatewayConfig{ + URL: f.httpServer.URL, Token: f.token, ConnectTimeout: "2s", ResponseTimeout: "2s", ToolTimeout: "2s", + }}, provider, model) + if err != nil { + t.Fatal(err) + } + return providerClient +} + +func runLeaseTurn(t *testing.T, provider llm.Provider, ctx context.Context, req llm.Request) error { + t.Helper() + stream, err := provider.Stream(ctx, req) + if err != nil { + return err + } + defer stream.Close() + for { + _, err = stream.Recv() + if err == io.EOF { + return nil + } + if err != nil { + return err + } + } +} + +func leaseTurns(sessionID string) (llm.Request, llm.Request) { + first := llm.Request{Model: "model-a", SessionID: sessionID, Messages: []llm.Message{llm.UserText("first")}} + second := llm.Request{Model: "model-a", SessionID: sessionID, Messages: []llm.Message{ + llm.UserText("first"), llm.AssistantText("answer"), llm.UserText("second"), + }} + return first, second +} + +type leaseCleanupProvider struct { + resets atomic.Int32 + cleanups atomic.Int32 +} + +func (*leaseCleanupProvider) Name() string { return "lease-cleanup" } +func (*leaseCleanupProvider) Credential() string { return "mock" } +func (*leaseCleanupProvider) Capabilities() llm.Capabilities { return llm.Capabilities{} } +func (*leaseCleanupProvider) Stream(context.Context, llm.Request) (llm.Stream, error) { + return nil, errors.New("unused") +} +func (p *leaseCleanupProvider) ResetConversation() { p.resets.Add(1) } +func (p *leaseCleanupProvider) CleanupMCP() { p.cleanups.Add(1) } + +func TestGatewayProviderSessionCleanupUsesRetryForwarders(t *testing.T) { + provider := &leaseCleanupProvider{} + cleanupProviderSession(llm.WrapWithRetry(provider, llm.RetryConfig{MaxAttempts: 1})) + if provider.resets.Load() != 1 || provider.cleanups.Load() != 1 { + t.Fatalf("provider reset/cleanup calls = %d/%d, want 1/1", provider.resets.Load(), provider.cleanups.Load()) + } +} + +func TestGatewayProviderSessionZeroDisablesReuse(t *testing.T) { + recorder := newLeaseWebSocketRecorder(t, 0) + fixture := newLeaseGatewayFixture(t, 0, []string{"remote"}, func(*config.Config, string, string) (llm.Provider, error) { + return newLeaseWebSocketProvider(recorder, nil), nil + }) + provider := fixture.satellite(t, "remote", "model-a") + first, second := leaseTurns("disabled-session") + if err := runLeaseTurn(t, provider, t.Context(), first); err != nil { + t.Fatal(err) + } + if err := runLeaseTurn(t, provider, t.Context(), second); err != nil { + t.Fatal(err) + } + _, connections, _, _ := recorder.snapshot() + if fixture.factoryRuns.Load() != 2 || connections != 2 { + t.Fatalf("zero timeout providers/connections = %d/%d, want 2/2", fixture.factoryRuns.Load(), connections) + } +} + +func TestGatewayProviderSessionNegativeTimeoutRejected(t *testing.T) { + dir := t.TempDir() + clients, err := OpenClientStore(dir + "/clients.json") + if err != nil { + t.Fatal(err) + } + sealer, err := OpenStateSealer(dir + "/state.key") + if err != nil { + t.Fatal(err) + } + defaulted, err := NewServer(ServerConfig{ + Config: &config.Config{Providers: map[string]config.ProviderConfig{}}, Clients: clients, Sealer: sealer, + }) + if err != nil { + t.Fatal(err) + } + if defaulted.providerSessions == nil || defaulted.providerSessions.idleTimeout != DefaultProviderSessionIdleTimeout { + t.Fatalf("ServerConfig zero idle timeout = %#v, want default %s", defaulted.providerSessions, DefaultProviderSessionIdleTimeout) + } + defaulted.Close() + + _, err = NewServer(ServerConfig{ + Config: &config.Config{Providers: map[string]config.ProviderConfig{}}, Clients: clients, Sealer: sealer, + ProviderSessionIdleTimeout: -time.Second, + }) + if err == nil || !strings.Contains(err.Error(), "cannot be negative") { + t.Fatalf("negative provider session timeout error = %v", err) + } +} + +func TestGatewayProviderSessionWarmReuseAndIdleExpiry(t *testing.T) { + const handshakeDelay = 25 * time.Millisecond + recorder := newLeaseWebSocketRecorder(t, handshakeDelay) + var cleanups atomic.Int32 + fixture := newLeaseGatewayFixture(t, 60*time.Millisecond, []string{"remote"}, func(*config.Config, string, string) (llm.Provider, error) { + return newLeaseWebSocketProvider(recorder, &cleanups), nil + }) + provider := fixture.satellite(t, "remote", "model-a") + first, second := leaseTurns("session-warm") + + firstStarted := time.Now() + if err := runLeaseTurn(t, provider, t.Context(), first); err != nil { + t.Fatal(err) + } + firstElapsed := time.Since(firstStarted) + warmStarted := time.Now() + if err := runLeaseTurn(t, provider, t.Context(), second); err != nil { + t.Fatal(err) + } + warmElapsed := time.Since(warmStarted) + + requests, connections, _, _ := recorder.snapshot() + if fixture.factoryRuns.Load() != 1 || connections != 1 || len(requests) != 2 { + t.Fatalf("warm provider/connections/requests = %d/%d/%d, want 1/1/2", fixture.factoryRuns.Load(), connections, len(requests)) + } + if requests[1]["previous_response_id"] != "resp_parent" { + t.Fatalf("warm previous_response_id = %#v, want resp_parent (providers=%d connections=%d requests=%#v)", requests[1]["previous_response_id"], fixture.factoryRuns.Load(), connections, requests) + } + if input, ok := requests[1]["input"].([]any); !ok || len(input) != 1 || !strings.Contains(fmt.Sprint(input[0]), "second") { + t.Fatalf("warm continuation input = %#v, want only latest suffix", requests[1]["input"]) + } + + time.Sleep(90 * time.Millisecond) + coldStarted := time.Now() + if err := runLeaseTurn(t, provider, t.Context(), second); err != nil { + t.Fatal(err) + } + coldElapsed := time.Since(coldStarted) + requests, connections, _, _ = recorder.snapshot() + if fixture.factoryRuns.Load() != 2 || connections != 2 || len(requests) != 3 { + t.Fatalf("expired provider/connections/requests = %d/%d/%d, want 2/2/3", fixture.factoryRuns.Load(), connections, len(requests)) + } + if _, ok := requests[2]["previous_response_id"]; ok { + t.Fatalf("expired request retained previous_response_id: %#v", requests[2]) + } + if input, ok := requests[2]["input"].([]any); !ok || len(input) != 3 { + t.Fatalf("expired request input = %#v, want full transcript", requests[2]["input"]) + } + if cleanups.Load() < 1 { + t.Fatal("idle expiry did not reset/close the retained provider") + } + t.Logf("fake live WebSocket timings: first=%s warm-follow-up=%s expired-cold=%s; handshakes warm=1, after expiry=2", firstElapsed, warmElapsed, coldElapsed) +} + +func TestGatewayProviderSessionKeyIsolation(t *testing.T) { + recorder := newLeaseWebSocketRecorder(t, 0) + fixture := newLeaseGatewayFixture(t, time.Second, []string{"alpha", "beta"}, func(*config.Config, string, string) (llm.Provider, error) { + return newLeaseWebSocketProvider(recorder, nil), nil + }) + otherClient, otherToken, err := fixture.clients.Add("satellite-b", Policy{AllowProviders: []string{"alpha", "beta"}, MaxConcurrentInference: 8}) + if err != nil { + t.Fatal(err) + } + _ = otherClient + otherProvider, err := llm.NewGatewayProvider(&config.Config{Gateway: config.GatewayConfig{ + URL: fixture.httpServer.URL, Token: otherToken, ConnectTimeout: "2s", ResponseTimeout: "2s", + }}, "alpha", "model-a") + if err != nil { + t.Fatal(err) + } + + cases := []struct { + provider llm.Provider + session string + }{ + {fixture.satellite(t, "alpha", "model-a"), "same-session"}, + {fixture.satellite(t, "alpha", "model-a"), "other-session"}, + {fixture.satellite(t, "beta", "model-a"), "same-session"}, + {otherProvider, "same-session"}, + } + for _, tc := range cases { + first, second := leaseTurns(tc.session) + if err := runLeaseTurn(t, tc.provider, t.Context(), first); err != nil { + t.Fatal(err) + } + if err := runLeaseTurn(t, tc.provider, t.Context(), second); err != nil { + t.Fatal(err) + } + } + requests, connections, _, _ := recorder.snapshot() + if fixture.factoryRuns.Load() != int32(len(cases)) || connections != len(cases) || len(requests) != 2*len(cases) { + t.Fatalf("isolated provider/connections/requests = %d/%d/%d, want %d/%d/%d", fixture.factoryRuns.Load(), connections, len(requests), len(cases), len(cases), 2*len(cases)) + } + for i := 1; i < len(requests); i += 2 { + if requests[i]["previous_response_id"] != "resp_parent" { + t.Fatalf("isolated continuation %d lost its own parent: %#v", i, requests[i]) + } + } +} + +func TestGatewayProviderSessionsDifferentKeysRemainConcurrent(t *testing.T) { + release := make(chan struct{}) + recorder := newLeaseWebSocketRecorder(t, 0) + recorder.block = release + fixture := newLeaseGatewayFixture(t, time.Second, []string{"remote"}, func(*config.Config, string, string) (llm.Provider, error) { + return newLeaseWebSocketProvider(recorder, nil), nil + }) + provider := fixture.satellite(t, "remote", "model-a") + errs := make(chan error, 2) + for _, sessionID := range []string{"session-a", "session-b"} { + sessionID := sessionID + go func() { + first, _ := leaseTurns(sessionID) + errs <- runLeaseTurn(t, provider, context.Background(), first) + }() + } + deadline := time.Now().Add(time.Second) + for { + requests, connections, _, _ := recorder.snapshot() + if len(requests) == 2 && connections == 2 { + break + } + if time.Now().After(deadline) { + close(release) + t.Fatalf("different keys did not reach upstream concurrently: requests=%d connections=%d", len(requests), connections) + } + time.Sleep(time.Millisecond) + } + close(release) + for range 2 { + if err := <-errs; err != nil { + t.Fatal(err) + } + } +} + +func TestGatewayProviderSessionSameKeySerializes(t *testing.T) { + release := make(chan struct{}) + recorder := newLeaseWebSocketRecorder(t, 0) + recorder.block = release + fixture := newLeaseGatewayFixture(t, time.Second, []string{"remote"}, func(*config.Config, string, string) (llm.Provider, error) { + return newLeaseWebSocketProvider(recorder, nil), nil + }) + provider := fixture.satellite(t, "remote", "model-a") + first, second := leaseTurns("serialized") + errs := make(chan error, 2) + go func() { errs <- runLeaseTurn(t, provider, context.Background(), first) }() + select { + case <-recorder.requestSeen: + case <-time.After(time.Second): + t.Fatal("first upstream request did not start") + } + go func() { errs <- runLeaseTurn(t, provider, context.Background(), second) }() + + timer := time.NewTimer(40 * time.Millisecond) + <-timer.C + requests, connections, _, _ := recorder.snapshot() + if len(requests) != 1 || connections != 1 { + t.Fatalf("same-key concurrent turn reached upstream early: requests=%d connections=%d", len(requests), connections) + } + close(release) + for range 2 { + if err := <-errs; err != nil { + t.Fatal(err) + } + } + requests, connections, _, _ = recorder.snapshot() + if len(requests) != 2 || connections != 1 || fixture.factoryRuns.Load() != 1 { + t.Fatalf("serialized final requests/connections/providers = %d/%d/%d", len(requests), connections, fixture.factoryRuns.Load()) + } +} + +func TestGatewayProviderSessionWarmStateImportIsRejectedAndEvicted(t *testing.T) { + recorder := newLeaseWebSocketRecorder(t, 0) + var cleanups atomic.Int32 + fixture := newLeaseGatewayFixture(t, time.Second, []string{"remote"}, func(*config.Config, string, string) (llm.Provider, error) { + return newLeaseWebSocketProvider(recorder, &cleanups), nil + }) + provider := fixture.satellite(t, "remote", "model-a") + first, second := leaseTurns("invalid-state-session") + if err := runLeaseTurn(t, provider, t.Context(), first); err != nil { + t.Fatal(err) + } + importer, ok := provider.(llm.ProviderStateImporter) + if !ok { + t.Fatal("GatewayProvider does not expose sealed state import") + } + if err := importer.ImportProviderState([]byte("stale-sealed-state")); err != nil { + t.Fatal(err) + } + if err := runLeaseTurn(t, provider, t.Context(), second); err == nil || !strings.Contains(err.Error(), "must not be re-imported") { + t.Fatalf("warm stale-state error = %v", err) + } + fresh := fixture.satellite(t, "remote", "model-a") + if err := runLeaseTurn(t, fresh, t.Context(), second); err != nil { + t.Fatal(err) + } + requests, connections, _, _ := recorder.snapshot() + if fixture.factoryRuns.Load() != 2 || connections != 2 || len(requests) != 2 { + t.Fatalf("invalid-state providers/connections/requests = %d/%d/%d, want 2/2/2", fixture.factoryRuns.Load(), connections, len(requests)) + } + if _, ok := requests[1]["previous_response_id"]; ok { + t.Fatalf("post-invalid-state request reused response state: %#v", requests[1]) + } + if cleanups.Load() < 1 { + t.Fatal("invalid warm state did not clean the retained provider") + } +} + +func TestGatewayProviderSessionFailureEvicts(t *testing.T) { + recorder := newLeaseWebSocketRecorder(t, 0) + var cleanups atomic.Int32 + fixture := newLeaseGatewayFixture(t, time.Second, []string{"remote"}, func(*config.Config, string, string) (llm.Provider, error) { + return newLeaseWebSocketProvider(recorder, &cleanups), nil + }) + provider := fixture.satellite(t, "remote", "model-a") + failed := llm.Request{Model: "model-a", SessionID: "failed-session", Messages: []llm.Message{llm.UserText("fail")}} + if err := runLeaseTurn(t, provider, t.Context(), failed); err == nil { + t.Fatal("provider failure unexpectedly succeeded") + } + _, second := leaseTurns("failed-session") + if err := runLeaseTurn(t, provider, t.Context(), second); err != nil { + t.Fatal(err) + } + requests, connections, _, _ := recorder.snapshot() + if fixture.factoryRuns.Load() != 2 || connections != 2 || len(requests) != 2 { + t.Fatalf("post-failure providers/connections/requests = %d/%d/%d, want 2/2/2", fixture.factoryRuns.Load(), connections, len(requests)) + } + if _, ok := requests[1]["previous_response_id"]; ok { + t.Fatalf("post-failure request reused suspect response state: %#v", requests[1]) + } + if cleanups.Load() < 1 { + t.Fatal("failed provider was not reset/closed") + } +} + +type cancelLeaseProvider struct { + start chan struct{} + blocking bool + resets *atomic.Int32 +} + +func (*cancelLeaseProvider) Name() string { return "cancel-lease" } +func (*cancelLeaseProvider) Credential() string { return "mock" } +func (*cancelLeaseProvider) Capabilities() llm.Capabilities { return llm.Capabilities{} } +func (p *cancelLeaseProvider) Stream(ctx context.Context, _ llm.Request) (llm.Stream, error) { + if p.start != nil { + close(p.start) + } + if p.blocking { + return &blockingStream{ctx: ctx, canceled: make(chan struct{})}, nil + } + return &oneEventStream{ctx: ctx, event: llm.Event{Type: llm.EventTextDelta, Text: "ok"}}, nil +} +func (p *cancelLeaseProvider) ResetConversation() { p.resets.Add(1) } + +func TestGatewayProviderSessionCancellationEvicts(t *testing.T) { + started := make(chan struct{}) + var instances atomic.Int32 + var resets atomic.Int32 + fixture := newLeaseGatewayFixture(t, time.Second, []string{"remote"}, func(*config.Config, string, string) (llm.Provider, error) { + instance := instances.Add(1) + return &cancelLeaseProvider{start: map[bool]chan struct{}{true: started}[instance == 1], blocking: instance == 1, resets: &resets}, nil + }) + provider := fixture.satellite(t, "remote", "model-a") + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { + done <- runLeaseTurn(t, provider, ctx, llm.Request{Model: "model-a", SessionID: "cancel-session", Messages: []llm.Message{llm.UserText("wait")}}) + }() + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("provider did not start") + } + cancel() + if err := <-done; err == nil || !errors.Is(err, context.Canceled) && !strings.Contains(err.Error(), "canceled") { + t.Fatalf("canceled turn error = %v", err) + } + _, second := leaseTurns("cancel-session") + if err := runLeaseTurn(t, provider, t.Context(), second); err != nil { + t.Fatal(err) + } + if instances.Load() != 2 || resets.Load() < 1 { + t.Fatalf("post-cancel instances/resets = %d/%d, want 2/>=1", instances.Load(), resets.Load()) + } +} + +func TestGatewayProviderSessionShutdownClosesWebSockets(t *testing.T) { + recorder := newLeaseWebSocketRecorder(t, 0) + var cleanups atomic.Int32 + fixture := newLeaseGatewayFixture(t, time.Second, []string{"remote"}, func(*config.Config, string, string) (llm.Provider, error) { + return newLeaseWebSocketProvider(recorder, &cleanups), nil + }) + provider := fixture.satellite(t, "remote", "model-a") + first, _ := leaseTurns("shutdown-session") + if err := runLeaseTurn(t, provider, t.Context(), first); err != nil { + t.Fatal(err) + } + fixture.server.Close() + deadline := time.Now().Add(time.Second) + for { + _, _, active, closed := recorder.snapshot() + if active == 0 && closed == 1 { + break + } + if time.Now().After(deadline) { + t.Fatalf("shutdown connection lifecycle active/closed = %d/%d", active, closed) + } + time.Sleep(time.Millisecond) + } + if cleanups.Load() != 1 { + t.Fatalf("shutdown provider resets = %d, want 1", cleanups.Load()) + } +} + +func TestGatewayProviderSessionDirectAndGatewayPayloadEquivalence(t *testing.T) { + recorder := newLeaseWebSocketRecorder(t, 0) + first, second := leaseTurns("equivalent-session") + + direct := newLeaseWebSocketProvider(recorder, nil) + if err := runLeaseTurn(t, direct, t.Context(), first); err != nil { + t.Fatal(err) + } + if err := runLeaseTurn(t, direct, t.Context(), second); err != nil { + t.Fatal(err) + } + direct.ResetConversation() + + fixture := newLeaseGatewayFixture(t, time.Second, []string{"remote"}, func(*config.Config, string, string) (llm.Provider, error) { + return newLeaseWebSocketProvider(recorder, nil), nil + }) + gatewayProvider := fixture.satellite(t, "remote", "model-a") + if err := runLeaseTurn(t, gatewayProvider, t.Context(), first); err != nil { + t.Fatal(err) + } + if err := runLeaseTurn(t, gatewayProvider, t.Context(), second); err != nil { + t.Fatal(err) + } + + requests, _, _, _ := recorder.snapshot() + if len(requests) != 4 { + t.Fatalf("payload count = %d, want 4", len(requests)) + } + for _, index := range []int{0, 1} { + if !reflect.DeepEqual(requests[index], requests[index+2]) { + t.Fatalf("direct/gateway payload %d differs\ndirect: %#v\ngateway: %#v", index, requests[index], requests[index+2]) + } + } + + fixture.server.Close() + restarted := newLeaseGatewayFixture(t, time.Second, []string{"remote"}, func(*config.Config, string, string) (llm.Provider, error) { + return newLeaseWebSocketProvider(recorder, nil), nil + }) + if err := runLeaseTurn(t, restarted.satellite(t, "remote", "model-a"), t.Context(), second); err != nil { + t.Fatal(err) + } + requests, _, _, _ = recorder.snapshot() + cold := requests[len(requests)-1] + if _, ok := cold["previous_response_id"]; ok { + t.Fatalf("cold restart retained previous_response_id: %#v", cold) + } + if input, ok := cold["input"].([]any); !ok || len(input) != 3 { + t.Fatalf("cold restart input = %#v, want full-history fallback", cold["input"]) + } +} + +func TestGatewayProviderSessionModelSwapKeepsLeaseButStartsSafeChain(t *testing.T) { + recorder := newLeaseWebSocketRecorder(t, 0) + fixture := newLeaseGatewayFixture(t, time.Second, []string{"remote"}, func(*config.Config, string, string) (llm.Provider, error) { + return newLeaseWebSocketProvider(recorder, nil), nil + }) + provider := fixture.satellite(t, "remote", "model-a") + first, second := leaseTurns("model-swap") + second.Model = "model-b" + if err := runLeaseTurn(t, provider, t.Context(), first); err != nil { + t.Fatal(err) + } + if err := runLeaseTurn(t, provider, t.Context(), second); err != nil { + t.Fatal(err) + } + requests, connections, _, _ := recorder.snapshot() + if fixture.factoryRuns.Load() != 1 || connections != 1 || len(requests) != 2 { + t.Fatalf("model swap providers/connections/requests = %d/%d/%d, want 1/1/2", fixture.factoryRuns.Load(), connections, len(requests)) + } + if _, ok := requests[1]["previous_response_id"]; ok { + t.Fatalf("incompatible model swap reused previous_response_id: %#v", requests[1]) + } + if input, ok := requests[1]["input"].([]any); !ok || len(input) != 3 { + t.Fatalf("model swap input = %#v, want direct-provider full-history chain", requests[1]["input"]) + } +} + +func TestGatewayResponsesEdgeNeverUsesSatelliteProviderSessions(t *testing.T) { + recorder := newLeaseWebSocketRecorder(t, 0) + fixture := newLeaseGatewayFixture(t, time.Second, []string{"remote"}, func(*config.Config, string, string) (llm.Provider, error) { + return newLeaseWebSocketProvider(recorder, nil), nil + }) + for _, prompt := range []string{"first", "second"} { + body := strings.NewReader(fmt.Sprintf(`{"model":"remote/model-a","input":%q}`, prompt)) + req, err := http.NewRequestWithContext(t.Context(), http.MethodPost, fixture.httpServer.URL+"/v1/responses", body) + if err != nil { + t.Fatal(err) + } + req.Header.Set("Authorization", "Bearer "+fixture.token) + req.Header.Set("Content-Type", "application/json") + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + data, _ := io.ReadAll(resp.Body) + _ = resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("Responses edge status/body = %d/%s", resp.StatusCode, data) + } + } + requests, connections, _, _ := recorder.snapshot() + if fixture.factoryRuns.Load() != 2 || connections != 2 || len(requests) != 2 { + t.Fatalf("stateless Responses providers/connections/requests = %d/%d/%d, want 2/2/2", fixture.factoryRuns.Load(), connections, len(requests)) + } + for _, request := range requests { + if _, ok := request["previous_response_id"]; ok { + t.Fatalf("stateless Responses request acquired gateway continuation: %#v", request) + } + } +} + +func TestGatewayProviderSessionEligibilityAndConfigFingerprint(t *testing.T) { + recorder := newLeaseWebSocketRecorder(t, 0) + fixture := newLeaseGatewayFixture(t, time.Second, []string{"remote"}, func(*config.Config, string, string) (llm.Provider, error) { + return newLeaseWebSocketProvider(recorder, nil), nil + }) + provider := fixture.satellite(t, "remote", "model-a") + first, second := leaseTurns("config-session") + if err := runLeaseTurn(t, provider, t.Context(), first); err != nil { + t.Fatal(err) + } + + fixture.server.configMu.Lock() + updated := *fixture.server.config + updated.Providers = make(map[string]config.ProviderConfig, len(fixture.server.config.Providers)) + for key, value := range fixture.server.config.Providers { + updated.Providers[key] = value + } + providerConfig := updated.Providers["remote"] + providerConfig.APIKey = "rotated-key" + updated.Providers["remote"] = providerConfig + fixture.server.config = &updated + fixture.server.configGeneration++ + fixture.server.configMu.Unlock() + + if err := runLeaseTurn(t, provider, t.Context(), second); err != nil { + t.Fatal(err) + } + requests, connections, _, _ := recorder.snapshot() + if connections != 2 || fixture.factoryRuns.Load() != 2 { + t.Fatalf("credential/config change connections/providers = %d/%d, want 2/2", connections, fixture.factoryRuns.Load()) + } + if _, ok := requests[1]["previous_response_id"]; ok { + t.Fatal("config-changed request reused stale continuation state") + } + + providerConfig.UseWebSocket = false + updated.Providers["remote"] = providerConfig + fixture.server.configMu.Lock() + fixture.server.config = &updated + fixture.server.configGeneration++ + fixture.server.configMu.Unlock() + second.Ephemeral = true + if err := runLeaseTurn(t, provider, t.Context(), second); err != nil { + t.Fatal(err) + } + if fixture.factoryRuns.Load() != 3 { + t.Fatalf("non-WebSocket/ephemeral request factory runs = %d, want 3", fixture.factoryRuns.Load()) + } +} diff --git a/internal/gateway/responses.go b/internal/gateway/responses.go index bc736048d..f887d2f01 100644 --- a/internal/gateway/responses.go +++ b/internal/gateway/responses.go @@ -468,13 +468,14 @@ func (s *Server) handleResponses(w http.ResponseWriter, r *http.Request, client return } envelope := protocol.InferenceRequest{Version: protocol.Version, RequestID: requestID, Provider: provider} - execution, requestErr := s.startInference(r.Context(), client, envelope, providerReq, true) + execution, requestErr := s.startInference(r.Context(), client, envelope, providerReq, true, false) if requestErr != nil { s.writeResponsesError(w, requestErr.Status, requestErr.Code, requestErr.Message, responsesErrorParam(requestErr.Code)) return } errorCode := "" - defer func() { execution.close(errorCode) }() + successful := false + defer func() { execution.close(errorCode, successful) }() responseID, err := randomSecret("resp", 16) if err != nil { @@ -583,6 +584,7 @@ func (s *Server) handleResponses(w http.ResponseWriter, r *http.Request, client if !request.Stream { s.writeResponsesJSON(w, http.StatusOK, &accumulator.document) } + successful = true } func responsesErrorParam(code string) string { diff --git a/internal/gateway/server.go b/internal/gateway/server.go index 33a4a7b80..53de69725 100644 --- a/internal/gateway/server.go +++ b/internal/gateway/server.go @@ -32,32 +32,37 @@ import ( ) const ( - defaultMaxBodyBytes = 64 << 20 - DefaultUpstreamRetryAttempts = 3 - DefaultUpstreamRetryElapsed = 20 * time.Second + defaultMaxBodyBytes = 64 << 20 + DefaultUpstreamRetryAttempts = 3 + DefaultUpstreamRetryElapsed = 20 * time.Second + DefaultProviderSessionIdleTimeout = 30 * time.Second ) type ProviderFactory func(*config.Config, string, string) (llm.Provider, error) type ConfigLoader func() (*config.Config, error) type ServerConfig struct { - Config *config.Config - ConfigLoader ConfigLoader - Clients *ClientStore - Sealer *StateSealer - Usage UsageRecorder - ProviderFactory ProviderFactory - Searcher search.Searcher - FetchTool *llm.ReadURLTool - Policy Policy - MaxBodyBytes int64 - IdleTimeout time.Duration - ToolTimeout time.Duration - CatalogTTL time.Duration - ModelListTimeout time.Duration - UpstreamRetryAttempts int - UpstreamRetryMaxElapsed time.Duration - RunTempRoot string + Config *config.Config + ConfigLoader ConfigLoader + Clients *ClientStore + Sealer *StateSealer + Usage UsageRecorder + ProviderFactory ProviderFactory + Searcher search.Searcher + FetchTool *llm.ReadURLTool + Policy Policy + MaxBodyBytes int64 + IdleTimeout time.Duration + ToolTimeout time.Duration + CatalogTTL time.Duration + ModelListTimeout time.Duration + UpstreamRetryAttempts int + UpstreamRetryMaxElapsed time.Duration + ProviderSessionIdleTimeout time.Duration + // DisableProviderSessionReuse explicitly disables reuse. A zero idle timeout + // otherwise selects DefaultProviderSessionIdleTimeout. + DisableProviderSessionReuse bool + RunTempRoot string } type runState struct { @@ -76,20 +81,23 @@ type clientLimits struct { } type inferenceExecution struct { - server *Server - client Client - envelope protocol.InferenceRequest - request llm.Request - entry protocol.CatalogEntry - provider llm.Provider - stream llm.Stream - ctx context.Context - cancel context.CancelFunc - release func() - tempDir string - started time.Time - total llm.Usage - finish sync.Once + server *Server + client Client + envelope protocol.InferenceRequest + request llm.Request + entry protocol.CatalogEntry + provider llm.Provider + stream llm.Stream + ctx context.Context + cancel context.CancelFunc + release func() + contextStop func() bool + providerSession *providerSessionLease + cleanupProvider bool + tempDir string + started time.Time + total llm.Usage + finish sync.Once } type inferenceRequestError struct { @@ -100,7 +108,7 @@ type inferenceRequestError struct { func (e *inferenceExecution) addUsage(use llm.Usage) { e.total.Add(use) } -func (e *inferenceExecution) close(errorCode string) { +func (e *inferenceExecution) close(errorCode string, successful bool) { if e == nil { return } @@ -108,11 +116,20 @@ func (e *inferenceExecution) close(errorCode string) { if errorCode == "" && errors.Is(e.ctx.Err(), context.Canceled) { errorCode = "canceled" } + successful = successful && errorCode == "" && e.ctx.Err() == nil _ = e.stream.Close() e.cancel() + if e.contextStop != nil { + e.contextStop() + } if e.tempDir != "" { _ = os.RemoveAll(e.tempDir) } + if e.providerSession != nil { + e.providerSession.release(successful) + } else if e.cleanupProvider { + cleanupProviderSession(e.provider) + } if e.release != nil { e.release() } @@ -123,8 +140,14 @@ func (e *inferenceExecution) close(errorCode string) { type Server struct { cfg ServerConfig - configMu sync.RWMutex - config *config.Config + configMu sync.RWMutex + config *config.Config + configGeneration uint64 + + providerSessions *providerSessionCache + lifecycleCtx context.Context + lifecycleCancel context.CancelFunc + closeOnce sync.Once catalogMu sync.RWMutex catalog protocol.Catalog @@ -169,17 +192,44 @@ func NewServer(cfg ServerConfig) (*Server, error) { if cfg.UpstreamRetryMaxElapsed <= 0 { cfg.UpstreamRetryMaxElapsed = DefaultUpstreamRetryElapsed } + providerSessionIdleTimeout := cfg.ProviderSessionIdleTimeout + if providerSessionIdleTimeout < 0 { + return nil, fmt.Errorf("gateway provider session idle timeout cannot be negative") + } + if providerSessionIdleTimeout == 0 { + providerSessionIdleTimeout = DefaultProviderSessionIdleTimeout + } if strings.TrimSpace(cfg.RunTempRoot) != "" { cfg.RunTempRoot = filepath.Clean(cfg.RunTempRoot) if err := prepareRunTempRoot(cfg.RunTempRoot); err != nil { return nil, err } } - return &Server{ - cfg: cfg, config: cfg.Config, configFetchedAt: time.Now().UTC(), + lifecycleCtx, lifecycleCancel := context.WithCancel(context.Background()) + server := &Server{ + cfg: cfg, config: cfg.Config, configGeneration: 1, configFetchedAt: time.Now().UTC(), catalog: protocol.Catalog{Version: protocol.Version}, catalogProviderFetched: make(map[string]time.Time), runs: make(map[string]*runState), limits: make(map[string]*clientLimits), - }, nil + lifecycleCtx: lifecycleCtx, lifecycleCancel: lifecycleCancel, + } + if !cfg.DisableProviderSessionReuse { + server.providerSessions = newProviderSessionCache(providerSessionIdleTimeout) + } + return server, nil +} + +// Close cancels active gateway inference and closes all retained provider +// sessions, including their WebSocket and provider-specific resources. +func (s *Server) Close() { + if s == nil { + return + } + s.closeOnce.Do(func() { + s.lifecycleCancel() + if s.providerSessions != nil { + s.providerSessions.Close() + } + }) } func (s *Server) Handler() http.Handler { return http.HandlerFunc(s.serveHTTP) } @@ -339,13 +389,14 @@ func (s *Server) handleInference(w http.ResponseWriter, r *http.Request, client s.writeError(w, http.StatusBadRequest, "invalid_request", err.Error(), envelope.RequestID) return } - execution, requestErr := s.startInference(r.Context(), client, envelope, providerReq, false) + execution, requestErr := s.startInference(r.Context(), client, envelope, providerReq, false, true) if requestErr != nil { s.writeError(w, requestErr.Status, requestErr.Code, requestErr.Message, envelope.RequestID) return } errorCode := "" - defer func() { execution.close(errorCode) }() + successful := false + defer func() { execution.close(errorCode, successful) }() providerReq = execution.request provider := execution.provider stream := execution.stream @@ -484,8 +535,12 @@ func (s *Server) handleInference(w http.ResponseWriter, r *http.Request, client } } if errorCode == "" { - _ = writeSSE(w, protocol.StreamRecord{Version: protocol.Version, Type: "done", RequestID: envelope.RequestID, RunID: runID}) - flusher.Flush() + if writeSSE(w, protocol.StreamRecord{Version: protocol.Version, Type: "done", RequestID: envelope.RequestID, RunID: runID}) { + flusher.Flush() + successful = true + } else { + errorCode = "canceled" + } } } @@ -656,15 +711,32 @@ func (s *Server) newRunTempDir() (string, error) { } func (s *Server) centralConfig() *config.Config { - clone := *s.currentConfig() + current := s.currentConfig() + clone := *current clone.Gateway = config.GatewayConfig{} + clone.Providers = make(map[string]config.ProviderConfig, len(current.Providers)) + for key, providerConfig := range current.Providers { + if providerConfig.Env != nil { + env := make(map[string]string, len(providerConfig.Env)) + for name, value := range providerConfig.Env { + env[name] = value + } + providerConfig.Env = env + } + clone.Providers[key] = providerConfig + } return &clone } func (s *Server) currentConfig() *config.Config { + cfg, _ := s.currentConfigAndGeneration() + return cfg +} + +func (s *Server) currentConfigAndGeneration() (*config.Config, uint64) { s.configMu.RLock() defer s.configMu.RUnlock() - return s.config + return s.config, s.configGeneration } func (s *Server) decodeJSON(w http.ResponseWriter, r *http.Request, target any) bool { @@ -776,6 +848,7 @@ func (s *Server) refreshConfigIfStale() error { clone.Gateway = config.GatewayConfig{} s.configMu.Lock() s.config = &clone + s.configGeneration++ s.configFetchedAt = time.Now().UTC() s.configMu.Unlock() } diff --git a/internal/llm/chatgpt_test.go b/internal/llm/chatgpt_test.go index 5f8eaa75b..0553f4db7 100644 --- a/internal/llm/chatgpt_test.go +++ b/internal/llm/chatgpt_test.go @@ -525,7 +525,7 @@ func TestChatGPTHTTPSecondTurnMatchesReconstructedGatewayProvider(t *testing.T) } } -func TestChatGPTWebSocketDirectContinuationGatewayDifferenceAndFullHistoryFallback(t *testing.T) { +func TestChatGPTWebSocketWarmGatewayEquivalenceAndColdFullHistoryFallback(t *testing.T) { var mu sync.Mutex connections := 0 var captured []map[string]any @@ -585,8 +585,28 @@ func TestChatGPTWebSocketDirectContinuationGatewayDifferenceAndFullHistoryFallba return } - // Gateway-equivalent reconstruction has no connection-local response ID, - // so its second-turn transcript is complete on its first frame. + if connection == 2 { + // A warm gateway lease retains the same provider instance, so its first + // and continuation frames are transport-equivalent to direct mode. + _, data, err := conn.ReadMessage() + if err != nil { + t.Errorf("read warm gateway first request: %v", err) + return + } + capture(data) + _ = conn.WriteJSON(map[string]any{"type": "response.completed", "response": map[string]any{"id": "resp_direct_1"}}) + _, data, err = conn.ReadMessage() + if err != nil { + t.Errorf("read warm gateway continuation: %v", err) + return + } + capture(data) + _ = conn.WriteJSON(map[string]any{"type": "response.completed", "response": map[string]any{"id": "resp_warm_2"}}) + return + } + + // A cold gateway reconstruction has no connection-local response ID, so + // its second-turn transcript is complete on its first frame. _, data, err := conn.ReadMessage() if err != nil { t.Errorf("read reconstructed request: %v", err) @@ -622,8 +642,22 @@ func TestChatGPTWebSocketDirectContinuationGatewayDifferenceAndFullHistoryFallba drainStreamToDone(t, stream) _ = stream.Close() - reconstructed := newWSProvider() - stream, err = reconstructed.Stream(t.Context(), second) + warmGateway := newWSProvider() + stream, err = warmGateway.Stream(t.Context(), first) + if err != nil { + t.Fatal(err) + } + drainStreamToDone(t, stream) + _ = stream.Close() + stream, err = warmGateway.Stream(t.Context(), second) + if err != nil { + t.Fatal(err) + } + drainStreamToDone(t, stream) + _ = stream.Close() + + coldGateway := newWSProvider() + stream, err = coldGateway.Stream(t.Context(), second) if err != nil { t.Fatal(err) } @@ -633,8 +667,8 @@ func TestChatGPTWebSocketDirectContinuationGatewayDifferenceAndFullHistoryFallba mu.Lock() requests := append([]map[string]any(nil), captured...) mu.Unlock() - if len(requests) != 4 { - t.Fatalf("captured WebSocket requests = %d, want 4", len(requests)) + if len(requests) != 6 { + t.Fatalf("captured WebSocket requests = %d, want 6", len(requests)) } directContinuation := requests[1] if directContinuation["previous_response_id"] != "resp_direct_1" { @@ -643,7 +677,10 @@ func TestChatGPTWebSocketDirectContinuationGatewayDifferenceAndFullHistoryFallba if input, ok := directContinuation["input"].([]any); !ok || len(input) != 1 { t.Fatalf("direct continuation input = %#v, want only new suffix", directContinuation["input"]) } - for name, request := range map[string]map[string]any{"direct fallback": requests[2], "gateway reconstruction": requests[3]} { + if !reflect.DeepEqual(requests[0], requests[3]) || !reflect.DeepEqual(requests[1], requests[4]) { + t.Fatalf("direct and warm-gateway ChatGPT payloads differ\ndirect first: %#v\nwarm first: %#v\ndirect continuation: %#v\nwarm continuation: %#v", requests[0], requests[3], requests[1], requests[4]) + } + for name, request := range map[string]map[string]any{"direct fallback": requests[2], "cold gateway reconstruction": requests[5]} { if _, ok := request["previous_response_id"]; ok { t.Fatalf("%s retained previous_response_id: %#v", name, request) } @@ -651,7 +688,7 @@ func TestChatGPTWebSocketDirectContinuationGatewayDifferenceAndFullHistoryFallba t.Fatalf("%s input = %#v, want full transcript", name, request["input"]) } } - if !reflect.DeepEqual(requests[2], requests[3]) { - t.Fatalf("full-history WS fallback and gateway reconstruction differ\nfallback: %#v\ngateway: %#v", requests[2], requests[3]) + if !reflect.DeepEqual(requests[2], requests[5]) { + t.Fatalf("full-history WS fallback and cold gateway reconstruction differ\nfallback: %#v\ngateway: %#v", requests[2], requests[5]) } }