diff --git a/docs/formats.md b/docs/formats.md index 3575041..cfe4b5c 100644 --- a/docs/formats.md +++ b/docs/formats.md @@ -166,7 +166,7 @@ extra_params: | --- | --- | --- | | `id` | yes | Profile identifier, trimmed before selection and publication. It must be non-empty after trimming and unique within one source after normalization. | | `backend` | unless `endpoint` is present | Backend registry ID. It is trimmed and registry membership is checked when the profile is prepared or inspected. | -| `endpoint` | unless `backend` is present | Non-empty OpenAI-compatible base URL, including an API version path when required. When both connection fields are present, this overrides the backend endpoint without changing backend identity. | +| `endpoint` | unless `backend` is present | OpenAI-compatible base URL, including an API version path when required. A nonempty value is trimmed and must be absolute HTTP or HTTPS with a host and without user information, a query, or a fragment. When both connection fields are present, this overrides the backend endpoint without changing backend identity. | | `model` | yes | Non-empty provider model name. | | `temperature` | no | Number from 0 through 2. | | `max_tokens` | no | Integer zero or greater. | diff --git a/docs/integrations/openai-compatible-chat.md b/docs/integrations/openai-compatible-chat.md index 5714b8d..aa4ca75 100644 --- a/docs/integrations/openai-compatible-chat.md +++ b/docs/integrations/openai-compatible-chat.md @@ -14,10 +14,16 @@ that produce these outbound settings. Generation sends an HTTP `POST` with `Content-Type: application/json`. Before the client is called, the engine resolves framework, backend, profile, -and request values into one execution target. A non-empty endpoint from that -target overrides the client's configured base URL. After trailing slashes are -removed, `/chat/completions` is appended. Generation fails before sending when -neither source supplies an endpoint. +and request values into one execution target. Endpoint configuration is trimmed +and must be an absolute HTTP or HTTPS URL with a host and without user +information, a query, or a fragment. A non-empty endpoint from the target +overrides the client's configured base URL. The final selected endpoint is +validated again before transport. + +The completion URL is composed through parsed URL path operations. Nested base +paths are retained, repeated trailing slashes are normalized, and the result +has exactly one appended `/chat/completions` suffix. Generation fails before +sending when neither source supplies a valid endpoint. The target's backend ID is routing metadata for prepared values, results, and injected clients. The built-in client does not derive the URL from that ID and diff --git a/docs/internal/llm.md b/docs/internal/llm.md index 2f917f6..301f306 100644 --- a/docs/internal/llm.md +++ b/docs/internal/llm.md @@ -25,16 +25,19 @@ and request precedence. The client uses its endpoint, credential metadata, generation fields, and extra parameters. `BackendID` remains routing metadata for the generation boundary and is not mapped into the provider payload. -Construction validates the configured base URL and clones any supplied -`http.Client` so Promptkit can apply its timeout default without mutating the -caller's client. Generation then: +Construction trims and validates a nonempty configured base URL and clones any +supplied `http.Client` so Promptkit can apply its timeout default without +mutating the caller's client. An empty configured base remains valid because a +resolved request target may supply the endpoint. Generation then: -1. validates shared execution-setting invariants and endpoint requirements; +1. validates shared execution-setting invariants and the final selected base + endpoint; 2. maps the internal request into the OpenAI-compatible chat payload; 3. validates and merges extra parameters; -4. resolves authentication; -5. performs the outbound request under the applicable deadlines; and -6. decodes the first response choice and token usage. +4. composes `/chat/completions` through parsed URL path operations; +5. resolves authentication; +6. performs the outbound request under the applicable deadlines; and +7. decodes the first response choice and token usage. `internal/llm` owns the set of reserved OpenAI-compatible request fields used when validating extra parameters. Backend registration consumes the same rule @@ -65,6 +68,10 @@ configuration, invalid generation requests, request execution failures, non-success provider statuses, and malformed successful responses. Provider response bodies are not included in non-success errors. +Invalid nonempty configured endpoints are configuration failures. A missing or +invalid final selected endpoint is an invalid generation request and is +rejected before transport. + An `http.Client.Do` failure is represented by a redacting multi-cause error: the package request-failure sentinel and the exact returned transport error are both available through `errors.Is` and `errors.As`, while the rendered text @@ -80,7 +87,8 @@ The [OpenAI-compatible client tests](../../internal/llm/openai_compatible_client_test.go) own configuration, client cloning, deterministic deadline precedence, authentication, request and response mapping, malformed data, error identity, -cancellation, endpoint selection, and response-body suppression. The root +cancellation, endpoint selection and composition, pre-transport rejection, and +response-body suppression. The root transport contract tests also verify that resolved backend settings reach this client without serializing backend identity and that ordinary-run cancellation retains its public generation and context identities. All use local test diff --git a/docs/internal/overview.md b/docs/internal/overview.md index 4c93054..f9b1997 100644 --- a/docs/internal/overview.md +++ b/docs/internal/overview.md @@ -16,7 +16,7 @@ contributor workflow and validation. | `examples/go-library/run` | Demonstrates an offline downstream consumer using a prompt file, in-memory profile, inline input, an injected deterministic model client, and `Run`. It is not a public library package. | [Example program](../../examples/go-library/run/main.go) | | `internal/backend` | Constructs each engine's immutable registry from the built-in OpenRouter definition and consumer additions, validates and defensively copies definitions through the shared JSON-value package, and consumes the LLM-owned OpenAI-compatible reserved request-field rule. | [Backend registry](../../internal/backend/registry.go) | | `internal/capacity` | Owns engine-local bounded execution admission and FIFO model-generation permits for limited backend IDs, including cancellation-safe waiter removal and client wrapping. | [Internal capacity management](capacity.md) | -| `internal/domain` | Defines internal framework values for requests, artifacts, prompt definitions, profiles, execution targets, rendering, generation, and validation, and owns source-neutral invariants for shared execution settings, session identifiers, and output contracts. Source parsing, required fields, source-specific normalization and defaulting, and boundary-specific error classification remain with their callers. | [Domain declarations](../../internal/domain/domain.go) | +| `internal/domain` | Defines internal framework values for requests, artifacts, prompt definitions, profiles, execution targets, rendering, generation, and validation, and owns source-neutral invariants for shared execution settings, OpenAI-compatible base endpoints, session identifiers, and output contracts. Source parsing, required fields, other source-specific normalization, defaulting, and boundary-specific error classification remain with their callers. | [Domain declarations](../../internal/domain/domain.go), [endpoint invariant](../../internal/domain/endpoint.go) | | `internal/defaults` | Defines application-neutral framework constants and constructs the default execution target. It contains no CLI, server, or inbound HTTP limits. | [Framework defaults](../../internal/defaults/defaults.go) | | `internal/filecatalog` | Provides deterministic YAML discovery and path helpers for operating-system filesystems and `fs.FS` sources. | [File catalog](../../internal/filecatalog/catalog.go) | | `internal/jsonvalue` | Validates and deeply copies bounded JSON-compatible extra-parameter and prepared-schema trees while preserving supported concrete value types and rejecting cycles or excessive depth and work. | [JSON values](../../internal/jsonvalue/jsonvalue.go) | diff --git a/docs/policy/architecture.md b/docs/policy/architecture.md index 1f2f8f2..9774d6e 100644 --- a/docs/policy/architecture.md +++ b/docs/policy/architecture.md @@ -94,9 +94,10 @@ on consumers or on Scriptorium. `internal/domain` owns source-neutral invariants for values shared across multiple input and execution boundaries, including execution-setting bounds, -session identifiers, and output-contract legality. Callers retain source -parsing, required-field rules, source-specific normalization, defaulting, -error classification, and other policy specific to their own boundary. +OpenAI-compatible base endpoints, session identifiers, and output-contract +legality. Callers retain source parsing, required-field rules, other +source-specific normalization, defaulting, error classification, and policy +specific to their own boundary. ## Repository And Consumer Boundary diff --git a/engine_test.go b/engine_test.go index b2e878e..632ab79 100644 --- a/engine_test.go +++ b/engine_test.go @@ -1857,7 +1857,7 @@ func TestPrepareWorksWithInMemoryProfilesWithoutProfileFiles(t *testing.T) { SchemaDir: frameworkSchemaDir, }, promptkit.WithProfiles(promptkit.Profile{ ID: "memory-profile", - Endpoint: "http://memory-profile/v1", + Endpoint: " https://memory-profile/nested/v1 ", Model: "memory-model", })) if err != nil { @@ -1878,6 +1878,9 @@ func TestPrepareWorksWithInMemoryProfilesWithoutProfileFiles(t *testing.T) { if prepared.EffectiveModelParams.Model != "memory-model" { t.Fatalf("expected in-memory profile model, got %q", prepared.EffectiveModelParams.Model) } + if prepared.EffectiveModelParams.Endpoint != "https://memory-profile/nested/v1" { + t.Fatalf("expected normalized in-memory profile endpoint, got %q", prepared.EffectiveModelParams.Endpoint) + } } func TestInMemoryProfilesOverrideBuiltInsAndProfileSources(t *testing.T) { @@ -1961,6 +1964,79 @@ func TestWithProfilesRejectsInvalidExecutionSettings(t *testing.T) { } } +func TestEndpointValidationMapsPublicErrorCategories(t *testing.T) { + t.Run("backend registration is invalid configuration", func(t *testing.T) { + _, err := promptkit.NewEngine( + promptkit.Config{PromptDir: frameworkPromptDir}, + promptkit.WithBackend(promptkit.Backend{ID: "invalid", Endpoint: "/v1"}), + ) + if !errors.Is(err, promptkit.ErrInvalidConfig) { + t.Fatalf("expected ErrInvalidConfig, got %v", err) + } + }) + + t.Run("in-memory profile is invalid configuration", func(t *testing.T) { + _, err := promptkit.NewEngine( + promptkit.Config{PromptDir: frameworkPromptDir}, + promptkit.WithProfiles(promptkit.Profile{ + ID: "invalid", + Endpoint: "https://provider.example/v1?mode=chat", + Model: "model", + }), + ) + if !errors.Is(err, promptkit.ErrInvalidConfig) { + t.Fatalf("expected ErrInvalidConfig, got %v", err) + } + }) + + t.Run("file profile is a profile-load failure", func(t *testing.T) { + profileFS := fstest.MapFS{ + "profiles/invalid.yaml": &fstest.MapFile{Data: []byte(` +id: invalid +endpoint: 'https://provider.example/v1#chat' +model: model +`)}, + } + engine, err := promptkit.NewEngine( + promptkit.Config{PromptDir: frameworkPromptDir}, + promptkit.WithProfileFS(profileFS, "profiles"), + ) + if err != nil { + t.Fatalf("construct engine: %v", err) + } + _, err = engine.InspectProfile(context.Background(), "invalid") + if !errors.Is(err, promptkit.ErrProfileLoad) { + t.Fatalf("expected ErrProfileLoad, got %v", err) + } + }) + + t.Run("request override is an invalid request", func(t *testing.T) { + engine, err := promptkit.NewEngine( + promptkit.Config{PromptDir: frameworkPromptDir}, + promptkit.WithProfiles(promptkit.Profile{ + ID: "valid", + Endpoint: "https://provider.example/v1", + Model: "model", + }), + ) + if err != nil { + t.Fatalf("construct engine: %v", err) + } + _, err = engine.Prepare(context.Background(), promptkit.RunRequest{ + PromptID: frameworkMarkdownSummaryPromptID, + ProfileID: "valid", + Inputs: map[string]promptkit.ArtifactRef{ + "transcript": promptkit.Inline("Rin opens the gate."), + "glossary": promptkit.Inline("gate: A guarded passage."), + }, + Execution: &promptkit.ExecutionTargetOverride{Endpoint: "ftp://provider.example/v1"}, + }) + if !errors.Is(err, promptkit.ErrInvalidRequest) { + t.Fatalf("expected ErrInvalidRequest, got %v", err) + } + }) +} + func TestOpenAICompatibleProfileRunsThroughNormalProfilePath(t *testing.T) { fake := &fakeLLMClient{response: &promptkit.GenerateResponse{Content: "ok"}} prof := promptkit.OpenAICompatibleProfile(promptkit.OpenAICompatibleProfileConfig{ diff --git a/internal/backend/registry.go b/internal/backend/registry.go index 1382920..951a5b6 100644 --- a/internal/backend/registry.go +++ b/internal/backend/registry.go @@ -5,7 +5,6 @@ package backend import ( "errors" "fmt" - "net/url" "regexp" "sort" "strings" @@ -109,10 +108,11 @@ func (r *Registry) CapacityPolicies() map[string]domain.BackendCapacityPolicy { } func normalizeBackend(definition domain.Backend) (domain.Backend, error) { - definition.Endpoint = strings.TrimSpace(definition.Endpoint) - if err := validateEndpoint(definition.Endpoint); err != nil { + endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(definition.Endpoint) + if err != nil { return domain.Backend{}, fmt.Errorf("backend %q endpoint: %w", definition.ID, err) } + definition.Endpoint = endpoint definition.APIKeyEnv = strings.TrimSpace(definition.APIKeyEnv) if definition.APIKeyEnv != "" && !environmentVariableName.MatchString(definition.APIKeyEnv) { @@ -182,31 +182,3 @@ func normalizeBackend(definition domain.Backend) (domain.Backend, error) { definition.ExtraParams = extraParams return definition, nil } - -func validateEndpoint(endpoint string) error { - if endpoint == "" { - return errors.New("must not be blank") - } - if strings.Contains(endpoint, "#") { - return errors.New("must not contain a fragment") - } - - parsed, err := url.Parse(endpoint) - if err != nil { - return fmt.Errorf("must be a valid URL: %w", err) - } - scheme := strings.ToLower(parsed.Scheme) - if scheme != "http" && scheme != "https" { - return errors.New("must use http or https") - } - if !parsed.IsAbs() || parsed.Hostname() == "" { - return errors.New("must be absolute and include a host") - } - if parsed.User != nil { - return errors.New("must not contain user information") - } - if parsed.RawQuery != "" || parsed.ForceQuery { - return errors.New("must not contain a query string") - } - return nil -} diff --git a/internal/domain/endpoint.go b/internal/domain/endpoint.go new file mode 100644 index 0000000..26a86ce --- /dev/null +++ b/internal/domain/endpoint.go @@ -0,0 +1,39 @@ +package domain + +import ( + "errors" + "net/url" + "strings" +) + +// NormalizeOpenAICompatibleBaseEndpoint trims and validates a source-neutral +// OpenAI-compatible provider base endpoint. +func NormalizeOpenAICompatibleBaseEndpoint(endpoint string) (string, error) { + endpoint = strings.TrimSpace(endpoint) + if endpoint == "" { + return "", errors.New("endpoint must not be blank") + } + if strings.Contains(endpoint, "#") { + return "", errors.New("endpoint must not contain a fragment") + } + + parsed, err := url.Parse(endpoint) + if err != nil { + return "", errors.New("endpoint must be a valid URL") + } + parsed.Scheme = strings.ToLower(parsed.Scheme) + if parsed.Scheme != "http" && parsed.Scheme != "https" { + return "", errors.New("endpoint must use http or https") + } + if !parsed.IsAbs() || parsed.Hostname() == "" { + return "", errors.New("endpoint must be absolute and include a host") + } + if parsed.User != nil { + return "", errors.New("endpoint must not contain user information") + } + if parsed.RawQuery != "" || parsed.ForceQuery { + return "", errors.New("endpoint must not contain a query string") + } + + return parsed.String(), nil +} diff --git a/internal/domain/endpoint_test.go b/internal/domain/endpoint_test.go new file mode 100644 index 0000000..cb396e5 --- /dev/null +++ b/internal/domain/endpoint_test.go @@ -0,0 +1,47 @@ +package domain + +import "testing" + +func TestNormalizeOpenAICompatibleBaseEndpoint(t *testing.T) { + tests := []struct { + name string + endpoint string + want string + wantErr bool + }{ + {name: "http host", endpoint: "http://provider.example", want: "http://provider.example"}, + {name: "https nested path and whitespace", endpoint: " HTTPS://provider.example/api/openai/v1 ", want: "https://provider.example/api/openai/v1"}, + {name: "IPv4 host and port", endpoint: "http://127.0.0.1:8080/v1", want: "http://127.0.0.1:8080/v1"}, + {name: "IPv6 host and port", endpoint: "https://[::1]:8443/v1", want: "https://[::1]:8443/v1"}, + {name: "repeated trailing slashes", endpoint: "https://provider.example/v1///", want: "https://provider.example/v1///"}, + {name: "blank", endpoint: " \t\n ", wantErr: true}, + {name: "relative path", endpoint: "/api/v1", wantErr: true}, + {name: "scheme relative", endpoint: "//provider.example/v1", wantErr: true}, + {name: "missing host", endpoint: "https:///v1", wantErr: true}, + {name: "unsupported scheme", endpoint: "ftp://provider.example/v1", wantErr: true}, + {name: "user information", endpoint: "https://user:secret@provider.example/v1", wantErr: true}, + {name: "query", endpoint: "https://provider.example/v1?mode=chat", wantErr: true}, + {name: "empty query", endpoint: "https://provider.example/v1?", wantErr: true}, + {name: "fragment", endpoint: "https://provider.example/v1#chat", wantErr: true}, + {name: "empty fragment", endpoint: "https://provider.example/v1#", wantErr: true}, + {name: "malformed URL", endpoint: "https://provider.example/%zz", wantErr: true}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got, err := NormalizeOpenAICompatibleBaseEndpoint(tc.endpoint) + if tc.wantErr { + if err == nil { + t.Fatalf("expected endpoint error, got %q", got) + } + return + } + if err != nil { + t.Fatalf("normalize endpoint: %v", err) + } + if got != tc.want { + t.Fatalf("normalized endpoint = %q, want %q", got, tc.want) + } + }) + } +} diff --git a/internal/llm/openai_compatible_client.go b/internal/llm/openai_compatible_client.go index c666565..48ab417 100644 --- a/internal/llm/openai_compatible_client.go +++ b/internal/llm/openai_compatible_client.go @@ -51,9 +51,11 @@ type OpenAICompatibleClient struct { } func NewOpenAICompatibleClient(cfg OpenAICompatibleConfig) (*OpenAICompatibleClient, error) { - baseURL := strings.TrimSpace(cfg.BaseURL) - if baseURL != "" { - if _, err := url.ParseRequestURI(baseURL); err != nil { + baseURL := "" + if strings.TrimSpace(cfg.BaseURL) != "" { + var err error + baseURL, err = domain.NormalizeOpenAICompatibleBaseEndpoint(cfg.BaseURL) + if err != nil { return nil, fmt.Errorf("%w: invalid base URL: %v", ErrInvalidConfig, err) } } @@ -75,7 +77,7 @@ func NewOpenAICompatibleClient(cfg OpenAICompatibleConfig) (*OpenAICompatibleCli } return &OpenAICompatibleClient{ - baseURL: strings.TrimRight(baseURL, "/"), + baseURL: baseURL, defaultModel: cfg.Model, httpClient: client, }, nil @@ -86,14 +88,18 @@ func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.Genera return nil, fmt.Errorf("%w: %v", ErrInvalidRequest, err) } - endpoint := strings.TrimSpace(req.Target.Endpoint) - if endpoint == "" { - endpoint = c.baseURL + selectedEndpoint := req.Target.Endpoint + if strings.TrimSpace(selectedEndpoint) == "" { + selectedEndpoint = c.baseURL } - if endpoint == "" { - return nil, fmt.Errorf("%w: endpoint is required", ErrInvalidRequest) + endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(selectedEndpoint) + if err != nil { + return nil, fmt.Errorf("%w: invalid endpoint: %v", ErrInvalidRequest, err) + } + endpoint, err = url.JoinPath(endpoint, defaults.OpenAIChatCompletionsPath) + if err != nil { + return nil, fmt.Errorf("%w: invalid endpoint path: %v", ErrInvalidRequest, err) } - endpoint = strings.TrimRight(endpoint, "/") + defaults.OpenAIChatCompletionsPath wireReq, err := openAIChatRequestFromGenerateRequest(req, c.defaultModel) if err != nil { diff --git a/internal/llm/openai_compatible_client_test.go b/internal/llm/openai_compatible_client_test.go index 0a64eb5..25f9378 100644 --- a/internal/llm/openai_compatible_client_test.go +++ b/internal/llm/openai_compatible_client_test.go @@ -64,14 +64,21 @@ func assertDeadlineNear(t *testing.T, deadline, before, after time.Time, duratio } func TestNewOpenAICompatibleClientRejectsInvalidBaseURL(t *testing.T) { - _, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ - BaseURL: "://invalid", - }) - if err == nil { - t.Fatal("expected invalid configuration error") - } - if !errors.Is(err, ErrInvalidConfig) { - t.Fatalf("expected ErrInvalidConfig, got %v", err) + for _, endpoint := range []string{ + "://invalid", + "/v1", + "https:///v1", + "ftp://provider.example/v1", + "https://user@provider.example/v1", + "https://provider.example/v1?mode=chat", + "https://provider.example/v1#chat", + } { + t.Run(endpoint, func(t *testing.T) { + _, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: endpoint}) + if !errors.Is(err, ErrInvalidConfig) { + t.Fatalf("expected ErrInvalidConfig, got %v", err) + } + }) } } @@ -1312,6 +1319,94 @@ func TestOpenAICompatibleClientAllowsEmptyConfiguredBaseURL(t *testing.T) { } } +func TestOpenAICompatibleClientComposesCompletionURL(t *testing.T) { + tests := []struct { + name string + baseURL string + wantURL string + }{ + {name: "HTTP host", baseURL: "http://provider.example", wantURL: "http://provider.example/chat/completions"}, + {name: "HTTPS host", baseURL: "https://provider.example", wantURL: "https://provider.example/chat/completions"}, + {name: "nested path", baseURL: "https://provider.example/api/openai/v1", wantURL: "https://provider.example/api/openai/v1/chat/completions"}, + {name: "trailing slash", baseURL: "https://provider.example/v1/", wantURL: "https://provider.example/v1/chat/completions"}, + {name: "repeated trailing slashes", baseURL: " https://provider.example/api/v1/// ", wantURL: "https://provider.example/api/v1/chat/completions"}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + var selectedURL string + client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ + BaseURL: tc.baseURL, + Model: "m", + HTTPClient: &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + selectedURL = req.URL.String() + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"choices":[{"message":{"content":"ok"}}]}`)), + Request: req, + }, nil + })}, + }) + if err != nil { + t.Fatalf("construct client: %v", err) + } + + _, err = client.Generate(context.Background(), domain.GenerateRequest{ + Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, + }) + if err != nil { + t.Fatalf("generate: %v", err) + } + if selectedURL != tc.wantURL { + t.Fatalf("selected URL = %q, want %q", selectedURL, tc.wantURL) + } + }) + } +} + +func TestOpenAICompatibleClientRejectsInvalidSelectedEndpointBeforeTransport(t *testing.T) { + invalidEndpoints := []string{ + "/v1", + "https:///v1", + "ftp://provider.example/v1", + "https://user@provider.example/v1", + "https://provider.example/v1?mode=chat", + "https://provider.example/v1#chat", + "https://sensitive-endpoint.example/%zz", + } + for _, endpoint := range invalidEndpoints { + t.Run(endpoint, func(t *testing.T) { + transportCalls := 0 + client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ + BaseURL: "https://configured.example/v1", + Model: "m", + HTTPClient: &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + transportCalls++ + return nil, errors.New("transport must not be called") + })}, + }) + if err != nil { + t.Fatalf("construct client: %v", err) + } + + _, err = client.Generate(context.Background(), domain.GenerateRequest{ + Prompt: domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}}}, + Target: domain.ExecutionTarget{Endpoint: endpoint}, + }) + if !errors.Is(err, ErrInvalidRequest) { + t.Fatalf("expected ErrInvalidRequest, got %v", err) + } + if strings.Contains(err.Error(), endpoint) { + t.Fatalf("error exposed selected endpoint %q: %v", endpoint, err) + } + if transportCalls != 0 { + t.Fatalf("transport calls = %d, want 0", transportCalls) + } + }) + } +} + func TestOpenAICompatibleClientRequiresEndpointWhenUnsetEverywhere(t *testing.T) { client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{ BaseURL: "", diff --git a/internal/profile/filesystem_repository.go b/internal/profile/filesystem_repository.go index bee533a..e8761d7 100644 --- a/internal/profile/filesystem_repository.go +++ b/internal/profile/filesystem_repository.go @@ -132,7 +132,7 @@ func loadProfile(ctx context.Context, fsys fs.FS, root string, id string) (*doma if err != nil { return nil, fmt.Errorf("%w: %s: %v", ErrInvalidProfile, relPath, err) } - if err := validateProfile(prof); err != nil { + if err := normalizeAndValidateProfile(prof); err != nil { if errors.Is(err, ErrRawAPIKeyNotAllowed) { return nil, fmt.Errorf("%w: %s", err, relPath) } @@ -256,13 +256,21 @@ func requireYAMLStreamEnd(decoder *yaml.Decoder) error { return errors.New("profile file must contain exactly one YAML document") } -func validateProfile(p *domain.ExecutionProfile) error { +func normalizeAndValidateProfile(p *domain.ExecutionProfile) error { if strings.TrimSpace(p.ID) == "" { return errors.New("id is required") } - if strings.TrimSpace(p.BackendID) == "" && strings.TrimSpace(p.Endpoint) == "" { + p.Endpoint = strings.TrimSpace(p.Endpoint) + if strings.TrimSpace(p.BackendID) == "" && p.Endpoint == "" { return errors.New("backend or endpoint is required") } + if p.Endpoint != "" { + endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(p.Endpoint) + if err != nil { + return err + } + p.Endpoint = endpoint + } if strings.TrimSpace(p.Model) == "" { return errors.New("model is required") } diff --git a/internal/profile/repository_test.go b/internal/profile/repository_test.go index 8bfb9a9..5d289c3 100644 --- a/internal/profile/repository_test.go +++ b/internal/profile/repository_test.go @@ -63,7 +63,7 @@ func TestFilesystemRepository_GetProfile(t *testing.T) { wantErr bool }{ {name: "backend only", connection: "backend: ' openrouter '", wantBackend: "openrouter"}, - {name: "endpoint only", connection: "endpoint: http://localhost:8000/v1", wantEndpoint: "http://localhost:8000/v1"}, + {name: "endpoint only", connection: "endpoint: ' https://localhost:8000/nested/v1 '", wantEndpoint: "https://localhost:8000/nested/v1"}, {name: "both", connection: "backend: openrouter\nendpoint: http://localhost:8000/v1", wantBackend: "openrouter", wantEndpoint: "http://localhost:8000/v1"}, {name: "neither", wantErr: true}, {name: "blank backend", connection: "backend: ' '", wantErr: true}, @@ -351,6 +351,43 @@ model: second }) } +func TestProfileRepositoriesRejectInvalidEndpoints(t *testing.T) { + tests := []struct { + name string + endpoint string + withBackend bool + }{ + {name: "relative", endpoint: "/v1"}, + {name: "missing host", endpoint: "https:///v1"}, + {name: "unsupported scheme", endpoint: "ftp://provider.example/v1"}, + {name: "user information", endpoint: "https://user@provider.example/v1"}, + {name: "query", endpoint: "https://provider.example/v1?mode=chat"}, + {name: "fragment", endpoint: "https://provider.example/v1#chat"}, + {name: "backend with invalid override", endpoint: "/v1", withBackend: true}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + backend := "" + if tc.withBackend { + backend = "backend: openrouter\n" + } + repo := NewFSRepository(fstest.MapFS{ + "profiles/invalid.yaml": profileMapFile(fmt.Sprintf( + "id: invalid-endpoint\nmodel: model\n%sendpoint: %q\n", + backend, + tc.endpoint, + )), + }, "profiles") + + _, err := repo.GetProfile(context.Background(), "invalid-endpoint") + if !errors.Is(err, ErrInvalidProfile) { + t.Fatalf("expected ErrInvalidProfile, got %v", err) + } + }) + } +} + func TestProfileRepositoriesValidateExtraParams(t *testing.T) { const validProfile = ` id: selected-profile diff --git a/internal/usecase/execution_settings_test.go b/internal/usecase/execution_settings_test.go index a0bc7d7..912d046 100644 --- a/internal/usecase/execution_settings_test.go +++ b/internal/usecase/execution_settings_test.go @@ -31,3 +31,53 @@ func TestRunnerPrepareExecutionRejectsInvalidExecutionSettings(t *testing.T) { t.Fatalf("expected ErrInvalidRequest, got %v", err) } } + +func TestRunnerPrepareExecutionValidatesAndNormalizesRequestEndpoints(t *testing.T) { + newRunner := func() *Runner { + return NewRunner( + &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)}, + &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}}, + nil, + defaultArtifactReader(), + defaultRenderer(), + &fakeLLM{forbid: true}, + nil, + nil, + ) + } + + invalidEndpoints := []string{ + "/v1", + "https:///v1", + "ftp://provider.example/v1", + "https://user@provider.example/v1", + "https://provider.example/v1?mode=chat", + "https://provider.example/v1#chat", + } + for _, endpoint := range invalidEndpoints { + t.Run(endpoint, func(t *testing.T) { + _, err := newRunner().PrepareExecution(context.Background(), domain.RunRequest{ + PromptID: "p", + ProfileID: "exec", + Inputs: singleInputRef(), + Execution: &domain.ExecutionTargetOverride{Endpoint: endpoint}, + }) + if !errors.Is(err, ErrInvalidRequest) { + t.Fatalf("expected ErrInvalidRequest, got %v", err) + } + }) + } + + prepared, err := newRunner().PrepareExecution(context.Background(), domain.RunRequest{ + PromptID: "p", + ProfileID: "exec", + Inputs: singleInputRef(), + Execution: &domain.ExecutionTargetOverride{Endpoint: " https://provider.example/nested/v1 "}, + }) + if err != nil { + t.Fatalf("prepare normalized endpoint: %v", err) + } + if got := prepared.Details().EffectiveModelParams.Endpoint; got != "https://provider.example/nested/v1" { + t.Fatalf("effective endpoint = %q", got) + } +} diff --git a/internal/usecase/profile_inspection.go b/internal/usecase/profile_inspection.go index f3ddd86..3d84026 100644 --- a/internal/usecase/profile_inspection.go +++ b/internal/usecase/profile_inspection.go @@ -57,14 +57,19 @@ func (r *Runner) resolveProfileSelection( }, nil } -func validateResolvedExecutionTarget(target domain.ExecutionTarget) error { - if strings.TrimSpace(target.Endpoint) == "" { - return errors.New("execution endpoint is required") +func normalizeResolvedExecutionTarget(target domain.ExecutionTarget) (domain.ExecutionTarget, error) { + endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(target.Endpoint) + if err != nil { + return domain.ExecutionTarget{}, fmt.Errorf("execution endpoint: %w", err) } + target.Endpoint = endpoint if strings.TrimSpace(target.Model) == "" { - return errors.New("execution model is required") + return domain.ExecutionTarget{}, errors.New("execution model is required") } - return domain.ValidateExecutionTargetSettings(target) + if err := domain.ValidateExecutionTargetSettings(target); err != nil { + return domain.ExecutionTarget{}, err + } + return target, nil } // InspectProfile resolves one explicit profile without prompt or execution work. @@ -87,7 +92,8 @@ func (r *Runner) InspectProfile( return nil, err } target, _ := resolveExecutionTarget(selection.backend, selection.profile, nil) - if err := validateResolvedExecutionTarget(target); err != nil { + target, err = normalizeResolvedExecutionTarget(target) + if err != nil { return nil, fmt.Errorf("%w: %w", ErrProfileLoad, err) } target.APIKey = "" diff --git a/internal/usecase/runner.go b/internal/usecase/runner.go index c63d127..ff88113 100644 --- a/internal/usecase/runner.go +++ b/internal/usecase/runner.go @@ -317,7 +317,8 @@ func (r *Runner) resolvePreparation( effectiveModel, targetPresence := resolveExecutionTarget(selection.backend, selection.profile, req.Execution) effectiveModel.APIKey = req.APIKey - if err := validateResolvedExecutionTarget(effectiveModel); err != nil { + effectiveModel, err = normalizeResolvedExecutionTarget(effectiveModel) + if err != nil { return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err) } if err := validateAPIKey(effectiveModel.APIKeyEnv, effectiveModel.APIKey, effectiveModel.APIKeyRequired); err != nil { diff --git a/profiles.go b/profiles.go index 04b649c..1f7248e 100644 --- a/profiles.go +++ b/profiles.go @@ -100,19 +100,27 @@ func toDomainProfile(publicProfile Profile) (domain.ExecutionProfile, error) { APIKeyRequired: publicProfile.APIKeyRequired, ExtraParams: extraParams, } - if err := validatePublicProfile(prof); err != nil { + if err := normalizeAndValidatePublicProfile(&prof); err != nil { return domain.ExecutionProfile{}, err } return prof, nil } -func validatePublicProfile(prof domain.ExecutionProfile) error { +func normalizeAndValidatePublicProfile(prof *domain.ExecutionProfile) error { if strings.TrimSpace(prof.ID) == "" { return errors.New("id is required") } - if strings.TrimSpace(prof.BackendID) == "" && strings.TrimSpace(prof.Endpoint) == "" { + prof.Endpoint = strings.TrimSpace(prof.Endpoint) + if strings.TrimSpace(prof.BackendID) == "" && prof.Endpoint == "" { return errors.New("backend or endpoint is required") } + if prof.Endpoint != "" { + endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(prof.Endpoint) + if err != nil { + return err + } + prof.Endpoint = endpoint + } if strings.TrimSpace(prof.Model) == "" { return errors.New("model is required") } diff --git a/types.go b/types.go index f00722b..87f17e5 100644 --- a/types.go +++ b/types.go @@ -297,7 +297,8 @@ type ExecutionTarget struct { // empty for endpoint-only profiles. It is supplied to injected LLMClient // implementations as part of the effective target. BackendID string `json:"backend_id,omitempty"` - // Endpoint is the model-provider base URL. + // Endpoint is the normalized absolute HTTP or HTTPS model-provider base URL. + // It has a host and no user information, query, or fragment. Endpoint string `json:"endpoint"` // Model is the provider model identifier. Model string `json:"model"` @@ -396,7 +397,9 @@ type PromptInspection struct { // framework deadline when no higher-precedence value is present. type ExecutionTargetOverride struct { // Endpoint replaces the profile or backend endpoint when non-empty without - // changing the effective BackendID. + // changing the effective BackendID. Preparation trims it and requires an + // absolute HTTP or HTTPS URL with a host and no user information, query, or + // fragment. Endpoint string // Model replaces the profile model when non-empty. Model string @@ -457,7 +460,8 @@ type Profile struct { BackendID string // Endpoint is the model-provider base URL. It is required only when // BackendID is blank and otherwise overrides the backend endpoint when - // non-blank. + // non-blank. WithProfiles trims it and requires an absolute HTTP or HTTPS URL + // with a host and no user information, query, or fragment. Endpoint string // Model is the required non-blank provider model identifier. Model string