From 19997842a722d7175b24813e1fb7fe804d6aff4b Mon Sep 17 00:00:00 2001 From: NaDdjg Date: Thu, 10 Sep 2026 15:11:34 +0100 Subject: [PATCH 1/6] integrated edenai --- .env.template | 12 +- README.md | 1 + config/config.example.yaml | 15 + config/config.go | 2 + config/config_test.go | 2 +- docs/advanced/configuration.mdx | 3 +- docs/docs.json | 1 + docs/features/passthrough-api.mdx | 6 +- docs/providers/edenai.mdx | 164 ++++++ docs/providers/overview.mdx | 11 + internal/providers/config_test.go | 46 ++ internal/providers/edenai/edenai.go | 161 ++++++ internal/providers/edenai/edenai_test.go | 468 ++++++++++++++++++ internal/providers/edenai/models.go | 262 ++++++++++ internal/providers/edenai/models_test.go | 425 ++++++++++++++++ .../providers/edenai/passthrough_semantics.go | 14 + .../edenai/passthrough_semantics_test.go | 81 +++ internal/providers/edenai/response.go | 126 +++++ internal/providers/edenai/response_test.go | 225 +++++++++ .../registry_provider_pricing_test.go | 145 ++++++ internal/server/handlers_test.go | 2 +- internal/server/passthrough_support.go | 2 +- internal/server/passthrough_support_test.go | 43 ++ internal/usage/cost.go | 32 ++ internal/usage/cost_test.go | 97 ++++ internal/usage/extractor_test.go | 64 +++ internal/usage/stream_observer.go | 22 + internal/usage/stream_observer_test.go | 143 ++++++ run/lifecycle_test.go | 18 + run/providers.go | 2 + run/providers_test.go | 2 +- web/dashboard/src/lib/utils/providerDocs.js | 1 + web/dashboard/tests/provider-docs.test.js | 1 + 33 files changed, 2589 insertions(+), 10 deletions(-) create mode 100644 docs/providers/edenai.mdx create mode 100644 internal/providers/edenai/edenai.go create mode 100644 internal/providers/edenai/edenai_test.go create mode 100644 internal/providers/edenai/models.go create mode 100644 internal/providers/edenai/models_test.go create mode 100644 internal/providers/edenai/passthrough_semantics.go create mode 100644 internal/providers/edenai/passthrough_semantics_test.go create mode 100644 internal/providers/edenai/response.go create mode 100644 internal/providers/edenai/response_test.go create mode 100644 internal/providers/registry_provider_pricing_test.go diff --git a/.env.template b/.env.template index 5dda58db0..79f198563 100644 --- a/.env.template +++ b/.env.template @@ -82,9 +82,9 @@ # Allow optional /p/{provider}/v1/... passthrough aliases while keeping /p/{provider}/... canonical (default: true) # ALLOW_PASSTHROUGH_V1_ALIAS=true -# Comma-separated list of provider types enabled for /p/{provider}/... passthrough (default: openai,anthropic,openrouter,kilo,zai,sglang,vllm,llamacpp,llmd,deepseek,hetzner) +# Comma-separated list of provider types enabled for /p/{provider}/... passthrough (default: openai,anthropic,openrouter,kilo,zai,sglang,vllm,llamacpp,llmd,deepseek,hetzner,edenai) # Cohere native passthrough is opt-in; add cohere when those routes are needed. -# ENABLED_PASSTHROUGH_PROVIDERS=openai,anthropic,cohere,openrouter,kilo,zai,sglang,vllm,llamacpp,llmd,deepseek,hetzner +# ENABLED_PASSTHROUGH_PROVIDERS=openai,anthropic,cohere,openrouter,kilo,zai,sglang,vllm,llamacpp,llmd,deepseek,hetzner,edenai # Enable the realtime (speech-to-speech) endpoints (default: true): the /v1/realtime # websocket (and /p/{provider}/v1/realtime passthrough upgrade), the WebRTC SDP @@ -648,6 +648,14 @@ # Optional configured model list; see CONFIGURED_PROVIDER_MODELS_MODE below # KILO_MODELS=anthropic/claude-sonnet-4.5,openai/gpt-5.5 +# Eden AI (default base URL: https://api.edenai.run/v3) +# Multi-provider gateway. Model IDs use provider/model and pass through unchanged. +# EDENAI_BASE_URL is optional: the default above is used when it is unset. +# EDENAI_API_KEY= +# EDENAI_BASE_URL=https://api.edenai.run/v3 +# Optional configured model list; see CONFIGURED_PROVIDER_MODELS_MODE below +# EDENAI_MODELS=openai/gpt-4,anthropic/claude-sonnet-latest + # Z.ai (default base URL: https://api.z.ai/api/paas/v4) # For GLM Coding Plan, use: https://api.z.ai/api/coding/paas/v4 # ZAI_API_KEY=... diff --git a/README.md b/README.md index 73c27a459..6d2e865ee 100644 --- a/README.md +++ b/README.md @@ -120,6 +120,7 @@ The official SDKs therefore work unchanged. Configure their base URLs as follows - Z.ai - Alibaba Cloud Model Studio (Bailian) - Kilo AI +- Eden AI - MiniMax - Xiaomi MiMo - OpenCode Go diff --git a/config/config.example.yaml b/config/config.example.yaml index 28260e3f7..426a260bd 100644 --- a/config/config.example.yaml +++ b/config/config.example.yaml @@ -532,6 +532,21 @@ providers: # base_url defaults to "https://llm.chutes.ai/v1". # Set base_url when using a different compatible endpoint. + edenai: + type: edenai + api_key: "${EDENAI_API_KEY}" + # base_url defaults to "https://api.edenai.run/v3". + # Multi-provider gateway: model IDs use provider/model notation and pass + # through unchanged. Models, context windows, capabilities, and per-token + # pricing are discovered from Eden's own /v3/models, and per-request cost + # comes from the exact USD "cost" Eden returns on each response, so no + # pricing metadata needs declaring here. + # GoModel's /v1/responses is served by translating to Eden's chat + # completions; Eden's own /v3/responses is a different API and is never used. + # models: + # - id: "openai/gpt-4" + # - id: "anthropic/claude-sonnet-latest" + elevenlabs: type: elevenlabs api_key: "${ELEVENLABS_API_KEY}" diff --git a/config/config.go b/config/config.go index 5d0b4be5f..97ab53273 100644 --- a/config/config.go +++ b/config/config.go @@ -127,6 +127,8 @@ func buildDefaultConfig() *Config { "llamacpp", "llmd", "deepseek", + "hetzner", + "edenai", }, }, Models: ModelsConfig{ diff --git a/config/config_test.go b/config/config_test.go index 98d9238a3..470fed960 100644 --- a/config/config_test.go +++ b/config/config_test.go @@ -147,7 +147,7 @@ func TestBuildDefaultConfig(t *testing.T) { if !cfg.Server.AllowPassthroughV1Alias { t.Error("expected Server.AllowPassthroughV1Alias=true") } - if got, want := cfg.Server.EnabledPassthroughProviders, []string{"openai", "anthropic", "openrouter", "kilo", "zai", "sglang", "vllm", "llamacpp", "llmd", "deepseek"}; !reflect.DeepEqual(got, want) { + if got, want := cfg.Server.EnabledPassthroughProviders, []string{"openai", "anthropic", "openrouter", "kilo", "zai", "sglang", "vllm", "llamacpp", "llmd", "deepseek", "hetzner", "edenai"}; !reflect.DeepEqual(got, want) { t.Errorf("expected Server.EnabledPassthroughProviders=%v, got %v", want, got) } if cfg.Models.ConfiguredProviderModelsMode != ConfiguredProviderModelsModeFallback { diff --git a/docs/advanced/configuration.mdx b/docs/advanced/configuration.mdx index 1eb6c4ad0..2956e0b5d 100644 --- a/docs/advanced/configuration.mdx +++ b/docs/advanced/configuration.mdx @@ -332,7 +332,7 @@ Set these to automatically register providers. No YAML configuration required. | `LLAMACPP_BASE_URL` | llama.cpp llama-server / LM Studio (no API key needed unless started with `--api-key`) | | `LLMD_BASE_URL` | llm-d Router/EPP (no API key needed unless its Gateway requires one) | -Most providers can use a custom base URL via `_BASE_URL` (for example `OPENAI_BASE_URL`). Chutes AI defaults to `https://llm.chutes.ai/v1` and can be overridden with `CHUTES_BASE_URL`. DeepSeek defaults to `https://api.deepseek.com`; set `DEEPSEEK_BASE_URL` only for a compatible proxy or alternate DeepSeek endpoint. OpenRouter defaults to `https://openrouter.ai/api/v1` and can be overridden with `OPENROUTER_BASE_URL`. Kilo AI defaults to `https://api.kilo.ai/api/gateway` and can be overridden with `KILO_BASE_URL`. Z.ai defaults to `https://api.z.ai/api/paas/v4`; set `ZAI_BASE_URL=https://api.z.ai/api/coding/paas/v4` for the GLM Coding Plan endpoint. SGLang defaults to `http://localhost:30000/v1` when `SGLANG_API_KEY` is set, but keyless deployments should set `SGLANG_BASE_URL` explicitly to register the provider. vLLM follows the same pattern at `http://localhost:8000/v1`. llama.cpp's `LLAMACPP_BASE_URL` is always required (llama-server's default port collides with GoModel's own 8080, so there is no default); `LLAMACPP_API_KEY` is optional. llm-d has no universal endpoint, so `LLMD_BASE_URL` is always required; `LLMD_API_KEY` is optional. Azure uses `AZURE_BASE_URL` for its deployment base URL and accepts an optional `AZURE_API_VERSION` override; otherwise it defaults to `2024-10-21`. Oracle requires `ORACLE_BASE_URL` because its OpenAI-compatible endpoint is region-specific. +Most providers can use a custom base URL via `_BASE_URL` (for example `OPENAI_BASE_URL`). Chutes AI defaults to `https://llm.chutes.ai/v1` and can be overridden with `CHUTES_BASE_URL`. DeepSeek defaults to `https://api.deepseek.com`; set `DEEPSEEK_BASE_URL` only for a compatible proxy or alternate DeepSeek endpoint. OpenRouter defaults to `https://openrouter.ai/api/v1` and can be overridden with `OPENROUTER_BASE_URL`. Kilo AI defaults to `https://api.kilo.ai/api/gateway` and can be overridden with `KILO_BASE_URL`. Eden AI defaults to `https://api.edenai.run/v3` and can be overridden with `EDENAI_BASE_URL`. Z.ai defaults to `https://api.z.ai/api/paas/v4`; set `ZAI_BASE_URL=https://api.z.ai/api/coding/paas/v4` for the GLM Coding Plan endpoint. SGLang defaults to `http://localhost:30000/v1` when `SGLANG_API_KEY` is set, but keyless deployments should set `SGLANG_BASE_URL` explicitly to register the provider. vLLM follows the same pattern at `http://localhost:8000/v1`. llama.cpp's `LLAMACPP_BASE_URL` is always required (llama-server's default port collides with GoModel's own 8080, so there is no default); `LLAMACPP_API_KEY` is optional. llm-d has no universal endpoint, so `LLMD_BASE_URL` is always required; `LLMD_API_KEY` is optional. Azure uses `AZURE_BASE_URL` for its deployment base URL and accepts an optional `AZURE_API_VERSION` override; otherwise it defaults to `2024-10-21`. Oracle requires `ORACLE_BASE_URL` because its OpenAI-compatible endpoint is region-specific. Every provider type also accepts a comma-separated configured model list via `_MODELS`, for example `OPENROUTER_MODELS`, `ORACLE_MODELS`, @@ -477,6 +477,7 @@ export GROQ_API_KEY="gsk_..." # Registers "groq" provider export CHUTES_API_KEY="cpk_..." # Registers "chutes" provider export OPENROUTER_API_KEY="sk-or-..." # Registers "openrouter" provider export KILO_API_KEY="..." # Registers "kilo" provider +export EDENAI_API_KEY="..." # Registers "edenai" provider export ZAI_API_KEY="..." # Registers "zai" provider # Optional: export ZAI_BASE_URL="https://api.z.ai/api/coding/paas/v4" export AZURE_API_KEY="..." # Registers "azure" provider when paired with AZURE_BASE_URL diff --git a/docs/docs.json b/docs/docs.json index 9f3b14914..8e834069f 100644 --- a/docs/docs.json +++ b/docs/docs.json @@ -219,6 +219,7 @@ "providers/multiple-ollama", "providers/kimicode", "providers/hetzner", + "providers/edenai", { "group": "Cloud Platforms", "icon": "cloud", diff --git a/docs/features/passthrough-api.mdx b/docs/features/passthrough-api.mdx index 5a432887f..aa37e1807 100644 --- a/docs/features/passthrough-api.mdx +++ b/docs/features/passthrough-api.mdx @@ -136,8 +136,8 @@ from passthrough requests before forwarding them upstream. Passthrough is intentionally narrow while the API is in beta. -- `openai`, `anthropic`, `openrouter`, `kilo`, `zai`, `sglang`, `vllm`, `llamacpp`, `llmd`, and `deepseek` - are enabled by default. +- `openai`, `anthropic`, `openrouter`, `kilo`, `zai`, `sglang`, `vllm`, `llamacpp`, `llmd`, `deepseek`, + `hetzner`, and `edenai` are enabled by default. - Chutes supports passthrough but requires explicit operator opt-in because passthrough can forward provider-native routes that do not identify a model. Add `chutes` to `ENABLED_PASSTHROUGH_PROVIDERS` only when you intend to expose @@ -155,7 +155,7 @@ Passthrough routes are enabled by default: ```env ENABLE_PASSTHROUGH_ROUTES=true ALLOW_PASSTHROUGH_V1_ALIAS=true -ENABLED_PASSTHROUGH_PROVIDERS=openai,anthropic,openrouter,kilo,zai,sglang,vllm,llamacpp,llmd,deepseek +ENABLED_PASSTHROUGH_PROVIDERS=openai,anthropic,openrouter,kilo,zai,sglang,vllm,llamacpp,llmd,deepseek,hetzner,edenai ``` Set `ENABLED_PASSTHROUGH_PROVIDERS` to the provider types you want to expose. diff --git a/docs/providers/edenai.mdx b/docs/providers/edenai.mdx new file mode 100644 index 000000000..3594ef079 --- /dev/null +++ b/docs/providers/edenai.mdx @@ -0,0 +1,164 @@ +--- +title: "Eden AI" +description: "Configure Eden AI's OpenAI-compatible multi-provider API in GoModel." +icon: "leaf" +keywords: ["Eden AI", "edenai", "multi-provider", "OpenAI-compatible", "provider setup"] +--- + +Eden AI is a multi-provider gateway exposing an OpenAI-compatible REST API at +`https://api.edenai.run/v3`. GoModel routes chat completions, streaming, model +listing, embeddings, and passthrough through the shared OpenAI adapter, so one +Eden key reaches models from OpenAI, Anthropic, Google, Mistral, Cohere, +DeepInfra, and others. + +## Configure + +Create an API key in the [Eden AI console](https://app.edenai.run/) and set: + +```bash +EDENAI_API_KEY= +``` + +`EDENAI_BASE_URL` is optional — the provider defaults to +`https://api.edenai.run/v3`. Set it only to reach a different Eden-compatible +endpoint: + +```bash +EDENAI_BASE_URL=https://api.edenai.run/v3 +``` + +Or in `config.yaml`: + +```yaml +providers: + edenai: + type: edenai + api_key: "${EDENAI_API_KEY}" +``` + +You can also add the credential from the **Providers** page in the admin +dashboard instead of using env vars. + +## Models + +GoModel discovers Eden's catalog from Eden's own `GET /v3/models` on startup +and on every registry refresh. **There is no built-in model list**: a model +Eden adds is routable as soon as the catalog refreshes, with no GoModel +upgrade and no configuration change. + +Model IDs use `provider/model` notation and are forwarded unchanged: + +```json +{ "model": "openai/gpt-4", "messages": [{ "role": "user", "content": "Hi" }] } +``` + +Because the ID already contains a slash, qualify it with the provider name when +another configured provider exposes the same raw ID: +`edenai/anthropic/claude-sonnet-latest`. GoModel strips only the outer +`edenai/` routing qualifier before forwarding. + +Optionally pin a configured subset: + +```bash +EDENAI_MODELS=openai/gpt-4,anthropic/claude-sonnet-latest +``` + +### Discovered metadata + +Each catalog entry contributes metadata that `GET /v1/models` returns and that +the router, filters, and cost strategies use: + +| Eden field | GoModel metadata | +| ---------- | ---------------- | +| `context_length` | context window | +| `capabilities.supports_*` | capabilities (`reasoning`, `function_calling`, `prompt_caching`, …) | +| `capabilities.input_modalities` | `vision` / `audio` / `video` capabilities | +| `capabilities.output_modalities` | modes and categories | +| `pricing` | per-model pricing (see below) | + +Output modalities also decide what GoModel advertises: Eden's image- and +speech-only models are filtered out of `/v1/models`, because this provider +serves chat, embeddings, and passthrough only. + +## Pricing + +Eden publishes per-token USD rates per model, and GoModel converts them to its +per-million-token representation (`input_cost_per_token: 6e-8` → `$0.06 / +MTok`). Rates come from Eden's `pricing` block, which is what the account is +actually charged; the undiscounted `list_pricing` block is ignored. + +Pricing is read live from Eden — nothing is hard-coded, and Eden models do not +need to be present in GoModel's central model catalog. This makes +`EDENAI_MODEL_FILTER_MAX_PRICE_PER_MTOK`, cost-based load balancing, and +price display work for Eden models. + +A rate Eden does not publish is left unset rather than treated as free, so a +partially priced model is never silently undercounted. + +## Request cost + +Eden returns the exact USD charge for each request as a top-level `cost` +member, and GoModel records that figure as the request's cost instead of +recomputing it from token counts. Eden reprices its upstreams automatically +and applies account-level discounts, so its own number is authoritative in a +way a rate-card reconstruction is not. + +This makes Eden spend visible to usage records, budgets, cost dashboards, and +observability, and the usage entry is labelled with the cost source +`edenai_cost`. If a response carries no usable cost, GoModel falls back to the +discovered per-model pricing above. + + + The `provider` member Eden returns (`"openai"`, `"deepinfra"`) names the + upstream Eden routed to. GoModel reports `edenai` as the executing provider — + that is the provider it called — and re-exposes Eden's value on the response + as `edenai_upstream_provider` so clients can still see which upstream served + the request. + + +## Responses API + +GoModel serves `/v1/responses` for Eden by **translating the request to Eden's +chat-completions endpoint**. + + + Eden's own `/v3/responses` route is **not** the OpenAI Responses API. It takes + Eden-specific inputs (`routing`, `router_candidates`, `fallbacks`) and returns + its own response object, so GoModel never forwards to it. Treat `/v1/responses` + support here as chat-completion translation, not native compatibility. + + +## Eden-specific request fields + +Eden accepts extra top-level fields on chat completions — `routing`, +`fallbacks`, `session_id`, `pre_hooks`, and `post_hooks`. GoModel preserves +unknown top-level JSON fields on chat requests, so these reach Eden unchanged: + +```json +{ + "model": "openai/gpt-4", + "messages": [{ "role": "user", "content": "Hi" }], + "fallbacks": ["anthropic/claude-sonnet-latest"] +} +``` + +## Embeddings + +Eden's `/embeddings` route is OpenAI-compatible and uses the same +`provider/model` IDs: + +```json +{ "model": "openai/text-embedding-3-small", "input": "hello" } +``` + +## Passthrough + +`edenai` is in the default `ENABLED_PASSTHROUGH_PROVIDERS` allowlist, so +`/p/edenai/...` routes work without operator opt-in. Passthrough is a generic +forwarder: it sends any path you give it to Eden unchanged, under the +gateway's own credential. + +## Unsupported surfaces + +Files, batches, and audio are not exposed for Eden. Requests to those gateway +endpoints will not route to this provider. diff --git a/docs/providers/overview.mdx b/docs/providers/overview.mdx index 7395ff6d9..15a5278f6 100644 --- a/docs/providers/overview.mdx +++ b/docs/providers/overview.mdx @@ -53,6 +53,7 @@ support, not every individual model capability exposed by an upstream provider. | Meta (Muse Spark) | `META_API_KEY` (`META_BASE_URL` optional) | `muse-spark-1.1` | ✅ | ✅ | ❌ | ❌ | ❌ | ✅ | — | | OpenRouter | `OPENROUTER_API_KEY` | `google/gemini-2.5-flash` | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | — | | Kilo AI | `KILO_API_KEY` (`KILO_BASE_URL` optional) | `anthropic/claude-sonnet-4.5` | ✅ | ✅ | ❌ | ❌ | ❌ | ✅ | — | +| Eden AI | `EDENAI_API_KEY` (`EDENAI_BASE_URL` optional) | `openai/gpt-4` | ✅ | ✅ | ✅ | ❌ | ❌ | ✅ | [Eden AI](/providers/edenai) | | Z.ai | `ZAI_API_KEY` (`ZAI_BASE_URL` optional) | `glm-5.1` | ✅ | ✅ | ✅ | ❌ | ❌ | ✅ | — | | xAI (Grok) | `XAI_API_KEY` | `grok-4.6` | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | [xAI (Grok)](/providers/xai) | | Alibaba Cloud Model Studio (Bailian) | `BAILIAN_API_KEY` (`BAILIAN_BASE_URL` optional) | `qwen3-max` | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | [Alibaba Cloud Model Studio](/providers/bailian) | @@ -118,6 +119,16 @@ support, not every individual model capability exposed by an upstream provider. another configured provider exposes the same raw model ID, select Kilo explicitly with `kilo/anthropic/claude-sonnet-4.5`; GoModel removes only the outer `kilo/` routing qualifier before forwarding. +- **Eden AI** — model IDs use `provider/model` (for example, `openai/gpt-4`) + and are forwarded unchanged; select Eden explicitly with + `edenai/openai/gpt-4` when another configured provider exposes the same raw + ID. GoModel discovers Eden's catalog, context windows, capabilities, and + per-token pricing from Eden's own `/v3/models`, so there is no built-in + model list and new Eden models need no upgrade. Per-request cost comes from + the exact USD `cost` Eden returns on each response rather than from token + math. GoModel serves `/v1/responses` by translating it to Eden's + chat-completions endpoint — Eden's own `/v3/responses` is a different API + and is never used. See the [Eden AI guide](/providers/edenai). - **Xiaomi MiMo** — TTS (`mimo-v2.5-tts*`) and ASR (`mimo-v2.5-asr`) are served through `/v1/audio/speech` and `/v1/audio/transcriptions` (translated to MiMo's chat-completions audio dialect) as well as directly via chat diff --git a/internal/providers/config_test.go b/internal/providers/config_test.go index aa46f8ee0..a20cce260 100644 --- a/internal/providers/config_test.go +++ b/internal/providers/config_test.go @@ -91,6 +91,9 @@ var testDiscoveryConfigs = map[string]DiscoveryConfig{ "hetzner": { DefaultBaseURL: "https://inference.hetzner.com/api/v1", }, + "edenai": { + DefaultBaseURL: "https://api.edenai.run/v3", + }, } // --- buildProviderConfig --- @@ -1960,6 +1963,49 @@ func TestBuildProviderConfig_Hetzner_ResolvesBaseURL(t *testing.T) { } } +// TestBuildProviderConfig_EdenAI_ResolvesBaseURL asserts that EDENAI_API_KEY +// alone registers the provider and resolves Eden's default endpoint. The env +// prefix is derived from the registered type "edenai" by the generic +// discovery, so this also pins the spelling: renaming the type to "eden-ai" +// would silently move the credential to EDEN_AI_API_KEY. +func TestBuildProviderConfig_EdenAI_ResolvesBaseURL(t *testing.T) { + t.Setenv("EDENAI_API_KEY", "edenai-test-key") + + got := applyProviderEnvVars(map[string]config.RawProviderConfig{}, testDiscoveryConfigs) + + p, exists := got["edenai"] + if !exists { + t.Fatal("edenai not discovered by config parser") + } + if p.Type != "edenai" { + t.Errorf("Type = %q, want edenai", p.Type) + } + if p.APIKey != "edenai-test-key" { + t.Errorf("APIKey = %q, want edenai-test-key", p.APIKey) + } + if p.BaseURL != "https://api.edenai.run/v3" { + t.Errorf("BaseURL = %q, want https://api.edenai.run/v3", p.BaseURL) + } +} + +// TestBuildProviderConfig_EdenAI_BaseURLOverride asserts EDENAI_BASE_URL wins +// over the registered default, so operators can point the provider at a +// different Eden-compatible endpoint. +func TestBuildProviderConfig_EdenAI_BaseURLOverride(t *testing.T) { + t.Setenv("EDENAI_API_KEY", "edenai-test-key") + t.Setenv("EDENAI_BASE_URL", "https://eden.internal.example/v3") + + got := applyProviderEnvVars(map[string]config.RawProviderConfig{}, testDiscoveryConfigs) + + p, exists := got["edenai"] + if !exists { + t.Fatal("edenai not discovered by config parser") + } + if p.BaseURL != "https://eden.internal.example/v3" { + t.Errorf("BaseURL = %q, want https://eden.internal.example/v3", p.BaseURL) + } +} + func TestApplyProviderEnvVars_ModelFilter(t *testing.T) { t.Setenv("OPENROUTER_API_KEY", "sk-openrouter") t.Setenv("OPENROUTER_MODEL_FILTER_INCLUDE", "*:free, *:nitro") diff --git a/internal/providers/edenai/edenai.go b/internal/providers/edenai/edenai.go new file mode 100644 index 000000000..610422e48 --- /dev/null +++ b/internal/providers/edenai/edenai.go @@ -0,0 +1,161 @@ +// Package edenai provides Eden AI API integration for the LLM gateway. +// +// Eden AI is a multi-provider gateway exposing an OpenAI-compatible API at +// https://api.edenai.run/v3. Chat completions, streaming, model listing, +// embeddings, and passthrough all go through the shared OpenAI-compatible +// transport, and model IDs use provider/model notation +// ("openai/gpt-4", "deepinfra/inclusionAI/Ling-3.0-flash-VL") which is +// forwarded unchanged. +// +// Eden also serves a /responses route, but it is not the OpenAI Responses +// API: it takes Eden-specific inputs (routing, router_candidates, fallbacks) +// and returns its own response object. Forwarding a GoModel Responses request +// there would hand the client a body it cannot parse, so Responses and +// StreamResponses are translated through chat completions instead and Eden's +// /responses route is never reached. +// +// The provider composes an unexported *openai.CompatibleProvider rather than +// embedding *openai.ChatCompatible. Composition is required for two reasons: +// ListModels needs the raw transport (CompatibleProvider.Do) to decode Eden's +// catalog metadata, and explicit delegation keeps the surface to exactly what +// Eden implements — Go embedding cannot subtract the batch, file, and audio +// methods the router discovers by interface assertion. +package edenai + +import ( + "context" + "io" + "net/http" + + "github.com/enterpilot/gomodel/internal/core" + "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/openai" +) + +const ( + defaultBaseURL = "https://api.edenai.run/v3" + providerType = "edenai" +) + +// Registration provides factory registration for the Eden AI provider. +var Registration = providers.Registration{ + Type: providerType, + New: New, + PassthroughSemanticEnricher: passthroughSemanticEnricher, + Discovery: providers.DiscoveryConfig{ + DefaultBaseURL: defaultBaseURL, + }, +} + +// Provider implements the core.Provider interface for Eden AI. Eden +// authenticates with a plain bearer token and exposes an OpenAI-shaped chat, +// models, and embeddings surface. Eden-only request fields such as routing +// and fallbacks reach the upstream unchanged through core.ChatRequest's +// unknown-field passthrough, so no request adaptation is needed. +type Provider struct { + compat *openai.CompatibleProvider +} + +var ( + _ core.Provider = (*Provider)(nil) + _ core.PassthroughProvider = (*Provider)(nil) +) + +// New creates a new Eden AI provider. +func New(cfg providers.ProviderConfig, opts providers.ProviderOptions) core.Provider { + return &Provider{compat: openai.NewCompatibleProvider(cfg.APIKey, opts, compatibleConfig( + providers.ResolveBaseURL(cfg.BaseURL, defaultBaseURL), + ))} +} + +// NewWithHTTPClient creates a new Eden AI provider with a custom HTTP client. +// If httpClient is nil, http.DefaultClient is used. +// +// The signature matches every other chat-compatible provider on main: +// (apiKey, baseURL, httpClient, hooks). +func NewWithHTTPClient(apiKey string, baseURL string, httpClient *http.Client, hooks llmclient.Hooks) *Provider { + return &Provider{compat: openai.NewCompatibleProviderWithHTTPClient(apiKey, httpClient, hooks, compatibleConfig( + providers.ResolveBaseURL(baseURL, defaultBaseURL), + ))} +} + +// compatibleConfig returns the shared OpenAI-compatible transport settings for Eden AI. +func compatibleConfig(baseURL string) openai.CompatibleProviderConfig { + return openai.CompatibleProviderConfig{ + ProviderName: providerType, + BaseURL: baseURL, + SetHeaders: setHeaders, + } +} + +// setHeaders applies Eden AI's bearer-token authentication. CompatibleProvider +// sends no credential when SetHeaders is nil (unlike ChatCompatible, which +// defaults to bearer), so this must stay wired up. +func setHeaders(req *http.Request, apiKey string) { + providers.SetAuthHeaders(req, apiKey, providers.AuthHeaderConfig{AuthScheme: "Bearer "}) +} + +// SetBaseURL changes the Eden AI API base URL. +func (p *Provider) SetBaseURL(baseURL string) { + p.compat.SetBaseURL(baseURL) +} + +// GetBaseURL returns the provider's current base URL. +func (p *Provider) GetBaseURL() string { + return p.compat.GetBaseURL() +} + +// ChatCompletion sends a chat completion request to Eden AI and normalizes +// the Eden-specific response members (see normalizeChatResponse). +func (p *Provider) ChatCompletion(ctx context.Context, req *core.ChatRequest) (*core.ChatResponse, error) { + resp, err := p.compat.ChatCompletion(ctx, req) + if err != nil { + return nil, err + } + normalizeChatResponse(resp) + return resp, nil +} + +// StreamChatCompletion sends a streaming chat completion request to Eden AI. +// +// The stream is forwarded verbatim, so unlike ChatCompletion there is no +// opportunity to relocate Eden's root-level cost before the usage pipeline +// sees it. The shared stream observer harvests a root-level "cost" from the +// chunk carrying usage (see usage.copyRootLevelCost), which covers Eden +// reporting cost the same way it does on the non-streaming response. Whether +// Eden actually emits cost on streamed chunks is not established by its +// published contract; when it does not, usage falls back to the per-model +// pricing discovered from /models. +func (p *Provider) StreamChatCompletion(ctx context.Context, req *core.ChatRequest) (io.ReadCloser, error) { + return p.compat.StreamChatCompletion(ctx, req) +} + +// Responses translates an OpenAI Responses request through Eden chat +// completions. Eden's native /responses route is a different API and is +// deliberately never called. +func (p *Provider) Responses(ctx context.Context, req *core.ResponsesRequest) (*core.ResponsesResponse, error) { + return providers.ResponsesViaChat(ctx, p, req) +} + +// StreamResponses translates a streaming Responses request through Eden chat +// completions, for the same reason as Responses. +func (p *Provider) StreamResponses(ctx context.Context, req *core.ResponsesRequest) (io.ReadCloser, error) { + return providers.StreamResponsesViaChat(ctx, p, req, providerType) +} + +// Embeddings sends an embeddings request to Eden AI. Eden's /embeddings route +// is OpenAI-compatible and takes the same provider/model IDs. +func (p *Provider) Embeddings(ctx context.Context, req *core.EmbeddingRequest) (*core.EmbeddingResponse, error) { + resp, err := p.compat.Embeddings(ctx, req) + if err != nil { + return nil, err + } + normalizeEmbeddingResponse(resp) + return resp, nil +} + +// Passthrough forwards an opaque request to Eden AI. +func (p *Provider) Passthrough(ctx context.Context, req *core.PassthroughRequest) (*core.PassthroughResponse, error) { + return p.compat.Passthrough(ctx, req) +} diff --git a/internal/providers/edenai/edenai_test.go b/internal/providers/edenai/edenai_test.go new file mode 100644 index 000000000..052737c66 --- /dev/null +++ b/internal/providers/edenai/edenai_test.go @@ -0,0 +1,468 @@ +package edenai + +import ( + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/enterpilot/gomodel/internal/core" + "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers" +) + +// slashedModel is an Eden AI model ID in provider/model notation. Every +// request test uses it so a regression that splits, rewrites, or filters the +// provider prefix fails immediately. +const slashedModel = "openai/gpt-4" + +// TestNew_ReturnsProvider asserts that New returns a non-nil *Provider whose +// composed CompatibleProvider is wired up. +func TestNew_ReturnsProvider(t *testing.T) { + provider := New(providers.ProviderConfig{APIKey: "test-api-key"}, providers.ProviderOptions{}) + + if provider == nil { + t.Fatal("provider should not be nil") + } + + concrete, ok := provider.(*Provider) + if !ok { + t.Fatalf("New() returned %T, want *edenai.Provider", provider) + } + if concrete.compat == nil { + t.Error("composed CompatibleProvider should not be nil") + } +} + +// TestNew_DefaultsBaseURL asserts that a config without an explicit base URL +// falls back to Eden's public endpoint rather than an empty target. +func TestNew_DefaultsBaseURL(t *testing.T) { + provider, ok := New(providers.ProviderConfig{APIKey: "test-api-key"}, providers.ProviderOptions{}).(*Provider) + if !ok { + t.Fatal("New() did not return *edenai.Provider") + } + if got := provider.GetBaseURL(); got != defaultBaseURL { + t.Errorf("GetBaseURL() = %q, want %q", got, defaultBaseURL) + } +} + +// TestNew_HonoursConfiguredBaseURL asserts that EDENAI_BASE_URL (surfaced here +// as ProviderConfig.BaseURL) overrides the default. +func TestNew_HonoursConfiguredBaseURL(t *testing.T) { + const custom = "https://eden.internal.example/v3" + provider, ok := New(providers.ProviderConfig{APIKey: "k", BaseURL: custom}, providers.ProviderOptions{}).(*Provider) + if !ok { + t.Fatal("New() did not return *edenai.Provider") + } + if got := provider.GetBaseURL(); got != custom { + t.Errorf("GetBaseURL() = %q, want %q", got, custom) + } +} + +// TestNewWithHTTPClient_ReturnsProvider asserts the explicit HTTP-client constructor +// returns a valid Provider. +func TestNewWithHTTPClient_ReturnsProvider(t *testing.T) { + provider := NewWithHTTPClient("test-api-key", "http://example.invalid", &http.Client{}, llmclient.Hooks{}) + + if provider == nil { + t.Fatal("provider should not be nil") + } + if provider.compat == nil { + t.Error("composed CompatibleProvider should not be nil") + } +} + +// TestNewWithHTTPClient_NilHTTPClientDoesNotPanic asserts that passing nil for the +// HTTP client falls back to http.DefaultClient without panicking. +func TestNewWithHTTPClient_NilHTTPClientDoesNotPanic(t *testing.T) { + defer func() { + if r := recover(); r != nil { + t.Fatalf("NewWithHTTPClient(nil, ...) panicked: %v", r) + } + }() + provider := NewWithHTTPClient("test-api-key", "http://example.invalid", nil, llmclient.Hooks{}) + if provider == nil { + t.Fatal("provider should not be nil") + } +} + +// TestNewWithHTTPClient_ZeroHooksDoesNotPanic asserts that the hooks argument can be +// an empty struct (no hooks registered) without panicking. +func TestNewWithHTTPClient_ZeroHooksDoesNotPanic(t *testing.T) { + defer func() { + if r := recover(); r != nil { + t.Fatalf("NewWithHTTPClient(..., llmclient.Hooks{}) panicked: %v", r) + } + }() + provider := NewWithHTTPClient("test-api-key", "http://example.invalid", &http.Client{}, llmclient.Hooks{}) + if provider == nil { + t.Fatal("provider should not be nil") + } +} + +// TestRegistration_TypeAndDiscovery asserts the Registration struct exposes the +// expected type, New function, and default base URL. The type spelling also +// fixes the env prefix the generic discovery derives (EDENAI_API_KEY, +// EDENAI_BASE_URL) and the provider gate in internal/usage, so it is asserted +// exactly. +func TestRegistration_TypeAndDiscovery(t *testing.T) { + if Registration.Type != "edenai" { + t.Errorf("Registration.Type = %q, want %q", Registration.Type, "edenai") + } + if Registration.New == nil { + t.Error("Registration.New should not be nil") + } + want := "https://api.edenai.run/v3" + if Registration.Discovery.DefaultBaseURL != want { + t.Errorf("Registration.Discovery.DefaultBaseURL = %q, want %q", Registration.Discovery.DefaultBaseURL, want) + } + if Registration.PassthroughSemanticEnricher == nil { + t.Error("Registration.PassthroughSemanticEnricher should not be nil") + } + if Registration.Discovery.RequireBaseURL { + t.Error("Discovery.RequireBaseURL should be false: Eden has a public default endpoint") + } + if Registration.Discovery.AllowAPIKeyless { + t.Error("Discovery.AllowAPIKeyless should be false: Eden always requires an API key") + } +} + +// TestProvider_ImplementsCoreProvider is a compile-time check that *Provider +// satisfies the core.Provider interface used by the factory. +func TestProvider_ImplementsCoreProvider(t *testing.T) { + var _ core.Provider = (*Provider)(nil) +} + +// TestChatCompletion_UsesBearerAuthAndForwardsModel asserts that ChatCompletion +// posts to /chat/completions with the Bearer header and forwards Eden's +// provider/model ID unchanged. The bearer assertion matters more than usual +// here: CompatibleProvider sends no credential at all when SetHeaders is nil. +func TestChatCompletion_UsesBearerAuthAndForwardsModel(t *testing.T) { + var gotPath string + var gotAuth string + var gotBody map[string]any + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotPath = r.URL.Path + gotAuth = r.Header.Get("Authorization") + if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { + http.Error(w, "decode error", http.StatusBadRequest) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{ + "id":"chatcmpl-edenai", + "created":1677652288, + "model":"openai/gpt-4", + "choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}], + "usage":{"prompt_tokens":5,"completion_tokens":1,"total_tokens":6} + }`)) + })) + defer server.Close() + + provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ + Model: slashedModel, + Messages: []core.Message{{Role: "user", Content: "hi"}}, + }) + if err != nil { + t.Fatalf("ChatCompletion() error = %v", err) + } + if gotPath != "/chat/completions" { + t.Fatalf("path = %q, want /chat/completions", gotPath) + } + if gotAuth != "Bearer edenai-key" { + t.Fatalf("authorization = %q, want Bearer edenai-key", gotAuth) + } + if gotBody["model"] != slashedModel { + t.Fatalf("request model = %#v, want %q (provider/model IDs must pass through unchanged)", gotBody["model"], slashedModel) + } + if resp.Model != slashedModel { + t.Fatalf("response model = %q, want %q", resp.Model, slashedModel) + } + if len(resp.Choices) != 1 || resp.Choices[0].Message.Content != "hello" { + t.Fatalf("unexpected response: %+v", resp) + } +} + +// TestChatCompletion_ForwardsEdenExtraFields asserts that Eden-only request +// fields (routing, fallbacks) survive the round trip through the generic +// unknown-field mechanism, so no Eden-specific request adapter is needed. +func TestChatCompletion_ForwardsEdenExtraFields(t *testing.T) { + var gotBody map[string]any + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { + http.Error(w, "decode error", http.StatusBadRequest) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"id":"chatcmpl-edenai","created":1677652288,"model":"openai/gpt-4",` + + `"choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}]}`)) + })) + defer server.Close() + + var req core.ChatRequest + raw := `{"model":"openai/gpt-4","messages":[{"role":"user","content":"hi"}],` + + `"fallbacks":["anthropic/claude-sonnet-latest"],"routing":{"strategy":"cost"}}` + if err := json.Unmarshal([]byte(raw), &req); err != nil { + t.Fatalf("Unmarshal() error = %v", err) + } + + provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + if _, err := provider.ChatCompletion(context.Background(), &req); err != nil { + t.Fatalf("ChatCompletion() error = %v", err) + } + + fallbacks, ok := gotBody["fallbacks"].([]any) + if !ok || len(fallbacks) != 1 || fallbacks[0] != "anthropic/claude-sonnet-latest" { + t.Fatalf("fallbacks = %#v, want Eden fallback list forwarded unchanged", gotBody["fallbacks"]) + } + routing, ok := gotBody["routing"].(map[string]any) + if !ok || routing["strategy"] != "cost" { + t.Fatalf("routing = %#v, want Eden routing object forwarded unchanged", gotBody["routing"]) + } +} + +// TestStreamChatCompletion_UsesSSE asserts that streaming requests go to +// /chat/completions with the Bearer header, set stream=true, and return SSE data +// the adapter normalizes. +func TestStreamChatCompletion_UsesSSE(t *testing.T) { + var gotPath string + var gotAuth string + var gotBody map[string]any + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotPath = r.URL.Path + gotAuth = r.Header.Get("Authorization") + if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { + http.Error(w, "decode error", http.StatusBadRequest) + return + } + w.Header().Set("Content-Type", "text/event-stream") + _, _ = io.WriteString(w, "data: {\"id\":\"chatcmpl-edenai\",\"object\":\"chat.completion.chunk\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"hi\"}}]}\n\ndata: [DONE]\n\n") + })) + defer server.Close() + + provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + stream, err := provider.StreamChatCompletion(context.Background(), &core.ChatRequest{ + Model: slashedModel, + Messages: []core.Message{{Role: "user", Content: "hi"}}, + }) + if err != nil { + t.Fatalf("StreamChatCompletion() error = %v", err) + } + defer stream.Close() + body, err := io.ReadAll(stream) + if err != nil { + t.Fatalf("ReadAll() error = %v", err) + } + if gotPath != "/chat/completions" { + t.Fatalf("path = %q, want /chat/completions", gotPath) + } + if gotAuth != "Bearer edenai-key" { + t.Fatalf("authorization = %q, want Bearer edenai-key", gotAuth) + } + if gotBody["model"] != slashedModel || gotBody["stream"] != true { + t.Fatalf("stream request body = %#v", gotBody) + } + if !strings.Contains(string(body), "data: [DONE]") { + t.Fatalf("stream body = %q, want SSE terminator", body) + } +} + +// TestEmbeddings_ForwardsToEmbeddingsEndpoint asserts that embeddings reach +// Eden's OpenAI-compatible /embeddings route. Unlike hetzner and kilo, Eden +// documents this endpoint, so it is served rather than rejected locally. +func TestEmbeddings_ForwardsToEmbeddingsEndpoint(t *testing.T) { + var gotPath string + var gotAuth string + var gotBody map[string]any + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotPath = r.URL.Path + gotAuth = r.Header.Get("Authorization") + if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { + http.Error(w, "decode error", http.StatusBadRequest) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{ + "object":"list", + "data":[{"object":"embedding","embedding":[0.1,0.2],"index":0}], + "model":"openai/text-embedding-3-small", + "usage":{"prompt_tokens":4,"total_tokens":4} + }`)) + })) + defer server.Close() + + provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + resp, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ + Model: "openai/text-embedding-3-small", + Input: "hello", + }) + if err != nil { + t.Fatalf("Embeddings() error = %v", err) + } + if gotPath != "/embeddings" { + t.Fatalf("path = %q, want /embeddings", gotPath) + } + if gotAuth != "Bearer edenai-key" { + t.Fatalf("authorization = %q, want Bearer edenai-key", gotAuth) + } + if gotBody["model"] != "openai/text-embedding-3-small" { + t.Fatalf("request model = %#v, want provider/model ID forwarded unchanged", gotBody["model"]) + } + if len(resp.Data) != 1 || resp.Data[0].Index != 0 || len(resp.Data[0].Embedding) == 0 { + t.Fatalf("embedding data = %+v, want one populated vector", resp.Data) + } + if resp.Model != "openai/text-embedding-3-small" { + t.Errorf("response model = %q, want openai/text-embedding-3-small", resp.Model) + } + if resp.Usage.PromptTokens != 4 || resp.Usage.TotalTokens != 4 { + t.Errorf("usage = %+v, want prompt 4 / total 4", resp.Usage) + } +} + +// TestResponses_TranslatesToChatCompletions is the load-bearing test for this +// provider's architecture: Eden's own /responses route is not the OpenAI +// Responses API, so a GoModel Responses request must be translated through +// chat completions and must never reach /responses upstream. +func TestResponses_TranslatesToChatCompletions(t *testing.T) { + var paths []string + var gotBody struct { + Model string `json:"model"` + } + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + paths = append(paths, r.URL.Path) + if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil { + http.Error(w, "decode error", http.StatusBadRequest) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{ + "id":"chatcmpl-edenai", + "created":1677652288, + "model":"openai/gpt-4", + "choices":[{"index":0,"message":{"role":"assistant","content":"translated"},"finish_reason":"stop"}], + "usage":{"prompt_tokens":3,"completion_tokens":2,"total_tokens":5} + }`)) + })) + defer server.Close() + + provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{ + Model: slashedModel, + Input: "hi", + }) + if err != nil { + t.Fatalf("Responses() error = %v", err) + } + if len(paths) != 1 || paths[0] != "/chat/completions" { + t.Fatalf("upstream paths = %v, want exactly [/chat/completions]", paths) + } + for _, path := range paths { + if strings.Contains(path, "/responses") { + t.Fatalf("request reached %q; Eden's native /responses is not the OpenAI Responses API and must never be used", path) + } + } + if gotBody.Model != slashedModel { + t.Fatalf("request model = %q, want %q", gotBody.Model, slashedModel) + } + if resp.Object != "response" || resp.Status != "completed" { + t.Fatalf("response metadata = object %q status %q, want response/completed", resp.Object, resp.Status) + } +} + +// TestStreamResponses_TranslatesToChatCompletions asserts the streaming +// Responses surface is translated the same way, and likewise never touches +// Eden's native /responses route. +func TestStreamResponses_TranslatesToChatCompletions(t *testing.T) { + var paths []string + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + paths = append(paths, r.URL.Path) + w.Header().Set("Content-Type", "text/event-stream") + _, _ = io.WriteString(w, "data: {\"id\":\"chatcmpl-edenai\",\"object\":\"chat.completion.chunk\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"hi\"}}]}\n\ndata: [DONE]\n\n") + })) + defer server.Close() + + provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + stream, err := provider.StreamResponses(context.Background(), &core.ResponsesRequest{ + Model: slashedModel, + Input: "hi", + }) + if err != nil { + t.Fatalf("StreamResponses() error = %v", err) + } + defer stream.Close() + if _, err := io.ReadAll(stream); err != nil { + t.Fatalf("ReadAll() error = %v", err) + } + if len(paths) != 1 || paths[0] != "/chat/completions" { + t.Fatalf("upstream paths = %v, want exactly [/chat/completions]", paths) + } +} + +// TestPassthrough_ForwardsOpaqueRequest asserts the provider forwards an opaque +// passthrough request to the given Eden path under the gateway's own credential. +func TestPassthrough_ForwardsOpaqueRequest(t *testing.T) { + var gotPath string + var gotAuth string + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotPath = r.URL.Path + gotAuth = r.Header.Get("Authorization") + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"ok":true}`)) + })) + defer server.Close() + + provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ + Method: http.MethodPost, + Endpoint: "chat/completions", + Body: io.NopCloser(strings.NewReader(`{"model":"openai/gpt-4"}`)), + }) + if err != nil { + t.Fatalf("Passthrough() error = %v", err) + } + defer resp.Body.Close() + if gotPath != "/chat/completions" { + t.Fatalf("path = %q, want /chat/completions", gotPath) + } + if gotAuth != "Bearer edenai-key" { + t.Fatalf("authorization = %q, want Bearer edenai-key", gotAuth) + } + if resp.StatusCode != http.StatusOK { + t.Fatalf("status = %d, want 200", resp.StatusCode) + } +} + +// TestProvider_DoesNotExposeOptionalOpenAICompatibleInterfaces guards the +// composition: *Provider delegates only the methods Eden implements, so it +// must not satisfy the optional native interfaces. Eden documents no +// OpenAI-shaped batches API, and its files and audio surfaces are unverified, +// so advertising those capabilities to the router would promise what the +// upstream cannot honour. This matters more under composition than it did +// under embedding: adding a delegation by mistake is all it would take. +func TestProvider_DoesNotExposeOptionalOpenAICompatibleInterfaces(t *testing.T) { + provider := NewWithHTTPClient("edenai-key", "", nil, llmclient.Hooks{}) + + if _, ok := any(provider).(core.NativeBatchProvider); ok { + t.Fatal("edenai provider should not implement native batch provider") + } + if _, ok := any(provider).(core.NativeFileProvider); ok { + t.Fatal("edenai provider should not implement native file provider") + } + if _, ok := any(provider).(core.AudioProvider); ok { + t.Fatal("edenai provider should not implement audio provider") + } + if _, ok := any(provider).(core.ImageProvider); ok { + t.Fatal("edenai provider should not implement image provider") + } +} diff --git a/internal/providers/edenai/models.go b/internal/providers/edenai/models.go new file mode 100644 index 000000000..d2a556812 --- /dev/null +++ b/internal/providers/edenai/models.go @@ -0,0 +1,262 @@ +package edenai + +import ( + "context" + "math" + "net/http" + "strings" + + "github.com/enterpilot/gomodel/internal/core" + "github.com/enterpilot/gomodel/internal/llmclient" +) + +// usdPerTokenToPerMtok scales Eden's per-token USD rates to the gateway's +// per-million-token representation. Eden names the unit in the field itself +// (input_cost_per_token), so the conversion is fixed rather than inferred. +const usdPerTokenToPerMtok = 1_000_000 + +type modelsResponse struct { + Object string `json:"object"` + Data []modelInfo `json:"data"` +} + +// modelInfo is one entry of Eden's /models catalog. Capabilities is decoded +// as a loose map on purpose: Eden publishes a growing set of supports_* flags, +// and a map keeps new ones flowing through without a source change here. +type modelInfo struct { + ID string `json:"id"` + Object string `json:"object"` + OwnedBy string `json:"owned_by"` + Description string `json:"description"` + Capabilities map[string]any `json:"capabilities"` + Pricing *modelPricing `json:"pricing"` + ListPricing *modelPricing `json:"list_pricing"` + Created int64 `json:"created"` + ContextLength int `json:"context_length"` +} + +// modelPricing holds one of Eden's per-token USD rate blocks. +type modelPricing struct { + InputCostPerToken *float64 `json:"input_cost_per_token"` + OutputCostPerToken *float64 `json:"output_cost_per_token"` + CacheReadInputTokenCost *float64 `json:"cache_read_input_token_cost"` +} + +// effectivePricing picks the rate card to publish. Eden's `pricing` is what +// the account is actually charged (any discount already applied), so it wins +// whenever it carries a usable rate. `list_pricing` is the undiscounted card +// and is used only as a fallback: an approximate rate still lets price +// filters and the cost load-balancing strategy rank the model, whereas no +// rate at all drops it from both. +func (m modelInfo) effectivePricing() *core.ModelPricing { + if pricing := m.Pricing.toCore(); pricing != nil { + return pricing + } + return m.ListPricing.toCore() +} + +// ListModels returns Eden's live catalog, retaining the context window, +// capability flags, modalities, and per-token pricing Eden publishes with it. +// Eden's /models response is the source of truth for this provider: no static +// list is kept here, so a model Eden adds appears at the next registry +// refresh without a source change. +func (p *Provider) ListModels(ctx context.Context) (*core.ModelsResponse, error) { + var upstream modelsResponse + if err := p.compat.Do(ctx, llmclient.Request{ + Method: http.MethodGet, + Endpoint: "/models", + }, &upstream); err != nil { + return nil, err + } + + result := &core.ModelsResponse{Object: "list"} + result.Data = make([]core.Model, 0, len(upstream.Data)) + for _, model := range upstream.Data { + if strings.TrimSpace(model.ID) == "" { + continue + } + result.Data = append(result.Data, model.toCore()) + } + return result, nil +} + +// toCore normalizes an Eden catalog entry into GoModel's provider-neutral +// model shape. +func (m modelInfo) toCore() core.Model { + object := strings.TrimSpace(m.Object) + if object == "" { + object = "model" + } + + metadata := &core.ModelMetadata{ + Description: strings.TrimSpace(m.Description), + Capabilities: m.capabilities(), + Pricing: m.effectivePricing(), + } + if modes := m.modes(); len(modes) > 0 { + metadata.Modes = modes + metadata.Categories = core.CategoriesForModes(modes) + } + if m.ContextLength > 0 { + metadata.ContextWindow = new(m.ContextLength) + } + if metadataEmpty(metadata) { + metadata = nil + } + + return core.Model{ + ID: strings.TrimSpace(m.ID), + Object: object, + OwnedBy: strings.TrimSpace(m.OwnedBy), + Created: m.Created, + Metadata: metadata, + } +} + +// metadataEmpty reports whether an entry contributed nothing worth attaching, +// so a bare catalog row leaves Metadata nil rather than an empty struct that +// enrichment would treat as a real provider report. +func metadataEmpty(metadata *core.ModelMetadata) bool { + return metadata.Description == "" && + len(metadata.Capabilities) == 0 && + metadata.Pricing == nil && + len(metadata.Modes) == 0 && + metadata.ContextWindow == nil +} + +// modes maps Eden's output modalities onto the gateway's mode vocabulary. +// Text models claim both "chat" and "responses": the Responses surface is +// served for them by translating through chat completions. Modalities the +// gateway has no Eden-backed surface for still produce their mode, so the +// registry can hide models this provider cannot actually serve (it implements +// neither core.AudioProvider nor core.ImageProvider). +func (m modelInfo) modes() []string { + modes := make([]string, 0, 2) + seen := make(map[string]struct{}, 2) + add := func(mode string) { + if _, ok := seen[mode]; ok { + return + } + seen[mode] = struct{}{} + modes = append(modes, mode) + } + for _, modality := range stringSlice(m.Capabilities["output_modalities"]) { + switch strings.ToLower(strings.TrimSpace(modality)) { + case "text": + add("chat") + add("responses") + case "image": + add("image_generation") + case "audio", "speech": + add("audio_speech") + case "embedding", "embeddings": + add("embedding") + case "video": + add("video_generation") + } + } + if len(modes) == 0 { + return nil + } + return modes +} + +// capabilities projects Eden's supports_* flags and input modalities into the +// gateway's capability map. The supports_ prefix is stripped so the names read +// the same way other providers report them ("function_calling", "reasoning"). +// Unknown keys are ignored rather than guessed at, but any future supports_* +// flag is picked up automatically. +func (m modelInfo) capabilities() map[string]bool { + capabilities := make(map[string]bool, len(m.Capabilities)) + for key, value := range m.Capabilities { + name, ok := strings.CutPrefix(strings.ToLower(strings.TrimSpace(key)), "supports_") + if !ok || name == "" { + continue + } + enabled, ok := value.(bool) + if !ok || !enabled { + continue + } + capabilities[name] = true + } + for _, modality := range stringSlice(m.Capabilities["input_modalities"]) { + switch strings.ToLower(strings.TrimSpace(modality)) { + case "image": + capabilities["vision"] = true + case "audio": + capabilities["audio"] = true + case "video": + capabilities["video"] = true + } + } + if len(capabilities) == 0 { + return nil + } + return capabilities +} + +// stringSlice reads a JSON string array out of the loosely decoded +// capabilities map, skipping non-string members. +func stringSlice(value any) []string { + raw, ok := value.([]any) + if !ok { + return nil + } + result := make([]string, 0, len(raw)) + for _, item := range raw { + if text, ok := item.(string); ok { + result = append(result, text) + } + } + return result +} + +// toCore converts one of Eden's per-token USD rate blocks into the gateway's +// per-million-token pricing. Eden names the unit in each field +// (input_cost_per_token), so the ×1e6 scaling is read off the contract rather +// than assumed. A rate Eden omits stays absent: costing a token type at zero +// because no price was published would understate spend, so only a rate Eden +// explicitly reports as 0 prices at zero. +func (p *modelPricing) toCore() *core.ModelPricing { + if p == nil { + return nil + } + input, hasInput := perMtok(p.InputCostPerToken) + output, hasOutput := perMtok(p.OutputCostPerToken) + cachedInput, hasCachedInput := perMtok(p.CacheReadInputTokenCost) + if !hasInput && !hasOutput && !hasCachedInput { + return nil + } + + pricing := &core.ModelPricing{Currency: "USD"} + if hasInput { + pricing.InputPerMtok = &input + } + if hasOutput { + pricing.OutputPerMtok = &output + } + if hasCachedInput { + pricing.CachedInputPerMtok = &cachedInput + } + return pricing +} + +// perMtok scales one per-token USD rate to per million tokens. Rates that are +// absent, negative, or non-finite report no price rather than a wrong one, and +// so does a rate large enough that scaling overflows to infinity: a corrupt +// number here would propagate into every price comparison, budget total, and +// cost-strategy decision downstream. +func perMtok(perToken *float64) (float64, bool) { + if perToken == nil { + return 0, false + } + rate := *perToken + if rate < 0 || math.IsNaN(rate) || math.IsInf(rate, 0) { + return 0, false + } + scaled := rate * usdPerTokenToPerMtok + if math.IsInf(scaled, 0) { + return 0, false + } + return scaled, true +} diff --git a/internal/providers/edenai/models_test.go b/internal/providers/edenai/models_test.go new file mode 100644 index 000000000..10c958739 --- /dev/null +++ b/internal/providers/edenai/models_test.go @@ -0,0 +1,425 @@ +package edenai + +import ( + "context" + "math" + "net/http" + "net/http/httptest" + "slices" + "testing" + + "github.com/enterpilot/gomodel/internal/core" + "github.com/enterpilot/gomodel/internal/llmclient" +) + +// edenCatalogEntry is a verbatim entry from a live Eden /v3/models response. +// Keeping the real shape (including the null members and both pricing blocks) +// is what makes the mapping assertions below meaningful. +const edenCatalogEntry = `{ + "id": "deepinfra/inclusionAI/Ling-3.0-flash-VL", + "object": "model", + "created": 1788880306, + "owned_by": "deepinfra", + "model_name": "inclusionAI/Ling-3.0-flash-VL", + "context_length": 131072, + "description": null, + "source": null, + "capabilities": { + "input_modalities": ["text", "image"], + "output_modalities": ["text"], + "supports_reasoning": true, + "supports_web_search": false, + "supports_tool_choice": false, + "supports_computer_use": false, + "supports_prompt_caching": true, + "supports_response_schema": false, + "supports_system_messages": false, + "supports_function_calling": false, + "supports_native_streaming": false, + "supports_assistant_prefill": false, + "supports_embedding_image_input": false, + "supports_parallel_function_calling": false + }, + "pricing": { + "input_cost_per_token": 6e-8, + "output_cost_per_token": 1.8e-7, + "cache_read_input_token_cost": 1.2e-8 + }, + "list_pricing": { + "input_cost_per_token": 9e-8, + "output_cost_per_token": 2.8e-7, + "cache_read_input_token_cost": 2.2e-8 + }, + "discount": null, + "regions": [{"code": "us", "name": "United States"}], + "alias_of": null +}` + +// modelsServer serves one /models payload and records the request. +func modelsServer(t *testing.T, payload string) (*Provider, *string, *string) { + t.Helper() + gotPath := new("") + gotAuth := new("") + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + *gotPath = r.URL.Path + *gotAuth = r.Header.Get("Authorization") + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(payload)) + })) + t.Cleanup(server.Close) + return NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}), gotPath, gotAuth +} + +func firstModel(t *testing.T, payload string) core.Model { + t.Helper() + provider, _, _ := modelsServer(t, payload) + resp, err := provider.ListModels(context.Background()) + if err != nil { + t.Fatalf("ListModels() error = %v", err) + } + if len(resp.Data) != 1 { + t.Fatalf("models = %+v, want exactly one entry", resp.Data) + } + return resp.Data[0] +} + +// TestListModels_MapsLiveCatalogEntry asserts the full mapping of a real Eden +// catalog entry: identity, context window, capabilities, modalities, derived +// categories, and per-token pricing scaled to per-million-token. +func TestListModels_MapsLiveCatalogEntry(t *testing.T) { + provider, gotPath, gotAuth := modelsServer(t, `{"object":"list","data":[`+edenCatalogEntry+`]}`) + + resp, err := provider.ListModels(context.Background()) + if err != nil { + t.Fatalf("ListModels() error = %v", err) + } + if *gotPath != "/models" { + t.Fatalf("path = %q, want /models", *gotPath) + } + if *gotAuth != "Bearer edenai-key" { + t.Fatalf("authorization = %q, want Bearer edenai-key", *gotAuth) + } + if resp.Object != "list" || len(resp.Data) != 1 { + t.Fatalf("response = %+v, want a one-entry list", resp) + } + + model := resp.Data[0] + if model.ID != "deepinfra/inclusionAI/Ling-3.0-flash-VL" { + t.Errorf("ID = %q, want the Eden ID forwarded unchanged", model.ID) + } + if model.Object != "model" || model.OwnedBy != "deepinfra" || model.Created != 1788880306 { + t.Errorf("identity = %+v, want object/owned_by/created preserved", model) + } + + meta := model.Metadata + if meta == nil { + t.Fatal("Metadata = nil, want Eden catalog metadata") + } + if meta.ContextWindow == nil || *meta.ContextWindow != 131072 { + t.Errorf("ContextWindow = %v, want 131072", meta.ContextWindow) + } + + // output_modalities ["text"] -> chat + responses (Responses is served by + // translating through chat completions). + if want := []string{"chat", "responses"}; !slices.Equal(meta.Modes, want) { + t.Errorf("Modes = %v, want %v", meta.Modes, want) + } + if len(meta.Categories) != 1 || meta.Categories[0] != core.CategoryTextGeneration { + t.Errorf("Categories = %v, want [%v]", meta.Categories, core.CategoryTextGeneration) + } + + // Only the true supports_* flags become capabilities, with the prefix + // stripped; input_modalities ["text","image"] adds vision. + wantCapabilities := map[string]bool{"reasoning": true, "prompt_caching": true, "vision": true} + if len(meta.Capabilities) != len(wantCapabilities) { + t.Errorf("Capabilities = %v, want %v", meta.Capabilities, wantCapabilities) + } + for name := range wantCapabilities { + if !meta.Capabilities[name] { + t.Errorf("Capabilities[%q] = false, want true", name) + } + } + if meta.Capabilities["web_search"] || meta.Capabilities["function_calling"] { + t.Errorf("Capabilities = %v, want false flags omitted", meta.Capabilities) + } + + // 6e-8 USD/token -> $0.06/MTok, 1.8e-7 -> $0.18, 1.2e-8 -> $0.012. + // The values come from `pricing`, not the higher `list_pricing` block. + assertPrice(t, "InputPerMtok", meta.Pricing.InputPerMtok, 0.06) + assertPrice(t, "OutputPerMtok", meta.Pricing.OutputPerMtok, 0.18) + assertPrice(t, "CachedInputPerMtok", meta.Pricing.CachedInputPerMtok, 0.012) + if meta.Pricing.Currency != "USD" { + t.Errorf("Currency = %q, want USD", meta.Pricing.Currency) + } +} + +// TestListModels_UsesDiscountedPricingNotListPricing pins the choice of block: +// cost metadata must describe what the account is actually charged. +func TestListModels_UsesDiscountedPricingNotListPricing(t *testing.T) { + model := firstModel(t, `{"object":"list","data":[`+edenCatalogEntry+`]}`) + assertPrice(t, "InputPerMtok", model.Metadata.Pricing.InputPerMtok, 0.06) + if got := *model.Metadata.Pricing.InputPerMtok; got == 0.09 { + t.Fatal("InputPerMtok took the list_pricing rate; want the discounted pricing block") + } +} + +// TestListModels_FallsBackToListPricing asserts the undiscounted rate card is +// used when Eden publishes no applicable pricing. An approximate rate still +// lets price filters and the cost strategy rank the model; no rate at all +// drops it from both. +func TestListModels_FallsBackToListPricing(t *testing.T) { + model := firstModel(t, `{"object":"list","data":[{ + "id":"openai/gpt-4","object":"model", + "pricing": null, + "list_pricing": {"input_cost_per_token": 9e-8, "output_cost_per_token": 2.8e-7} + }]}`) + + if model.Metadata == nil || model.Metadata.Pricing == nil { + t.Fatal("Pricing = nil, want the list_pricing fallback") + } + assertPrice(t, "InputPerMtok", model.Metadata.Pricing.InputPerMtok, 0.09) + assertPrice(t, "OutputPerMtok", model.Metadata.Pricing.OutputPerMtok, 0.28) +} + +// TestListModels_ListPricingDoesNotMaskApplicablePricing asserts the fallback +// never overrides a usable applicable rate, including a partial one: a model +// priced only for input keeps that rate rather than swapping in the full list +// card. +func TestListModels_ListPricingDoesNotMaskApplicablePricing(t *testing.T) { + model := firstModel(t, `{"object":"list","data":[{ + "id":"openai/gpt-4","object":"model", + "pricing": {"input_cost_per_token": 6e-8}, + "list_pricing": {"input_cost_per_token": 9e-8, "output_cost_per_token": 2.8e-7} + }]}`) + + assertPrice(t, "InputPerMtok", model.Metadata.Pricing.InputPerMtok, 0.06) + if model.Metadata.Pricing.OutputPerMtok != nil { + t.Errorf("OutputPerMtok = %v, want nil: the list card must not fill gaps in applicable pricing", + *model.Metadata.Pricing.OutputPerMtok) + } +} + +// TestListModels_PricingEdgeCases covers partial, zero, negative, and +// overflowing rates. A rate Eden omits must stay absent rather than being +// costed at zero; a rate Eden reports as zero must be honoured as free. +func TestListModels_PricingEdgeCases(t *testing.T) { + tests := []struct { + name string + pricing string + wantNil bool + wantInput *float64 + wantOutput *float64 + wantCached *float64 + }{ + { + name: "no pricing member", + pricing: `null`, + wantNil: true, + }, + { + name: "empty pricing object", + pricing: `{}`, + wantNil: true, + }, + { + name: "partial pricing keeps the reported rate and omits the rest", + pricing: `{"input_cost_per_token": 6e-8}`, + wantInput: new(0.06), + }, + { + name: "explicit zero prices at zero", + pricing: `{"input_cost_per_token": 0, "output_cost_per_token": 0}`, + wantInput: new(0.0), + wantOutput: new(0.0), + }, + { + name: "negative rates are rejected, valid siblings survive", + pricing: `{"input_cost_per_token": -1e-8, "output_cost_per_token": 1.8e-7}`, + wantOutput: new(0.18), + }, + { + name: "every rate invalid yields no pricing", + pricing: `{"input_cost_per_token": -1e-8, "output_cost_per_token": -2e-8}`, + wantNil: true, + }, + { + name: "a rate that overflows on scaling is dropped", + pricing: `{"input_cost_per_token": 1e308, "output_cost_per_token": 1.8e-7}`, + wantOutput: new(0.18), + }, + { + name: "cached input rate alone is still pricing", + pricing: `{"cache_read_input_token_cost": 1.2e-8}`, + wantCached: new(0.012), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + payload := `{"object":"list","data":[{"id":"openai/gpt-4","object":"model","pricing":` + tt.pricing + `}]}` + model := firstModel(t, payload) + + if tt.wantNil { + if model.Metadata != nil && model.Metadata.Pricing != nil { + t.Fatalf("Pricing = %+v, want nil", model.Metadata.Pricing) + } + return + } + if model.Metadata == nil || model.Metadata.Pricing == nil { + t.Fatal("Pricing = nil, want a partial pricing block") + } + assertOptionalPrice(t, "InputPerMtok", model.Metadata.Pricing.InputPerMtok, tt.wantInput) + assertOptionalPrice(t, "OutputPerMtok", model.Metadata.Pricing.OutputPerMtok, tt.wantOutput) + assertOptionalPrice(t, "CachedInputPerMtok", model.Metadata.Pricing.CachedInputPerMtok, tt.wantCached) + }) + } +} + +// TestPerMtok_RejectsNonFiniteRates covers NaN and infinity directly: neither +// can be expressed as a JSON number, so they can only arrive through a decoder +// that tolerates them, and both must be refused rather than scaled. +func TestPerMtok_RejectsNonFiniteRates(t *testing.T) { + for _, rate := range []float64{math.NaN(), math.Inf(1), math.Inf(-1), -1} { + if _, ok := perMtok(&rate); ok { + t.Errorf("perMtok(%v) reported a usable price, want rejected", rate) + } + } + if _, ok := perMtok(nil); ok { + t.Error("perMtok(nil) reported a usable price, want rejected") + } + value, ok := perMtok(new(6e-8)) + if !ok || math.Abs(value-0.06) > 1e-9 { + t.Errorf("perMtok(6e-8) = %v, %v; want 0.06, true", value, ok) + } +} + +// TestListModels_ModalityMapping asserts output modalities become modes (so +// the registry can hide models this provider cannot serve) and input +// modalities become capabilities. +func TestListModels_ModalityMapping(t *testing.T) { + tests := []struct { + name string + capabilities string + wantModes []string + wantCapabilities []string + }{ + { + name: "text only", + capabilities: `{"output_modalities":["text"]}`, + wantModes: []string{"chat", "responses"}, + }, + { + name: "image output becomes image_generation", + capabilities: `{"output_modalities":["image"]}`, + wantModes: []string{"image_generation"}, + }, + { + name: "audio output becomes audio_speech", + capabilities: `{"output_modalities":["audio"]}`, + wantModes: []string{"audio_speech"}, + }, + { + name: "embedding output", + capabilities: `{"output_modalities":["embeddings"]}`, + wantModes: []string{"embedding"}, + }, + { + name: "multimodal input adds capabilities", + capabilities: `{"output_modalities":["text"],"input_modalities":["text","image","audio","video"]}`, + wantModes: []string{"chat", "responses"}, + wantCapabilities: []string{"vision", "audio", "video"}, + }, + { + name: "unknown modality is ignored", + capabilities: `{"output_modalities":["telepathy"]}`, + wantModes: nil, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + payload := `{"object":"list","data":[{"id":"m","object":"model","capabilities":` + tt.capabilities + `}]}` + model := firstModel(t, payload) + + var modes []string + if model.Metadata != nil { + modes = model.Metadata.Modes + } + if !slices.Equal(modes, tt.wantModes) { + t.Errorf("Modes = %v, want %v", modes, tt.wantModes) + } + for _, capability := range tt.wantCapabilities { + if model.Metadata == nil || !model.Metadata.Capabilities[capability] { + t.Errorf("Capabilities missing %q", capability) + } + } + }) + } +} + +// TestListModels_SkipsInvalidEntriesAndKeepsBareOnes asserts a blank ID is +// dropped and an entry with nothing to enrich keeps Metadata nil rather than +// an empty struct that enrichment would read as a real provider report. +func TestListModels_SkipsInvalidEntriesAndKeepsBareOnes(t *testing.T) { + provider, _, _ := modelsServer(t, `{"object":"list","data":[ + {"id":" ","object":"model"}, + {"id":"openai/gpt-4","object":"model"}, + {"id":"openai/gpt-5","object":"","context_length":0} + ]}`) + + resp, err := provider.ListModels(context.Background()) + if err != nil { + t.Fatalf("ListModels() error = %v", err) + } + if len(resp.Data) != 2 { + t.Fatalf("models = %+v, want the blank ID dropped", resp.Data) + } + if resp.Data[0].ID != "openai/gpt-4" || resp.Data[0].Metadata != nil { + t.Errorf("bare entry = %+v, want nil Metadata", resp.Data[0]) + } + if resp.Data[1].Object != "model" { + t.Errorf("Object = %q, want the default \"model\" applied", resp.Data[1].Object) + } +} + +// TestListModels_PropagatesUpstreamError asserts a failed catalog fetch +// surfaces as an error, so the registry records the failure and keeps the +// provider registered for the next refresh instead of publishing an empty +// catalog as though Eden had no models. +func TestListModels_PropagatesUpstreamError(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Error(w, `{"error":{"message":"unauthorized"}}`, http.StatusUnauthorized) + })) + defer server.Close() + + provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + resp, err := provider.ListModels(context.Background()) + if err == nil { + t.Fatalf("ListModels() error = nil, want the upstream failure; resp = %+v", resp) + } + if resp != nil { + t.Errorf("ListModels() resp = %+v, want nil on error", resp) + } +} + +func assertPrice(t *testing.T, name string, got *float64, want float64) { + t.Helper() + if got == nil { + t.Errorf("%s = nil, want %v", name, want) + return + } + if math.Abs(*got-want) > 1e-9 { + t.Errorf("%s = %v, want %v", name, *got, want) + } +} + +func assertOptionalPrice(t *testing.T, name string, got, want *float64) { + t.Helper() + if want == nil { + if got != nil { + t.Errorf("%s = %v, want nil (an unreported rate must not be costed)", name, *got) + } + return + } + assertPrice(t, name, got, *want) +} diff --git a/internal/providers/edenai/passthrough_semantics.go b/internal/providers/edenai/passthrough_semantics.go new file mode 100644 index 000000000..70eca867e --- /dev/null +++ b/internal/providers/edenai/passthrough_semantics.go @@ -0,0 +1,14 @@ +package edenai + +import "github.com/enterpilot/gomodel/internal/providers" + +// Eden AI's chat completions and embeddings routes are OpenAI-shaped, so they +// carry OpenAI's semantics and audit paths. Eden's /responses route is +// deliberately absent: it is not the OpenAI Responses API, and labelling it as +// one would attach a /v1/responses audit path to a request whose body follows +// a different contract. Unlisted endpoints keep the generic /p/edenai/... +// audit path from SemanticEnricher. +var passthroughSemanticEnricher = providers.NewSemanticEnricher("edenai", map[string]providers.PassthroughEndpointSemantics{ + "/chat/completions": {Operation: "edenai.chat_completions", GenAIOperation: "chat", AuditPath: "/v1/chat/completions"}, + "/embeddings": {Operation: "edenai.embeddings", GenAIOperation: "embeddings", AuditPath: "/v1/embeddings"}, +}) diff --git a/internal/providers/edenai/passthrough_semantics_test.go b/internal/providers/edenai/passthrough_semantics_test.go new file mode 100644 index 000000000..698b4652f --- /dev/null +++ b/internal/providers/edenai/passthrough_semantics_test.go @@ -0,0 +1,81 @@ +package edenai + +import ( + "testing" + + "github.com/enterpilot/gomodel/internal/core" +) + +func TestPassthroughSemanticEnricher(t *testing.T) { + if got := passthroughSemanticEnricher.ProviderType(); got != "edenai" { + t.Fatalf("ProviderType() = %q, want edenai", got) + } + + tests := []struct { + name string + rawEndpoint string + normalizedEndpoint string + wantOperation string + wantGenAIOperation string + wantAuditPath string + }{ + { + name: "chat completions", + rawEndpoint: "v1/chat/completions", + normalizedEndpoint: "chat/completions", + wantOperation: "edenai.chat_completions", + wantGenAIOperation: "chat", + wantAuditPath: "/v1/chat/completions", + }, + { + name: "embeddings", + rawEndpoint: "v1/embeddings", + normalizedEndpoint: "embeddings", + wantOperation: "edenai.embeddings", + wantGenAIOperation: "embeddings", + wantAuditPath: "/v1/embeddings", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := passthroughSemanticEnricher.Enrich(nil, nil, &core.PassthroughRouteInfo{ + RawEndpoint: tt.rawEndpoint, + NormalizedEndpoint: tt.normalizedEndpoint, + }) + if got == nil { + t.Fatal("Enrich() returned nil") + } + if got.SemanticOperation != tt.wantOperation { + t.Errorf("SemanticOperation = %q, want %q", got.SemanticOperation, tt.wantOperation) + } + if got.GenAIOperation != tt.wantGenAIOperation { + t.Errorf("GenAIOperation = %q, want %q", got.GenAIOperation, tt.wantGenAIOperation) + } + if got.AuditPath != tt.wantAuditPath { + t.Errorf("AuditPath = %q, want %q", got.AuditPath, tt.wantAuditPath) + } + }) + } +} + +// TestPassthroughSemanticEnricher_ResponsesIsNotOpenAIShaped asserts Eden's +// native /responses route is deliberately absent from the table. Eden's +// /v3/responses is not the OpenAI Responses API, so labelling it +// "edenai.responses" with a /v1/responses audit path would record a request as +// something it is not. It must fall through to the generic /p/edenai/... path. +func TestPassthroughSemanticEnricher_ResponsesIsNotOpenAIShaped(t *testing.T) { + got := passthroughSemanticEnricher.Enrich(nil, nil, &core.PassthroughRouteInfo{ + RawEndpoint: "v1/responses", + NormalizedEndpoint: "responses", + }) + if got == nil { + t.Fatal("Enrich() returned nil") + } + if got.SemanticOperation != "" { + t.Errorf("SemanticOperation = %q, want empty: Eden /responses must not be advertised as OpenAI Responses", got.SemanticOperation) + } + if got.AuditPath != "/p/edenai/responses" { + t.Errorf("AuditPath = %q, want /p/edenai/responses", got.AuditPath) + } +} diff --git a/internal/providers/edenai/response.go b/internal/providers/edenai/response.go new file mode 100644 index 000000000..9d87b3100 --- /dev/null +++ b/internal/providers/edenai/response.go @@ -0,0 +1,126 @@ +package edenai + +import ( + "bytes" + "math" + "strings" + + "github.com/goccy/go-json" + + "github.com/enterpilot/gomodel/internal/core" +) + +const ( + // costField is Eden's per-request charge in USD. Eden reports it at the + // response root, next to "choices", rather than inside "usage" where + // OpenRouter and xAI put theirs. + costField = "cost" + // upstreamProviderField carries Eden's own "provider" member — the + // upstream that actually served the request — after it is moved off the + // typed field. See normalizeUpstreamProvider. + upstreamProviderField = "edenai_upstream_provider" +) + +// normalizeChatResponse reconciles Eden's response extensions with GoModel's +// response semantics. It is applied to every chat completion, which also +// covers the Responses surface because that is translated through chat. +func normalizeChatResponse(resp *core.ChatResponse) { + if resp == nil { + return + } + liftResponseCost(resp) + normalizeUpstreamProvider(resp) +} + +// liftResponseCost copies Eden's root-level "cost" into Usage.RawUsage, where +// the usage pipeline already looks for a provider-reported exact cost. +// +// internal/usage builds its rawData exclusively from the usage object, so a +// root-level member is invisible to cost accounting. Moving the value one +// level down here — rather than teaching the usage pipeline to read response +// roots — keeps the Eden-specific knowledge inside the Eden provider and lets +// the existing, well-tested usage.cost path do the accounting. +// +// The value is left in ExtraFields as well, so clients still receive Eden's +// cost member verbatim. An existing usage.cost wins: if Eden ever also +// reports it in the conventional place, that reading is the more specific one. +func liftResponseCost(resp *core.ChatResponse) { + cost, ok := decodeCost(resp.ExtraFields.Lookup(costField)) + if !ok { + return + } + if resp.Usage.RawUsage == nil { + resp.Usage.RawUsage = make(map[string]any, 1) + } + if _, exists := resp.Usage.RawUsage[costField]; exists { + return + } + resp.Usage.RawUsage[costField] = cost +} + +// decodeCost parses a raw JSON cost member, rejecting anything that would +// corrupt downstream accounting: absent, null, non-numeric, negative, NaN, or +// infinite. +// +// The null check is load-bearing rather than defensive: unmarshalling a JSON +// null into a float64 is a no-op that reports no error, so without it a +// `"cost": null` would be lifted as a real $0.00 charge and silently +// understate spend. +func decodeCost(raw json.RawMessage) (float64, bool) { + trimmed := bytes.TrimSpace(raw) + if core.IsJSONNull(trimmed) { + return 0, false + } + var cost float64 + if err := json.Unmarshal(trimmed, &cost); err != nil { + return 0, false + } + if cost < 0 || math.IsNaN(cost) || math.IsInf(cost, 0) { + return 0, false + } + return cost, true +} + +// normalizeUpstreamProvider clears Eden's "provider" member off the typed +// field and re-exposes it under an Eden-namespaced key. +// +// Eden reports the upstream that served the request ("openai", "deepinfra"), +// but core.ChatResponse.Provider means the provider GoModel executed against. +// The gateway treats a populated value as authoritative: it feeds +// gateway.ResponseProviderType, which labels provider-attempt telemetry and +// failover metadata, and it is echoed to the client. Leaving Eden's value +// there would report "openai" as the executing provider for a request that +// actually ran through Eden, mislabeling both. Clearing it lets +// ResponseProviderType fall back to the configured provider type ("edenai"), +// which is the accurate answer, while the namespaced extra field keeps the +// upstream visible to clients that want it. +func normalizeUpstreamProvider(resp *core.ChatResponse) { + upstream := strings.TrimSpace(resp.Provider) + resp.Provider = "" + if upstream == "" { + return + } + encoded, err := json.Marshal(upstream) + if err != nil { + return + } + merged, err := core.MergeUnknownJSONFields(resp.ExtraFields, map[string]json.RawMessage{ + upstreamProviderField: encoded, + }) + if err != nil { + return + } + resp.ExtraFields = merged +} + +// normalizeEmbeddingResponse applies the same provider-field reasoning to +// embeddings, which Eden also annotates with the upstream provider and which +// feeds the same gateway.ResponseProviderType labeling. core.EmbeddingResponse +// models no unknown-field container, so the upstream value is dropped rather +// than relocated. +func normalizeEmbeddingResponse(resp *core.EmbeddingResponse) { + if resp == nil { + return + } + resp.Provider = "" +} diff --git a/internal/providers/edenai/response_test.go b/internal/providers/edenai/response_test.go new file mode 100644 index 000000000..93c3aa0fe --- /dev/null +++ b/internal/providers/edenai/response_test.go @@ -0,0 +1,225 @@ +package edenai + +import ( + "context" + "math" + "net/http" + "net/http/httptest" + "testing" + + "github.com/goccy/go-json" + + "github.com/enterpilot/gomodel/internal/core" + "github.com/enterpilot/gomodel/internal/llmclient" +) + +// edenChatResponse is a verbatim non-streaming Eden /v3/chat/completions body. +// Note where Eden puts its extensions: `cost` and `provider` sit at the +// response root, not inside `usage`. +const edenChatResponse = `{ + "status": "success", + "id": "chatcmpl-eden", + "created": 1741015112, + "model": "gpt-4o-mini-2024-07-18", + "object": "chat.completion", + "choices": [{"index":0,"message":{"role":"assistant","content":"hi"},"finish_reason":"stop"}], + "usage": { + "completion_tokens": 99, + "prompt_tokens": 1170, + "total_tokens": 1269 + }, + "service_tier": "default", + "cost": 0.0002349, + "provider": "openai" +}` + +func chatResponseFrom(t *testing.T, payload string) *core.ChatResponse { + t.Helper() + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(payload)) + })) + t.Cleanup(server.Close) + + provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ + Model: slashedModel, + Messages: []core.Message{{Role: "user", Content: "hi"}}, + }) + if err != nil { + t.Fatalf("ChatCompletion() error = %v", err) + } + return resp +} + +// TestChatCompletion_LiftsRootCostIntoRawUsage is the load-bearing test for +// exact-cost accounting: Eden reports cost at the response root, but +// internal/usage builds its rawData from the usage object, so the provider +// must relocate it for the existing cost path to see it. +func TestChatCompletion_LiftsRootCostIntoRawUsage(t *testing.T) { + resp := chatResponseFrom(t, edenChatResponse) + + cost, ok := resp.Usage.RawUsage["cost"] + if !ok { + t.Fatalf("Usage.RawUsage = %v, want Eden's root-level cost lifted in", resp.Usage.RawUsage) + } + value, ok := cost.(float64) + if !ok || math.Abs(value-0.0002349) > 1e-12 { + t.Fatalf("Usage.RawUsage[\"cost\"] = %#v, want 0.0002349", cost) + } + if resp.Usage.PromptTokens != 1170 || resp.Usage.CompletionTokens != 99 { + t.Errorf("token counts = %+v, want the reported usage preserved", resp.Usage) + } +} + +// TestChatCompletion_KeepsCostVisibleToClients asserts lifting the value into +// RawUsage does not remove it from the response the client receives. +func TestChatCompletion_KeepsCostVisibleToClients(t *testing.T) { + resp := chatResponseFrom(t, edenChatResponse) + + encoded, err := json.Marshal(resp) + if err != nil { + t.Fatalf("Marshal() error = %v", err) + } + var decoded map[string]any + if err := json.Unmarshal(encoded, &decoded); err != nil { + t.Fatalf("Unmarshal() error = %v", err) + } + if cost, ok := decoded["cost"].(float64); !ok || math.Abs(cost-0.0002349) > 1e-12 { + t.Errorf("serialized cost = %#v, want Eden's cost preserved for the client", decoded["cost"]) + } +} + +// TestChatCompletion_RejectsUnusableCost asserts a cost that would corrupt +// accounting is dropped rather than lifted, leaving the request to fall back +// to token-derived pricing. +func TestChatCompletion_RejectsUnusableCost(t *testing.T) { + tests := []struct { + name string + cost string + }{ + {name: "absent", cost: ""}, + {name: "null", cost: `"cost": null,`}, + {name: "negative", cost: `"cost": -0.5,`}, + {name: "non-numeric", cost: `"cost": "free",`}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + payload := `{"id":"c","created":1,"model":"m","choices":[],` + tt.cost + + `"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}` + resp := chatResponseFrom(t, payload) + if _, ok := resp.Usage.RawUsage["cost"]; ok { + t.Fatalf("Usage.RawUsage = %v, want no cost lifted for an unusable value", resp.Usage.RawUsage) + } + }) + } +} + +// TestChatCompletion_UsageLevelCostWins asserts the conventional location +// takes precedence if Eden ever reports cost in both places. +func TestChatCompletion_UsageLevelCostWins(t *testing.T) { + payload := `{"id":"c","created":1,"model":"m","choices":[],"cost":9.99,` + + `"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2,"cost":0.5}}` + resp := chatResponseFrom(t, payload) + + value, ok := resp.Usage.RawUsage["cost"].(float64) + if !ok || value != 0.5 { + t.Fatalf("Usage.RawUsage[\"cost\"] = %#v, want the usage-level 0.5", resp.Usage.RawUsage["cost"]) + } +} + +// TestChatCompletion_DoesNotReportEdenUpstreamAsExecutingProvider asserts +// Eden's "provider" member is cleared off the typed field. The gateway treats +// a populated ChatResponse.Provider as the provider it executed against +// (gateway.ResponseProviderType feeds attempt telemetry and failover +// metadata), so leaving Eden's upstream there would report "openai" for a +// request that ran through Eden. Clearing it lets the gateway fall back to +// the configured provider type. +func TestChatCompletion_DoesNotReportEdenUpstreamAsExecutingProvider(t *testing.T) { + resp := chatResponseFrom(t, edenChatResponse) + + if resp.Provider != "" { + t.Fatalf("Provider = %q, want empty so the gateway labels the request edenai", resp.Provider) + } +} + +// TestChatCompletion_PreservesEdenUpstreamProvider asserts the upstream is not +// simply discarded: it is re-exposed under an Eden-namespaced key that cannot +// collide with the typed provider member. +func TestChatCompletion_PreservesEdenUpstreamProvider(t *testing.T) { + resp := chatResponseFrom(t, edenChatResponse) + + raw := resp.ExtraFields.Lookup(upstreamProviderField) + if len(raw) == 0 { + t.Fatalf("ExtraFields missing %q", upstreamProviderField) + } + var upstream string + if err := json.Unmarshal(raw, &upstream); err != nil { + t.Fatalf("Unmarshal(%s) error = %v", raw, err) + } + if upstream != "openai" { + t.Errorf("%s = %q, want openai", upstreamProviderField, upstream) + } +} + +// TestChatCompletion_NoUpstreamProviderLeavesNoMarker asserts a response +// without Eden's provider member does not gain an empty namespaced key. +func TestChatCompletion_NoUpstreamProviderLeavesNoMarker(t *testing.T) { + resp := chatResponseFrom(t, `{"id":"c","created":1,"model":"m","choices":[], + "usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`) + + if resp.Provider != "" { + t.Errorf("Provider = %q, want empty", resp.Provider) + } + if raw := resp.ExtraFields.Lookup(upstreamProviderField); len(raw) != 0 { + t.Errorf("ExtraFields[%q] = %s, want absent", upstreamProviderField, raw) + } +} + +// TestEmbeddings_DoesNotReportEdenUpstreamAsExecutingProvider applies the same +// reasoning to embeddings, which feed the same gateway provider labeling. +func TestEmbeddings_DoesNotReportEdenUpstreamAsExecutingProvider(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"object":"list","data":[{"object":"embedding","embedding":[0.1],"index":0}], + "model":"openai/text-embedding-3-small","provider":"openai","cost":0.0001, + "usage":{"prompt_tokens":4,"total_tokens":4}}`)) + })) + defer server.Close() + + provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + resp, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ + Model: "openai/text-embedding-3-small", + Input: "hello", + }) + if err != nil { + t.Fatalf("Embeddings() error = %v", err) + } + if resp.Provider != "" { + t.Fatalf("Provider = %q, want empty so the gateway labels the request edenai", resp.Provider) + } +} + +// TestResponses_InheritsCostLifting asserts the Responses surface picks up the +// same normalization, because it is translated through ChatCompletion. +func TestResponses_InheritsCostLifting(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(edenChatResponse)) + })) + defer server.Close() + + provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{ + Model: slashedModel, + Input: "hi", + }) + if err != nil { + t.Fatalf("Responses() error = %v", err) + } + value, ok := resp.Usage.RawUsage["cost"].(float64) + if !ok || math.Abs(value-0.0002349) > 1e-12 { + t.Fatalf("Responses usage cost = %#v, want 0.0002349 carried through the chat translation", resp.Usage.RawUsage["cost"]) + } +} diff --git a/internal/providers/registry_provider_pricing_test.go b/internal/providers/registry_provider_pricing_test.go new file mode 100644 index 000000000..7957abb35 --- /dev/null +++ b/internal/providers/registry_provider_pricing_test.go @@ -0,0 +1,145 @@ +package providers + +import ( + "context" + "testing" + + "github.com/enterpilot/gomodel/config" + "github.com/enterpilot/gomodel/internal/core" + "github.com/enterpilot/gomodel/internal/modeldata" +) + +// edenPricedModel mirrors what a dynamic-metadata provider (Eden AI, Chutes, +// OpenRouter) returns from ListModels: the model plus the pricing and context +// window it read off its own catalog. +func edenPricedModel() core.Model { + return core.Model{ + ID: "openai/gpt-4", + Object: "model", + OwnedBy: "openai", + Metadata: &core.ModelMetadata{ + Modes: []string{"chat", "responses"}, + Categories: core.CategoriesForModes([]string{"chat", "responses"}), + Capabilities: map[string]bool{"reasoning": true}, + ContextWindow: new(131072), + Pricing: &core.ModelPricing{ + Currency: "USD", + InputPerMtok: new(0.06), + OutputPerMtok: new(0.18), + CachedInputPerMtok: new(0.012), + }, + }, + } +} + +// TestInitialize_ProviderReportedPricingSurvivesCatalogMiss is the registry +// half of dynamic provider metadata. A gateway-of-gateways like Eden AI has +// none of its models in the central catalog (entries are keyed +// "/", so "edenai/openai/gpt-4" never resolves), and +// enrichment must therefore fall back to what the provider itself reported +// rather than dropping it. Without that, pricing never reaches +// ResolvePricing, and Eden spend stays invisible to budgets and cost routing. +func TestInitialize_ProviderReportedPricingSurvivesCatalogMiss(t *testing.T) { + registry := NewModelRegistry() + provider := ®istryMockProvider{ + name: "provider-edenai", + modelsResponse: &core.ModelsResponse{ + Object: "list", + Data: []core.Model{edenPricedModel()}, + }, + } + registry.RegisterProviderWithNameAndType(provider, "edenai", "edenai") + + if err := registry.Initialize(context.Background()); err != nil { + t.Fatalf("Initialize: %v", err) + } + + // A populated catalog that knows nothing about this provider: the exact + // condition Eden models are always in. + registry.setModelListAndEnrich(&modeldata.ModelList{ + Models: map[string]modeldata.ModelEntry{}, + ProviderModels: map[string]modeldata.ProviderModelEntry{ + "openai/gpt-4": {}, + }, + }, nil, "etag", "https://example.invalid/models.json") + + info := registry.GetModel("edenai/openai/gpt-4") + if info == nil { + t.Fatal("model not registered under its provider-qualified ID") + } + if info.Discovered == nil || info.Discovered.Pricing == nil { + t.Fatalf("Discovered = %+v, want the provider's own report retained", info.Discovered) + } + meta := info.Model.Metadata + if meta == nil || meta.Pricing == nil { + t.Fatalf("Metadata = %+v, want provider pricing to survive enrichment", meta) + } + assertPricePtr(t, "InputPerMtok", meta.Pricing.InputPerMtok, 0.06) + assertPricePtr(t, "OutputPerMtok", meta.Pricing.OutputPerMtok, 0.18) + assertPricePtr(t, "CachedInputPerMtok", meta.Pricing.CachedInputPerMtok, 0.012) + if meta.ContextWindow == nil || *meta.ContextWindow != 131072 { + t.Errorf("ContextWindow = %v, want 131072", meta.ContextWindow) + } + if !meta.Capabilities["reasoning"] { + t.Errorf("Capabilities = %v, want reasoning retained", meta.Capabilities) + } +} + +// TestResolvePricing_UsesProviderReportedPricing closes the loop to the cost +// path: ResolvePricing is what internal/usage calls per request, and it must +// find the dynamically discovered rates. +func TestResolvePricing_UsesProviderReportedPricing(t *testing.T) { + registry := NewModelRegistry() + provider := ®istryMockProvider{ + name: "provider-edenai", + modelsResponse: &core.ModelsResponse{ + Object: "list", + Data: []core.Model{edenPricedModel()}, + }, + } + registry.RegisterProviderWithNameAndType(provider, "edenai", "edenai") + + if err := registry.Initialize(context.Background()); err != nil { + t.Fatalf("Initialize: %v", err) + } + + for _, selector := range []string{"edenai/openai/gpt-4", "openai/gpt-4"} { + pricing := registry.ResolvePricing(selector, "edenai") + if pricing == nil { + t.Fatalf("ResolvePricing(%q) = nil, want the provider-reported rates", selector) + } + assertPricePtr(t, selector+" InputPerMtok", pricing.InputPerMtok, 0.06) + assertPricePtr(t, selector+" OutputPerMtok", pricing.OutputPerMtok, 0.18) + } +} + +// TestModelFilter_AdmitsProviderPricedModels pins the pre-request consumer of +// dynamic pricing: a max-price filter drops unpriced models by design, so a +// provider that reports no pricing has its whole catalog excluded. Reporting +// pricing is what makes the cap usable. +func TestModelFilter_AdmitsProviderPricedModels(t *testing.T) { + priced := edenPricedModel() + unpriced := core.Model{ID: "openai/gpt-5", Object: "model"} + + filter, active := newModelFilter(config.ModelFilter{MaxPricePerMtok: new(1.0)}) + if !active { + t.Fatal("newModelFilter reported an inactive filter for a price cap") + } + if !filter.keep(priced) { + t.Error("priced model rejected by a 1.0/MTok cap, want admitted (0.18 max rate)") + } + if filter.keep(unpriced) { + t.Error("unpriced model admitted by a price cap, want dropped") + } +} + +func assertPricePtr(t *testing.T, name string, got *float64, want float64) { + t.Helper() + if got == nil { + t.Errorf("%s = nil, want %v", name, want) + return + } + if *got != want { + t.Errorf("%s = %v, want %v", name, *got, want) + } +} diff --git a/internal/server/handlers_test.go b/internal/server/handlers_test.go index 5dbbe1b31..0f18343b1 100644 --- a/internal/server/handlers_test.go +++ b/internal/server/handlers_test.go @@ -7088,7 +7088,7 @@ func TestProviderPassthrough_RejectsUnsupportedProvider(t *testing.T) { if !strings.Contains(rec.Body.String(), `provider passthrough for \"groq\" is not enabled`) { t.Fatalf("unexpected error body: %s", rec.Body.String()) } - if !strings.Contains(rec.Body.String(), "anthropic, deepseek, hetzner, kilo, llamacpp, llmd, openai, openrouter, sglang, vllm, zai") { + if !strings.Contains(rec.Body.String(), "anthropic, deepseek, edenai, hetzner, kilo, llamacpp, llmd, openai, openrouter, sglang, vllm, zai") { t.Fatalf("unexpected error body: %s", rec.Body.String()) } } diff --git a/internal/server/passthrough_support.go b/internal/server/passthrough_support.go index ca48b69a0..bbd84eafd 100644 --- a/internal/server/passthrough_support.go +++ b/internal/server/passthrough_support.go @@ -20,7 +20,7 @@ import ( "github.com/enterpilot/gomodel/internal/usage" ) -var defaultEnabledPassthroughProviders = []string{"openai", "anthropic", "openrouter", "kilo", "zai", "sglang", "vllm", "llamacpp", "llmd", "deepseek", "hetzner"} +var defaultEnabledPassthroughProviders = []string{"openai", "anthropic", "openrouter", "kilo", "zai", "sglang", "vllm", "llamacpp", "llmd", "deepseek", "hetzner", "edenai"} const llmdDroppedReasonHeader = "X-Llm-D-Request-Dropped-Reason" diff --git a/internal/server/passthrough_support_test.go b/internal/server/passthrough_support_test.go index 53dd3aff0..e04fedf86 100644 --- a/internal/server/passthrough_support_test.go +++ b/internal/server/passthrough_support_test.go @@ -11,6 +11,7 @@ import ( "github.com/labstack/echo/v5" + "github.com/enterpilot/gomodel/config" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/usage" ) @@ -45,6 +46,48 @@ func TestDefaultEnabledPassthroughProvidersIncludesHetzner(t *testing.T) { } } +// TestDefaultEnabledPassthroughProvidersIncludesEdenAI asserts that the default +// allowlist contains edenai — the provider matrix marks edenai passthrough ✅, +// and the default handler must not reject those requests before contacting the +// upstream. +func TestDefaultEnabledPassthroughProvidersIncludesEdenAI(t *testing.T) { + found := slices.Contains(defaultEnabledPassthroughProviders, "edenai") + if !found { + t.Fatalf("defaultEnabledPassthroughProviders = %v, want edenai included", defaultEnabledPassthroughProviders) + } +} + +// TestDefaultEnabledPassthroughProvidersMatchesConfigDefault keeps the two +// passthrough allowlist defaults from drifting apart. +// +// This package's slice is only the fallback for a Handler built without +// config; the list a running gateway actually enforces comes from +// config.Config.Server.EnabledPassthroughProviders, which http.go applies over +// the fallback. Adding a provider to one and not the other compiles, passes +// every handler test (they construct Handlers directly and so read the +// fallback), and still rejects the provider at runtime with "passthrough for +// X is not enabled" — which is exactly how the edenai entry was first missed. +func TestDefaultEnabledPassthroughProvidersMatchesConfigDefault(t *testing.T) { + // Load() with no config file present yields the built-in defaults; any + // ENABLED_PASSTHROUGH_PROVIDERS in the environment would mask them. + t.Setenv("ENABLED_PASSTHROUGH_PROVIDERS", "") + loaded, err := config.Load() + if err != nil { + t.Fatalf("config.Load() error = %v", err) + } + + fromConfig := append([]string(nil), loaded.Config.Server.EnabledPassthroughProviders...) + fromServer := append([]string(nil), defaultEnabledPassthroughProviders...) + slices.Sort(fromConfig) + slices.Sort(fromServer) + + if !slices.Equal(fromConfig, fromServer) { + t.Fatalf("passthrough allowlist defaults disagree:\n config/config.go: %v\n internal/server: %v\n"+ + "both must list the same provider types, or the runtime default silently differs from the tested one", + fromConfig, fromServer) + } +} + // A successful non-streaming JSON passthrough response must produce a usage // entry from its usage member — the same accounting SSE streams get from the // stream usage observer. Covers the /p/{provider} surface directly. diff --git a/internal/usage/cost.go b/internal/usage/cost.go index c0be98f9e..a640acb8b 100644 --- a/internal/usage/cost.go +++ b/internal/usage/cost.go @@ -17,6 +17,7 @@ const ( CostSourceModelPricing = "model_pricing" CostSourceOpenRouterCredits = "openrouter_credits" CostSourceXAITicks = "xai_cost_in_usd_ticks" + CostSourceEdenAICost = "edenai_cost" ) // xAI reports usage.cost_in_usd_ticks as USD-denominated ticks, where 10^10 @@ -519,6 +520,9 @@ func CalculateUsageCost(inputTokens, outputTokens int, rawData map[string]any, p if result, ok := xaiTicksCost(rawData, providerType); ok { return result } + if result, ok := edenAICost(rawData, providerType); ok { + return result + } return CalculateGranularCost(inputTokens, outputTokens, rawData, providerType, pricing) } @@ -568,6 +572,34 @@ func isXAIProvider(providerType string) bool { return strings.EqualFold(strings.TrimSpace(providerType), "xai") } +// edenAICost uses Eden AI's own per-request charge instead of recomputing cost +// from token counts. Eden is a multi-provider gateway that reprices upstreams +// automatically and applies per-account discounts, so the figure it returns is +// authoritative in a way a rate-card reconstruction cannot be. +// +// Eden publishes cost at the response root rather than inside usage; the Eden +// provider moves it into RawUsage so it arrives here on the same path as +// OpenRouter's and xAI's. Eden reports no input/output split, so only the +// total is set — the same shape xaiTicksCost produces. +func edenAICost(rawData map[string]any, providerType string) (CostResult, bool) { + if !isEdenAIProvider(providerType) { + return CostResult{}, false + } + total, ok := extractFloat(rawData, "cost") + if !ok || !isFiniteCost(total) || total < 0 { + return CostResult{}, false + } + + return CostResult{ + TotalCost: &total, + Source: CostSourceEdenAICost, + }, true +} + +func isEdenAIProvider(providerType string) bool { + return strings.EqualFold(strings.TrimSpace(providerType), "edenai") +} + func openRouterCreditCostSplit(rawData map[string]any, total float64) (float64, float64, bool) { details, ok := nestedUsageMap(rawData["cost_details"]) if !ok { diff --git a/internal/usage/cost_test.go b/internal/usage/cost_test.go index 597048afe..f0e9ac960 100644 --- a/internal/usage/cost_test.go +++ b/internal/usage/cost_test.go @@ -820,3 +820,100 @@ func assertCostNear(t *testing.T, name string, got *float64, want float64) { t.Fatalf("%s = %f, want %f", name, *got, want) } } + +// --- Eden AI exact request cost --- + +// Eden AI is a multi-provider gateway that reports the exact USD charge for +// each request. That figure is authoritative — it already reflects Eden's +// upstream repricing and any account discount — so it must win over any rate +// card GoModel happens to hold for the model. +func TestCalculateUsageCost_EdenAICostOverridesStaticPricing(t *testing.T) { + pricing := &core.ModelPricing{ + InputPerMtok: new(100.0), + OutputPerMtok: new(100.0), + } + + result := CalculateUsageCost(1170, 99, map[string]any{"cost": 0.0002349}, "edenai", pricing) + + assertCostNear(t, "TotalCost", result.TotalCost, 0.0002349) + if result.Source != CostSourceEdenAICost { + t.Fatalf("Source = %q, want %q", result.Source, CostSourceEdenAICost) + } + // Eden reports no input/output split, so only the total is claimed rather + // than inventing a division of it. + if result.InputCost != nil || result.OutputCost != nil { + t.Fatalf("InputCost/OutputCost = %v/%v, want nil (Eden reports no split)", result.InputCost, result.OutputCost) + } +} + +func TestCalculateUsageCost_EdenAIAcceptsZeroCost(t *testing.T) { + result := CalculateUsageCost(10, 4, map[string]any{"cost": 0.0}, "edenai", nil) + + assertCostNear(t, "TotalCost", result.TotalCost, 0) + if result.Source != CostSourceEdenAICost { + t.Fatalf("Source = %q, want %q", result.Source, CostSourceEdenAICost) + } +} + +// Without a usable cost the request must fall back to the ordinary token math +// rather than recording nothing or a corrupt figure. +func TestCalculateUsageCost_EdenAIFallsBackToModelPricingWhenCostUnusable(t *testing.T) { + pricing := &core.ModelPricing{ + InputPerMtok: new(1.0), + OutputPerMtok: new(2.0), + } + + tests := []struct { + name string + rawData map[string]any + }{ + {name: "absent", rawData: map[string]any{}}, + {name: "nil raw data", rawData: nil}, + {name: "negative", rawData: map[string]any{"cost": -0.5}}, + {name: "not a number", rawData: map[string]any{"cost": "free"}}, + {name: "NaN", rawData: map[string]any{"cost": math.NaN()}}, + {name: "positive infinity", rawData: map[string]any{"cost": math.Inf(1)}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := CalculateUsageCost(1_000_000, 500_000, tt.rawData, "edenai", pricing) + + assertCostNear(t, "InputCost", result.InputCost, 1.0) + assertCostNear(t, "OutputCost", result.OutputCost, 1.0) + assertCostNear(t, "TotalCost", result.TotalCost, 2.0) + if result.Source != CostSourceModelPricing { + t.Fatalf("Source = %q, want %q", result.Source, CostSourceModelPricing) + } + }) + } +} + +// The Eden reading is provider-gated like OpenRouter's and xAI's: a "cost" +// member from any other provider keeps its existing meaning. +func TestCalculateUsageCost_EdenAICostIgnoredForOtherProviders(t *testing.T) { + pricing := &core.ModelPricing{ + InputPerMtok: new(1.0), + OutputPerMtok: new(2.0), + } + + result := CalculateUsageCost(1_000_000, 500_000, map[string]any{"cost": 0.0002349}, "openai", pricing) + + assertCostNear(t, "TotalCost", result.TotalCost, 2.0) + if result.Source != CostSourceModelPricing { + t.Fatalf("Source = %q, want %q", result.Source, CostSourceModelPricing) + } +} + +// With no pricing and no usable cost the entry stays uncosted rather than +// recording a fabricated zero. +func TestCalculateUsageCost_EdenAIWithoutCostOrPricingRecordsNothing(t *testing.T) { + result := CalculateUsageCost(10, 4, map[string]any{}, "edenai", nil) + + if result.TotalCost != nil { + t.Fatalf("TotalCost = %v, want nil", *result.TotalCost) + } + if result.Source != "" { + t.Fatalf("Source = %q, want empty", result.Source) + } +} diff --git a/internal/usage/extractor_test.go b/internal/usage/extractor_test.go index da12638f6..210df6666 100644 --- a/internal/usage/extractor_test.go +++ b/internal/usage/extractor_test.go @@ -1019,3 +1019,67 @@ func TestExtractFromEmbeddingResponse_NoUsageCaveat(t *testing.T) { t.Fatalf("repricing retained %q, want the caveat kept for tiered token rates", retained) } } + +// TestExtractFromChatResponse_EdenAIExactCostReachesTotalCost closes the seam +// between the Eden provider and this package. The provider lifts Eden's +// root-level cost into Usage.RawUsage; this asserts the extractor carries it +// into rawData and that the entry ends up priced from it, so Eden spend +// reaches usage records, budgets, and cost reporting. +func TestExtractFromChatResponse_EdenAIExactCostReachesTotalCost(t *testing.T) { + resp := &core.ChatResponse{ + ID: "chatcmpl-eden", + Model: "gpt-4o-mini-2024-07-18", + Usage: core.Usage{ + PromptTokens: 1170, + CompletionTokens: 99, + TotalTokens: 1269, + RawUsage: map[string]any{"cost": 0.0002349}, + }, + } + // Static pricing that would produce a very different number, to prove the + // exact charge wins rather than merely agreeing by coincidence. + pricing := &core.ModelPricing{InputPerMtok: new(100.0), OutputPerMtok: new(100.0)} + + entry := ExtractFromChatResponse(resp, "req-eden", "edenai", "/v1/chat/completions", pricing) + + if entry == nil { + t.Fatal("ExtractFromChatResponse() = nil") + } + if entry.RawData["cost"] != 0.0002349 { + t.Fatalf("RawData[cost] = %#v, want the lifted 0.0002349", entry.RawData["cost"]) + } + if entry.TotalCost == nil || math.Abs(*entry.TotalCost-0.0002349) > 1e-12 { + t.Fatalf("TotalCost = %v, want 0.0002349", entry.TotalCost) + } + if entry.CostSource != CostSourceEdenAICost { + t.Fatalf("CostSource = %q, want %q", entry.CostSource, CostSourceEdenAICost) + } + if entry.InputTokens != 1170 || entry.OutputTokens != 99 { + t.Errorf("token counts = %d/%d, want 1170/99", entry.InputTokens, entry.OutputTokens) + } +} + +// TestExtractFromChatResponse_EdenAIWithoutCostFallsBackToPricing asserts the +// fallback path: no exact charge means the discovered per-model pricing is +// used, rather than the entry going uncosted. +func TestExtractFromChatResponse_EdenAIWithoutCostFallsBackToPricing(t *testing.T) { + resp := &core.ChatResponse{ + ID: "chatcmpl-eden", + Model: "gpt-4o-mini", + Usage: core.Usage{PromptTokens: 1_000_000, CompletionTokens: 500_000, TotalTokens: 1_500_000}, + } + pricing := &core.ModelPricing{InputPerMtok: new(0.06), OutputPerMtok: new(0.18)} + + entry := ExtractFromChatResponse(resp, "req-eden", "edenai", "/v1/chat/completions", pricing) + + if entry == nil { + t.Fatal("ExtractFromChatResponse() = nil") + } + if entry.CostSource != CostSourceModelPricing { + t.Fatalf("CostSource = %q, want %q", entry.CostSource, CostSourceModelPricing) + } + // 1M * 0.06/1M + 0.5M * 0.18/1M = 0.06 + 0.09 + if entry.TotalCost == nil || math.Abs(*entry.TotalCost-0.15) > 1e-9 { + t.Fatalf("TotalCost = %v, want 0.15 from discovered per-model pricing", entry.TotalCost) + } +} diff --git a/internal/usage/stream_observer.go b/internal/usage/stream_observer.go index 67c572215..322fffba6 100644 --- a/internal/usage/stream_observer.go +++ b/internal/usage/stream_observer.go @@ -232,6 +232,7 @@ func (o *StreamUsageObserver) extractUsageFromEvent(chunk map[string]any) *Usage } copyExtendedUsageFields(rawData, usageMap) + copyRootLevelCost(rawData, chunk) copyUsageDetailsFields(rawData, usageMap["prompt_tokens_details"], "prompt_") copyUsageDetailsFields(rawData, usageMap["input_tokens_details"], "prompt_") copyUsageDetailsFields(rawData, usageMap["completion_tokens_details"], "completion_") @@ -292,6 +293,27 @@ func (o *StreamUsageObserver) pricingProvider() string { return strings.TrimSpace(o.provider) } +// copyRootLevelCost picks up a provider-reported per-request charge published +// at the chunk root rather than inside usage. Eden AI reports cost there, and +// a streamed chunk reaches the observer as raw JSON, so unlike the +// non-streaming path the provider has no opportunity to relocate it first. +// +// A usage-level cost stays authoritative: it is the conventional location, so +// a provider that reports both is taken at its more specific word. The value +// is only ever interpreted as USD for providers that CalculateUsageCost gates +// on, so harvesting it for everyone adds a raw-usage member without changing +// any other provider's cost math. +func copyRootLevelCost(rawData map[string]any, chunk map[string]any) { + if _, exists := rawData["cost"]; exists { + return + } + // A NaN fails this comparison and is dropped, matching how + // copyExtendedUsageFields screens the usage-level member. + if value, ok := numericFloat(chunk["cost"]); ok && value >= 0 { + rawData["cost"] = value + } +} + func copyExtendedUsageFields(rawData map[string]any, usageMap map[string]any) { for key, value := range usageMap { switch key { diff --git a/internal/usage/stream_observer_test.go b/internal/usage/stream_observer_test.go index a88095d6b..e4a36c921 100644 --- a/internal/usage/stream_observer_test.go +++ b/internal/usage/stream_observer_test.go @@ -811,3 +811,146 @@ func TestStreamUsageObserverAnthropicNativeEvents(t *testing.T) { t.Errorf("cache_read_input_tokens = %v, want 200", entry.RawData["cache_read_input_tokens"]) } } + +// Eden AI reports its per-request charge at the chunk root rather than inside +// usage. A streamed chunk reaches the observer as raw JSON, so unlike the +// non-streaming path the provider cannot relocate it first — the observer has +// to harvest it. +func TestStreamUsageObserverEdenAIRootLevelCost(t *testing.T) { + logger := &trackingLogger{enabled: true} + observer := NewStreamUsageObserver(logger, "openai/gpt-4o-mini", "edenai", "req-edenai", "/v1/chat/completions", nil) + observer.OnJSONEvent(map[string]any{ + "id": "chatcmpl-eden", + "model": "gpt-4o-mini-2024-07-18", + "cost": float64(0.0002349), + "provider": "openai", + "usage": map[string]any{ + "prompt_tokens": float64(1170), + "completion_tokens": float64(99), + "total_tokens": float64(1269), + }, + }) + observer.OnStreamClose() + + entries := logger.getEntries() + if len(entries) != 1 { + t.Fatalf("expected 1 entry, got %d", len(entries)) + } + entry := entries[0] + if entry.RawData == nil || entry.RawData["cost"] != 0.0002349 { + t.Fatalf("RawData[cost] = %#v, want 0.0002349 harvested from the chunk root", entry.RawData["cost"]) + } + if entry.TotalCost == nil || *entry.TotalCost != 0.0002349 { + t.Fatalf("TotalCost = %v, want 0.0002349", entry.TotalCost) + } + if entry.CostSource != CostSourceEdenAICost { + t.Fatalf("CostSource = %q, want %q", entry.CostSource, CostSourceEdenAICost) + } +} + +// A usage-level cost is the conventional location, so it stays authoritative +// when a provider reports both. +func TestStreamUsageObserverUsageCostWinsOverRootLevelCost(t *testing.T) { + logger := &trackingLogger{enabled: true} + observer := NewStreamUsageObserver(logger, "openai/gpt-4o-mini", "edenai", "req-edenai", "/v1/chat/completions", nil) + observer.OnJSONEvent(map[string]any{ + "id": "chatcmpl-eden", + "model": "gpt-4o-mini", + "cost": float64(9.99), + "usage": map[string]any{ + "prompt_tokens": float64(10), + "completion_tokens": float64(4), + "total_tokens": float64(14), + "cost": float64(0.5), + }, + }) + observer.OnStreamClose() + + entries := logger.getEntries() + if len(entries) != 1 { + t.Fatalf("expected 1 entry, got %d", len(entries)) + } + if got := entries[0].RawData["cost"]; got != 0.5 { + t.Fatalf("RawData[cost] = %#v, want the usage-level 0.5", got) + } +} + +// An unusable root-level cost must not reach rawData, so the entry falls back +// to token pricing instead of recording a corrupt figure. +func TestStreamUsageObserverRejectsUnusableRootLevelCost(t *testing.T) { + for name, cost := range map[string]any{ + "negative": float64(-1), + "NaN": math.NaN(), + "not a number": "free", + } { + t.Run(name, func(t *testing.T) { + logger := &trackingLogger{enabled: true} + resolver := &streamPricingCaptureResolver{pricing: &core.ModelPricing{ + InputPerMtok: new(1.0), + OutputPerMtok: new(2.0), + }} + observer := NewStreamUsageObserver(logger, "openai/gpt-4o-mini", "edenai", "req-edenai", "/v1/chat/completions", resolver) + observer.OnJSONEvent(map[string]any{ + "id": "chatcmpl-eden", + "model": "gpt-4o-mini", + "cost": cost, + "usage": map[string]any{ + "prompt_tokens": float64(1_000_000), + "completion_tokens": float64(500_000), + "total_tokens": float64(1_500_000), + }, + }) + observer.OnStreamClose() + + entries := logger.getEntries() + if len(entries) != 1 { + t.Fatalf("expected 1 entry, got %d", len(entries)) + } + entry := entries[0] + if _, ok := entry.RawData["cost"]; ok { + t.Fatalf("RawData[cost] = %#v, want the unusable value dropped", entry.RawData["cost"]) + } + if entry.CostSource != CostSourceModelPricing { + t.Fatalf("CostSource = %q, want %q", entry.CostSource, CostSourceModelPricing) + } + if entry.TotalCost == nil || math.Abs(*entry.TotalCost-2.0) > 1e-9 { + t.Fatalf("TotalCost = %v, want the token-priced 2.0", entry.TotalCost) + } + }) + } +} + +// Harvesting the root member must not change any other provider's cost math: +// the value is only ever interpreted as USD for providers CalculateUsageCost +// gates on. +func TestStreamUsageObserverRootLevelCostDoesNotRepriceOtherProviders(t *testing.T) { + logger := &trackingLogger{enabled: true} + resolver := &streamPricingCaptureResolver{pricing: &core.ModelPricing{ + InputPerMtok: new(1.0), + OutputPerMtok: new(2.0), + }} + observer := NewStreamUsageObserver(logger, "gpt-4o", "openai", "req-openai", "/v1/chat/completions", resolver) + observer.OnJSONEvent(map[string]any{ + "id": "chatcmpl-openai", + "model": "gpt-4o", + "cost": float64(9.99), + "usage": map[string]any{ + "prompt_tokens": float64(1_000_000), + "completion_tokens": float64(500_000), + "total_tokens": float64(1_500_000), + }, + }) + observer.OnStreamClose() + + entries := logger.getEntries() + if len(entries) != 1 { + t.Fatalf("expected 1 entry, got %d", len(entries)) + } + entry := entries[0] + if entry.CostSource != CostSourceModelPricing { + t.Fatalf("CostSource = %q, want %q", entry.CostSource, CostSourceModelPricing) + } + if entry.TotalCost == nil || math.Abs(*entry.TotalCost-2.0) > 1e-9 { + t.Fatalf("TotalCost = %v, want the token-priced 2.0", entry.TotalCost) + } +} diff --git a/run/lifecycle_test.go b/run/lifecycle_test.go index f92b6ab7a..70657cbb1 100644 --- a/run/lifecycle_test.go +++ b/run/lifecycle_test.go @@ -258,6 +258,24 @@ func TestMain_KimicodeProviderRegistration(t *testing.T) { } } +func TestMain_EdenAIProviderRegistration(t *testing.T) { + factory := defaultProviderFactory(&config.Config{}) + + registered := factory.RegisteredTypes() + found := slices.Contains(registered, "edenai") + if !found { + t.Fatalf("edenai not in RegisteredTypes() = %v", registered) + } + + provider, err := factory.Create(providers.ProviderConfig{Type: "edenai", APIKey: "test"}) + if err != nil { + t.Fatalf("factory.Create(edenai) error = %v, want nil", err) + } + if provider == nil { + t.Fatal("factory.Create(edenai) returned nil provider") + } +} + func TestMain_HetznerProviderRegistration(t *testing.T) { factory := defaultProviderFactory(&config.Config{}) diff --git a/run/providers.go b/run/providers.go index f5974173b..36bbd0ab4 100644 --- a/run/providers.go +++ b/run/providers.go @@ -13,6 +13,7 @@ import ( "github.com/enterpilot/gomodel/internal/providers/chutes" "github.com/enterpilot/gomodel/internal/providers/cohere" "github.com/enterpilot/gomodel/internal/providers/deepseek" + "github.com/enterpilot/gomodel/internal/providers/edenai" "github.com/enterpilot/gomodel/internal/providers/elevenlabs" "github.com/enterpilot/gomodel/internal/providers/fireworks" "github.com/enterpilot/gomodel/internal/providers/gemini" @@ -58,6 +59,7 @@ func defaultProviderFactory(cfg *config.Config) *providers.ProviderFactory { factory.Add(chutes.Registration) factory.Add(cohere.Registration) factory.Add(deepseek.Registration) + factory.Add(edenai.Registration) factory.Add(elevenlabs.Registration) factory.Add(fireworks.Registration) factory.Add(gemini.Registration) diff --git a/run/providers_test.go b/run/providers_test.go index 23fdadb52..abe9a9dd0 100644 --- a/run/providers_test.go +++ b/run/providers_test.go @@ -175,7 +175,7 @@ var credentialPayloadFields = []string{ func TestDefaultProviderFactoryRegistersAllProviderTypes(t *testing.T) { expected := []string{ - "anthropic", "azure", "bailian", "bedrock", "bedrock-mantle", "chatgpt", "chutes", "cohere", "deepseek", "elevenlabs", + "anthropic", "azure", "bailian", "bedrock", "bedrock-mantle", "chatgpt", "chutes", "cohere", "deepseek", "edenai", "elevenlabs", "fireworks", "gemini", "groq", "hetzner", "kilo", "kimicode", "llamacpp", "llmd", "meta", "minimax", "ollama", "openai", "opencode_go", "openrouter", "oracle", "sglang", "vertex", "vllm", "xai", "xiaomi", "zai", } diff --git a/web/dashboard/src/lib/utils/providerDocs.js b/web/dashboard/src/lib/utils/providerDocs.js index 23cc841eb..953ab1812 100644 --- a/web/dashboard/src/lib/utils/providerDocs.js +++ b/web/dashboard/src/lib/utils/providerDocs.js @@ -29,6 +29,7 @@ const PROVIDER_DOC_SLUGS = new Set([ "chatgpt", "cohere", "deepseek", + "edenai", "elevenlabs", "gemini", "hetzner", diff --git a/web/dashboard/tests/provider-docs.test.js b/web/dashboard/tests/provider-docs.test.js index 17e5f23cf..536c915c2 100644 --- a/web/dashboard/tests/provider-docs.test.js +++ b/web/dashboard/tests/provider-docs.test.js @@ -38,6 +38,7 @@ test("providerDocsUrl links every documented provider to its own page", () => { "chatgpt", "cohere", "deepseek", + "edenai", "elevenlabs", "gemini", "hetzner", From 3de8e0bb29cb956813f4e48b32198a419e7d3050 Mon Sep 17 00:00:00 2001 From: NaDdjg Date: Fri, 11 Sep 2026 11:47:44 +0100 Subject: [PATCH 2/6] fix: polish Eden AI provider integration --- cmd/gomodel/docs/docs.go | 4 + docs/advanced/configuration.mdx | 1 + docs/openapi.json | 4 + docs/providers/edenai.mdx | 68 ++- internal/core/types.go | 12 +- .../providers/edenai/capabilities_test.go | 212 ++++++++ internal/providers/edenai/edenai.go | 74 ++- .../providers/edenai/embeddings_cost_test.go | 271 ++++++++++ internal/providers/edenai/models.go | 148 ++++-- internal/providers/edenai/models_test.go | 189 ++++++- internal/providers/edenai/response.go | 48 +- internal/providers/edenai/response_test.go | 25 + internal/providers/edenai/transport.go | 91 ++++ internal/providers/edenai/transport_test.go | 501 ++++++++++++++++++ internal/usage/cost.go | 14 + internal/usage/extractor.go | 10 +- internal/usage/extractor_test.go | 114 ++++ 17 files changed, 1703 insertions(+), 83 deletions(-) create mode 100644 internal/providers/edenai/capabilities_test.go create mode 100644 internal/providers/edenai/embeddings_cost_test.go create mode 100644 internal/providers/edenai/transport.go create mode 100644 internal/providers/edenai/transport_test.go diff --git a/cmd/gomodel/docs/docs.go b/cmd/gomodel/docs/docs.go index 1c3a7ba37..3e9de104f 100644 --- a/cmd/gomodel/docs/docs.go +++ b/cmd/gomodel/docs/docs.go @@ -9732,6 +9732,10 @@ const docTemplate = `{ "prompt_tokens": { "type": "integer" }, + "raw_usage": { + "type": "object", + "additionalProperties": {} + }, "total_tokens": { "type": "integer" } diff --git a/docs/advanced/configuration.mdx b/docs/advanced/configuration.mdx index 2956e0b5d..c0fc086e8 100644 --- a/docs/advanced/configuration.mdx +++ b/docs/advanced/configuration.mdx @@ -320,6 +320,7 @@ Set these to automatically register providers. No YAML configuration required. | `DEEPSEEK_API_KEY` | DeepSeek | | `OPENROUTER_API_KEY` | OpenRouter | | `KILO_API_KEY` | Kilo AI Gateway | +| `EDENAI_API_KEY` | Eden AI (`EDENAI_BASE_URL` optional) | | `ZAI_API_KEY` | Z.ai | | `XAI_API_KEY` | xAI (Grok) | | `GROQ_API_KEY` | Groq | diff --git a/docs/openapi.json b/docs/openapi.json index 876b7491e..9b9a612c7 100644 --- a/docs/openapi.json +++ b/docs/openapi.json @@ -13472,6 +13472,10 @@ "prompt_tokens": { "type": "integer" }, + "raw_usage": { + "type": "object", + "additionalProperties": {} + }, "total_tokens": { "type": "integer" } diff --git a/docs/providers/edenai.mdx b/docs/providers/edenai.mdx index 3594ef079..80fafdb4d 100644 --- a/docs/providers/edenai.mdx +++ b/docs/providers/edenai.mdx @@ -27,6 +27,15 @@ endpoint: EDENAI_BASE_URL=https://api.edenai.run/v3 ``` + + Use an `https://` endpoint. GoModel withholds the Eden API key from any + cleartext destination, so an `http://` base URL reaches Eden + unauthenticated and fails with a 401 rather than sending the key in the + clear. The same applies to redirects: an HTTPS endpoint that redirects to + `http://` is refused instead of followed. Cleartext to `localhost` is + allowed, which is what keeps a local Eden-compatible proxy usable. + + Or in `config.yaml`: ```yaml @@ -34,6 +43,7 @@ providers: edenai: type: edenai api_key: "${EDENAI_API_KEY}" + # base_url: "https://api.edenai.run/v3" ``` You can also add the credential from the **Providers** page in the admin @@ -71,37 +81,65 @@ the router, filters, and cost strategies use: | Eden field | GoModel metadata | | ---------- | ---------------- | | `context_length` | context window | -| `capabilities.supports_*` | capabilities (`reasoning`, `function_calling`, `prompt_caching`, …) | +| `capabilities.supports_*` | capabilities, with the prefix stripped (`reasoning`, `function_calling`, `prompt_caching`, …) | | `capabilities.input_modalities` | `vision` / `audio` / `video` capabilities | | `capabilities.output_modalities` | modes and categories | | `pricing` | per-model pricing (see below) | -Output modalities also decide what GoModel advertises: Eden's image- and -speech-only models are filtered out of `/v1/models`, because this provider -serves chat, embeddings, and passthrough only. +Output modalities also decide what GoModel advertises: a model whose only +output is audio or images is left out of `/v1/models`, because Eden here +serves chat, embeddings, and passthrough only. In practice Eden's catalog +publishes text output for every model, including the handful that also return +images, so nothing is currently filtered. ## Pricing Eden publishes per-token USD rates per model, and GoModel converts them to its per-million-token representation (`input_cost_per_token: 6e-8` → `$0.06 / -MTok`). Rates come from Eden's `pricing` block, which is what the account is -actually charged; the undiscounted `list_pricing` block is ignored. +MTok`). + +Rates come from Eden's `pricing` block, which is what the account is actually +charged — the undiscounted `list_pricing` with any account discount already +applied. Each rate is resolved on its own: a usable account rate always wins, +and `list_pricing` supplies only the individual rates `pricing` does not carry. +Eden currently publishes the same rates in both blocks, so in practice +everything resolves from `pricing`; the per-rate fallback is what keeps a +partially priced model from losing the rates it is missing, which would +otherwise bill those token types at $0. + +A rate Eden reports as `0` is treated as genuinely free rather than missing, +and a rate neither block publishes usably is left unset rather than invented. + +These rates cover input, output, cache reads, and cache writes. Eden's +context-length-tiered rates (`input_cost_per_token_above_200k_tokens`), its +`tiered_pricing` list, and its per-query search fees have no GoModel +equivalent and are not read. + +Eden's reasoning and audio rates (`output_cost_per_reasoning_token`, +`input_cost_per_audio_token`) are also left unread. GoModel would price those +token types by subtracting the base rate, which assumes the counts are already +part of the base totals — and Eden's usage object reports only +`prompt_tokens`, `completion_tokens`, and `total_tokens`, so there is no +breakdown to confirm that against. Those token types fall back to the base +input/output rates instead. Pricing is read live from Eden — nothing is hard-coded, and Eden models do not need to be present in GoModel's central model catalog. This makes `EDENAI_MODEL_FILTER_MAX_PRICE_PER_MTOK`, cost-based load balancing, and price display work for Eden models. -A rate Eden does not publish is left unset rather than treated as free, so a -partially priced model is never silently undercounted. +This discovered pricing is what model metadata shows, and what cost accounting +falls back to when a response carries no exact charge. The exact per-request +cost Eden returns takes precedence whenever it is available (see below). ## Request cost Eden returns the exact USD charge for each request as a top-level `cost` -member, and GoModel records that figure as the request's cost instead of -recomputing it from token counts. Eden reprices its upstreams automatically -and applies account-level discounts, so its own number is authoritative in a -way a rate-card reconstruction is not. +member — on chat completions and on embeddings alike — and GoModel +records that figure as the request's cost instead of recomputing it from token +counts. Eden reprices its upstreams automatically and applies account-level +discounts, so its own number is authoritative in a way a rate-card +reconstruction is not. This makes Eden spend visible to usage records, budgets, cost dashboards, and observability, and the usage entry is labelled with the cost source @@ -151,6 +189,12 @@ Eden's `/embeddings` route is OpenAI-compatible and uses the same { "model": "openai/text-embedding-3-small", "input": "hello" } ``` +Embeddings responses carry the same Eden extensions as chat completions, so the +exact `cost` Eden reports is recorded for them too. Eden's LLM catalog +(`GET /v3/models`) does not list embedding models, so embedding IDs are +forwarded without discovered metadata; pin them under `models:` if you want +them advertised on `/v1/models`. + ## Passthrough `edenai` is in the default `ENABLED_PASSTHROUGH_PROVIDERS` allowlist, so diff --git a/internal/core/types.go b/internal/core/types.go index ee91e9870..357c16707 100644 --- a/internal/core/types.go +++ b/internal/core/types.go @@ -516,7 +516,15 @@ type EmbeddingData struct { } // EmbeddingUsage represents token usage information for embeddings. +// +// RawUsage carries provider-reported usage members that have no typed field, +// mirroring Usage.RawUsage on the chat surface. usage.ExtractFromEmbeddingResponse +// forwards it to the cost pipeline, which is what lets a provider that returns +// an exact per-request charge (Eden AI reports one as a root-level "cost") +// have that figure recorded instead of a rate-card reconstruction. Providers +// that report nothing extra leave it nil. type EmbeddingUsage struct { - PromptTokens int `json:"prompt_tokens"` - TotalTokens int `json:"total_tokens"` + PromptTokens int `json:"prompt_tokens"` + TotalTokens int `json:"total_tokens"` + RawUsage map[string]any `json:"raw_usage,omitempty"` } diff --git a/internal/providers/edenai/capabilities_test.go b/internal/providers/edenai/capabilities_test.go new file mode 100644 index 000000000..a96311d4f --- /dev/null +++ b/internal/providers/edenai/capabilities_test.go @@ -0,0 +1,212 @@ +package edenai + +import ( + "testing" +) + +// TestCapabilities_LiveSchemaOnly pins the one capability schema Eden actually +// publishes, checked against the live /v3/models catalog: a `capabilities` +// object holding `input_modalities`, `output_modalities`, and a growing set of +// `supports_*` booleans. +// +// Eden does not publish bare capability flags such as "pdf", "tool_calling", +// or "web_search" alongside those — the equivalents are supports_pdf_input, +// supports_tools/supports_function_calling, and supports_web_search — so this +// provider deliberately reads only the prefixed form. Accepting unprefixed +// names as well would mean advertising capabilities on the strength of a +// schema Eden does not serve. +func TestCapabilities_LiveSchemaOnly(t *testing.T) { + model := firstModel(t, `{"object":"list","data":[{ + "id":"openai/gpt-4","object":"model", + "capabilities": { + "input_modalities": ["text", "image"], + "output_modalities": ["text"], + "supports_reasoning": true, + "supports_web_search": true, + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_tool_choice": false, + "supports_computer_use": false + } + }]}`) + + capabilities := model.Metadata.Capabilities + // The supports_ prefix is stripped so the names read the way other + // providers report them. + for _, name := range []string{"reasoning", "web_search", "function_calling", "prompt_caching"} { + if !capabilities[name] { + t.Errorf("capability %q = false, want true", name) + } + } + // An image input modality is reported as vision, the gateway's name for it. + if !capabilities["vision"] { + t.Error(`capability "vision" = false, want true for an image input modality`) + } + // A flag Eden reports as false must not be advertised at all. + for _, name := range []string{"tool_choice", "computer_use"} { + if _, present := capabilities[name]; present { + t.Errorf("capability %q is present, want absent: Eden reported it false", name) + } + } + // The prefixed keys must not leak through under their raw Eden names. + for _, name := range []string{"supports_reasoning", "supports_web_search", "input_modalities", "output_modalities"} { + if _, present := capabilities[name]; present { + t.Errorf("capability %q is present, want the Eden key name not to leak", name) + } + } +} + +// TestCapabilities_RareSupportsFlagsFlowThrough asserts the long tail of +// supports_* flags Eden publishes on only a handful of models is picked up +// without a source change here, and lands on the gateway's own capability +// names where one already exists. +// +// supports_vision and supports_tools matter most: "vision" and "tools" are +// names other providers already report, so stripping the prefix is what keeps +// Eden's vocabulary aligned rather than parallel. +func TestCapabilities_RareSupportsFlagsFlowThrough(t *testing.T) { + model := firstModel(t, `{"object":"list","data":[{ + "id":"openai/gpt-4","object":"model", + "capabilities": { + "output_modalities": ["text"], + "supports_vision": true, + "supports_tools": true, + "supports_pdf_input": true, + "supports_audio_input": true, + "supports_structured_output": true, + "supports_responses_api": true + } + }]}`) + + for _, name := range []string{ + "vision", + "tools", + "pdf_input", + "audio_input", + "structured_output", + "responses_api", + } { + if !model.Metadata.Capabilities[name] { + t.Errorf("capability %q = false, want true", name) + } + } +} + +// TestCapabilities_ObjectValuedReasoningIsNotAFlag guards a real shape in the +// live catalog: most Eden entries carry a `reasoning` member that is an object +// describing reasoning options, sitting next to the separate +// `supports_reasoning` boolean. +// +// It must not become a capability. It has no supports_ prefix and is not a +// bool, so both screens in capabilities() have to hold — otherwise every model +// with reasoning metadata would claim a "reasoning" capability regardless of +// what supports_reasoning actually says. +func TestCapabilities_ObjectValuedReasoningIsNotAFlag(t *testing.T) { + model := firstModel(t, `{"object":"list","data":[{ + "id":"openai/gpt-4","object":"model", + "capabilities": { + "output_modalities": ["text"], + "supports_reasoning": false, + "reasoning": {"effort": ["low", "medium", "high"]}, + "supports_web_search": true + } + }]}`) + + capabilities := model.Metadata.Capabilities + if _, present := capabilities["reasoning"]; present { + t.Error(`capability "reasoning" is present, but Eden reported supports_reasoning: false; the object-valued "reasoning" member must not be read as a flag`) + } + if !capabilities["web_search"] { + t.Error(`capability "web_search" = false, want true: a sibling object member must not stop the real flags being read`) + } +} + +// TestCapabilities_NonBooleanSupportsValuesIgnored asserts a supports_* member +// Eden ever publishes as something other than a boolean is skipped rather than +// coerced into a true. +func TestCapabilities_NonBooleanSupportsValuesIgnored(t *testing.T) { + model := firstModel(t, `{"object":"list","data":[{ + "id":"openai/gpt-4","object":"model", + "capabilities": { + "output_modalities": ["text"], + "supports_web_search": "yes", + "supports_tool_choice": 1, + "supports_prompt_caching": null, + "supports_reasoning": true + } + }]}`) + + capabilities := model.Metadata.Capabilities + for _, name := range []string{"web_search", "tool_choice", "prompt_caching"} { + if _, present := capabilities[name]; present { + t.Errorf("capability %q is present, want absent: Eden did not report a boolean", name) + } + } + if !capabilities["reasoning"] { + t.Error(`capability "reasoning" = false, want true`) + } +} + +// TestCapabilities_BareSupportsPrefixIgnored asserts a key that is exactly the +// prefix contributes no empty-named capability. +func TestCapabilities_BareSupportsPrefixIgnored(t *testing.T) { + model := firstModel(t, `{"object":"list","data":[{ + "id":"openai/gpt-4","object":"model", + "capabilities": {"output_modalities": ["text"], "supports_": true} + }]}`) + + if _, present := model.Metadata.Capabilities[""]; present { + t.Error(`an empty-named capability was recorded for the bare "supports_" key`) + } +} + +// TestCapabilities_InputModalityMapping pins the input modalities the live +// catalog actually contains (text, image, file, video, audio) onto the +// gateway's capability names. "file" has no gateway capability and is skipped +// rather than guessed at. +func TestCapabilities_InputModalityMapping(t *testing.T) { + model := firstModel(t, `{"object":"list","data":[{ + "id":"google/gemini-3.8-flash","object":"model", + "capabilities": { + "input_modalities": ["text", "image", "video", "file", "audio"], + "output_modalities": ["text"] + } + }]}`) + + capabilities := model.Metadata.Capabilities + for _, name := range []string{"vision", "video", "audio"} { + if !capabilities[name] { + t.Errorf("capability %q = false, want true", name) + } + } + for _, name := range []string{"file", "text"} { + if _, present := capabilities[name]; present { + t.Errorf("capability %q is present; modalities with no gateway capability must be skipped", name) + } + } +} + +// TestCapabilities_VideoInputIsNotVideoOutput is the distinction finding 7 +// turned on. In the live catalog `video` appears only as an *input* modality +// (143 of 1049 models) and never as an output one, so a video-understanding +// model is an ordinary chat model that happens to accept video. +// +// It must keep its chat modes and stay advertised; only a model whose *output* +// is video is unservable here. +func TestCapabilities_VideoInputIsNotVideoOutput(t *testing.T) { + model := firstModel(t, `{"object":"list","data":[{ + "id":"qwen/qwen3.8-max-0902","object":"model", + "capabilities": { + "input_modalities": ["text", "image", "video"], + "output_modalities": ["text"] + } + }]}`) + + if !model.Metadata.Capabilities["video"] { + t.Error(`capability "video" = false, want true for a video input modality`) + } + modes := model.Metadata.Modes + if len(modes) != 2 || modes[0] != "chat" || modes[1] != "responses" { + t.Errorf("modes = %v, want [chat responses]: video input must not produce a video mode", modes) + } +} diff --git a/internal/providers/edenai/edenai.go b/internal/providers/edenai/edenai.go index 610422e48..81655c4c0 100644 --- a/internal/providers/edenai/edenai.go +++ b/internal/providers/edenai/edenai.go @@ -62,29 +62,54 @@ var ( _ core.PassthroughProvider = (*Provider)(nil) ) -// New creates a new Eden AI provider. +// New creates a new Eden AI provider. NewCompatibleProvider takes its +// transport from CompatibleProviderConfig.HTTPClient, so the guarded client +// compatibleConfig installs is the one that reaches the network. func New(cfg providers.ProviderConfig, opts providers.ProviderOptions) core.Provider { return &Provider{compat: openai.NewCompatibleProvider(cfg.APIKey, opts, compatibleConfig( providers.ResolveBaseURL(cfg.BaseURL, defaultBaseURL), + nil, ))} } // NewWithHTTPClient creates a new Eden AI provider with a custom HTTP client. // If httpClient is nil, http.DefaultClient is used. // +// The nil default is applied here rather than left to +// NewCompatibleProviderWithHTTPClient so it stays the same default every other +// chat-compatible provider gets from that helper. Eden always hands it a +// non-nil client, so the helper's own nil branch is never reached. +// +// Either way the client carries Eden's redirect guard (see guardedHTTPClient), +// installed on a copy so a shared client -- http.DefaultClient above all -- is +// never modified. +// +// Unlike NewCompatibleProvider, NewCompatibleProviderWithHTTPClient takes its +// transport from the positional argument and ignores +// CompatibleProviderConfig.HTTPClient entirely, so the guarded client has to be +// handed over there as well. Passing the raw caller client here would leave +// this construction path — and only this one — following credential-leaking +// redirects. +// // The signature matches every other chat-compatible provider on main: // (apiKey, baseURL, httpClient, hooks). func NewWithHTTPClient(apiKey string, baseURL string, httpClient *http.Client, hooks llmclient.Hooks) *Provider { - return &Provider{compat: openai.NewCompatibleProviderWithHTTPClient(apiKey, httpClient, hooks, compatibleConfig( - providers.ResolveBaseURL(baseURL, defaultBaseURL), - ))} + if httpClient == nil { + httpClient = http.DefaultClient + } + cfg := compatibleConfig(providers.ResolveBaseURL(baseURL, defaultBaseURL), httpClient) + return &Provider{compat: openai.NewCompatibleProviderWithHTTPClient(apiKey, cfg.HTTPClient, hooks, cfg)} } -// compatibleConfig returns the shared OpenAI-compatible transport settings for Eden AI. -func compatibleConfig(baseURL string) openai.CompatibleProviderConfig { +// compatibleConfig returns the shared OpenAI-compatible transport settings for +// Eden AI. httpClient is the caller-supplied client, or nil to build the +// gateway default; either way it is wrapped by guardedHTTPClient, so both +// constructors get the same redirect policy from one place. +func compatibleConfig(baseURL string, httpClient *http.Client) openai.CompatibleProviderConfig { return openai.CompatibleProviderConfig{ ProviderName: providerType, BaseURL: baseURL, + HTTPClient: guardedHTTPClient(httpClient), SetHeaders: setHeaders, } } @@ -92,7 +117,17 @@ func compatibleConfig(baseURL string) openai.CompatibleProviderConfig { // setHeaders applies Eden AI's bearer-token authentication. CompatibleProvider // sends no credential when SetHeaders is nil (unlike ChatCompatible, which // defaults to bearer), so this must stay wired up. +// +// The credential is withheld from a destination that would carry it in +// cleartext. Eden's base URL is operator-supplied, so an http:// override — +// whether set by mistake or by a downgrade attempt — would otherwise put the +// gateway's Eden key on the wire in plain text. Loopback is exempt, which is +// what keeps local proxies and this package's httptest servers working. A +// withheld credential yields an Eden 401 rather than a leaked key. func setHeaders(req *http.Request, apiKey string) { + if !credentialSafeURL(req.URL) { + return + } providers.SetAuthHeaders(req, apiKey, providers.AuthHeaderConfig{AuthScheme: "Bearer "}) } @@ -146,13 +181,32 @@ func (p *Provider) StreamResponses(ctx context.Context, req *core.ResponsesReque // Embeddings sends an embeddings request to Eden AI. Eden's /embeddings route // is OpenAI-compatible and takes the same provider/model IDs. +// +// The request is issued through the raw transport rather than +// CompatibleProvider.Embeddings because Eden annotates the embeddings envelope +// with the same two extensions it puts on chat completions — a root-level +// "cost" and the upstream "provider" — and core.EmbeddingResponse models no +// unknown-field container, so decoding straight into it would discard the +// exact charge before anything could read it. Everything else matches what the +// shared helper does: same endpoint, same operation label, and the same +// EnsureModel backfill for a response that omits the model. func (p *Provider) Embeddings(ctx context.Context, req *core.EmbeddingRequest) (*core.EmbeddingResponse, error) { - resp, err := p.compat.Embeddings(ctx, req) - if err != nil { + if req == nil { + return nil, core.NewInvalidRequestError("embedding request is required", nil) + } + var resp embeddingResponse + if err := p.compat.Do(ctx, llmclient.Request{ + Method: http.MethodPost, + Endpoint: "/embeddings", + Operation: llmclient.OperationEmbeddings, + Model: req.Model, + Body: req, + }, &resp); err != nil { return nil, err } - normalizeEmbeddingResponse(resp) - return resp, nil + core.EnsureModel(&resp.Model, req.Model) + normalizeEmbeddingResponse(&resp) + return &resp.EmbeddingResponse, nil } // Passthrough forwards an opaque request to Eden AI. diff --git a/internal/providers/edenai/embeddings_cost_test.go b/internal/providers/edenai/embeddings_cost_test.go new file mode 100644 index 000000000..fb350312b --- /dev/null +++ b/internal/providers/edenai/embeddings_cost_test.go @@ -0,0 +1,271 @@ +package edenai + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + + "github.com/enterpilot/gomodel/internal/core" + "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/usage" +) + +// edenEmbeddingBody is the shape Eden documents for /v3/embeddings: the +// OpenAI-compatible envelope plus the two members Eden adds at the root, the +// upstream "provider" and the exact per-request "cost" in USD. +const edenEmbeddingBody = `{ + "object": "list", + "model": "openai/text-embedding-3-small", + "provider": "openai", + "data": [{"object": "embedding", "index": 0, "embedding": [0.0123, -0.0456]}], + "usage": {"prompt_tokens": 9, "total_tokens": 9}, + "cost": 0.0000012 +}` + +// embeddingsProvider serves one /embeddings payload. +func embeddingsProvider(t *testing.T, payload string) *Provider { + t.Helper() + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(payload)) + })) + t.Cleanup(server.Close) + return NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) +} + +// TestEmbeddings_LiftsRootLevelCostIntoUsage asserts Eden's root-level +// embeddings cost survives decoding. +// +// core.EmbeddingResponse models no unknown-field container, so a root member +// is dropped outright unless the provider decodes the envelope itself. Landing +// the value in Usage.RawUsage is what puts it on the same path the chat +// surface already uses. +func TestEmbeddings_LiftsRootLevelCostIntoUsage(t *testing.T) { + provider := embeddingsProvider(t, edenEmbeddingBody) + + resp, err := provider.Embeddings(context.Background(), embeddingRequest()) + if err != nil { + t.Fatalf("Embeddings() error = %v", err) + } + + cost, ok := resp.Usage.RawUsage["cost"] + if !ok { + t.Fatalf("Usage.RawUsage = %v, want Eden's root-level cost lifted into it", resp.Usage.RawUsage) + } + if cost != 0.0000012 { + t.Errorf("RawUsage[cost] = %v, want 0.0000012 exactly", cost) + } + // The rest of the envelope must still decode normally. + if resp.Usage.PromptTokens != 9 || resp.Usage.TotalTokens != 9 { + t.Errorf("usage tokens = %d/%d, want 9/9", resp.Usage.PromptTokens, resp.Usage.TotalTokens) + } + if len(resp.Data) != 1 || len(resp.Data[0].Embedding) == 0 { + t.Errorf("data = %+v, want one embedding", resp.Data) + } + if resp.Model != "openai/text-embedding-3-small" { + t.Errorf("model = %q, want the Eden model ID forwarded verbatim", resp.Model) + } + // Eden's upstream name must not be reported as the executing provider. + if resp.Provider != "" { + t.Errorf("Provider = %q, want empty so the gateway reports edenai", resp.Provider) + } +} + +// TestEmbeddings_ExactCostReachesRecordedTotal is the end-to-end assertion for +// the whole chain: Eden's root-level cost, through the provider, through usage +// extraction, to the cost recorded against the request. +// +// Catalog pricing is supplied deliberately and is wrong on purpose. If the +// exact figure were lost anywhere along the way, the entry would silently fall +// back to these token rates and the recorded total would be a rate-card +// reconstruction rather than what Eden actually charged. +func TestEmbeddings_ExactCostReachesRecordedTotal(t *testing.T) { + provider := embeddingsProvider(t, edenEmbeddingBody) + + resp, err := provider.Embeddings(context.Background(), embeddingRequest()) + if err != nil { + t.Fatalf("Embeddings() error = %v", err) + } + + wrongRate := 100.0 // $100/MTok would price 9 tokens at $0.0009, not $0.0000012. + pricing := &core.ModelPricing{Currency: "USD", InputPerMtok: &wrongRate} + + entry := usage.ExtractFromEmbeddingResponse(resp, "req-1", providerType, "/v1/embeddings", pricing) + if entry == nil { + t.Fatal("ExtractFromEmbeddingResponse returned nil") + } + if entry.TotalCost == nil { + t.Fatal("TotalCost = nil, want Eden's exact charge recorded") + } + if *entry.TotalCost != 0.0000012 { + t.Errorf("TotalCost = %v, want 0.0000012 exactly (Eden's reported charge, not a token-rate estimate)", *entry.TotalCost) + } + if entry.CostSource != usage.CostSourceEdenAICost { + t.Errorf("CostSource = %q, want %q", entry.CostSource, usage.CostSourceEdenAICost) + } + if entry.CostsCalculationCaveat != "" { + t.Errorf("CostsCalculationCaveat = %q, want empty for an exact provider-reported cost", entry.CostsCalculationCaveat) + } +} + +// TestEmbeddings_FallsBackToTokenPricingWithoutCost asserts the fallback still +// works: an Eden embeddings response carrying no cost is priced from the +// discovered per-model rates, exactly as before. +func TestEmbeddings_FallsBackToTokenPricingWithoutCost(t *testing.T) { + provider := embeddingsProvider(t, `{ + "object": "list", + "model": "openai/text-embedding-3-small", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1]}], + "usage": {"prompt_tokens": 1000, "total_tokens": 1000} + }`) + + resp, err := provider.Embeddings(context.Background(), embeddingRequest()) + if err != nil { + t.Fatalf("Embeddings() error = %v", err) + } + if len(resp.Usage.RawUsage) != 0 { + t.Errorf("Usage.RawUsage = %v, want empty when Eden reports no cost", resp.Usage.RawUsage) + } + + rate := 0.02 // $0.02/MTok * 1000 tokens = $0.00002 + pricing := &core.ModelPricing{Currency: "USD", InputPerMtok: &rate} + + entry := usage.ExtractFromEmbeddingResponse(resp, "req-2", providerType, "/v1/embeddings", pricing) + if entry == nil { + t.Fatal("ExtractFromEmbeddingResponse returned nil") + } + if entry.TotalCost == nil { + t.Fatal("TotalCost = nil, want the token-rate fallback to apply") + } + if *entry.TotalCost != 0.00002 { + t.Errorf("TotalCost = %v, want 0.00002 from the catalog rate", *entry.TotalCost) + } + if entry.CostSource != usage.CostSourceModelPricing { + t.Errorf("CostSource = %q, want %q", entry.CostSource, usage.CostSourceModelPricing) + } +} + +// TestEmbeddings_RejectsUnusableCost asserts a cost member that would corrupt +// accounting is ignored rather than recorded, leaving the token-rate fallback +// in charge. +// +// The null case is the one that matters most: unmarshalling a JSON null into a +// float64 succeeds without error, so an unscreened null would be recorded as a +// real $0.00 charge and understate spend. +func TestEmbeddings_RejectsUnusableCost(t *testing.T) { + tests := []struct { + name string + cost string + }{ + {"null", `null`}, + {"negative", `-0.5`}, + {"string", `"0.0000012"`}, + {"object", `{"amount": 0.1}`}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + provider := embeddingsProvider(t, `{ + "object": "list", + "model": "openai/text-embedding-3-small", + "data": [], + "usage": {"prompt_tokens": 1000, "total_tokens": 1000}, + "cost": `+tc.cost+` + }`) + + resp, err := provider.Embeddings(context.Background(), embeddingRequest()) + if err != nil { + t.Fatalf("Embeddings() error = %v", err) + } + if _, present := resp.Usage.RawUsage["cost"]; present { + t.Fatalf("RawUsage[cost] = %v, want absent for an unusable cost member", resp.Usage.RawUsage["cost"]) + } + + rate := 0.02 + entry := usage.ExtractFromEmbeddingResponse(resp, "req", providerType, "/v1/embeddings", + &core.ModelPricing{Currency: "USD", InputPerMtok: &rate}) + if entry.CostSource != usage.CostSourceModelPricing { + t.Errorf("CostSource = %q, want the token-rate fallback %q", entry.CostSource, usage.CostSourceModelPricing) + } + }) + } +} + +// TestEmbeddings_ZeroCostIsRecordedAsFree asserts an explicit zero is a real +// price, not a missing one: Eden reporting no charge must be recorded as $0 +// rather than silently repriced from the rate card. +func TestEmbeddings_ZeroCostIsRecordedAsFree(t *testing.T) { + provider := embeddingsProvider(t, `{ + "object": "list", + "model": "openai/text-embedding-3-small", + "data": [], + "usage": {"prompt_tokens": 1000, "total_tokens": 1000}, + "cost": 0 + }`) + + resp, err := provider.Embeddings(context.Background(), embeddingRequest()) + if err != nil { + t.Fatalf("Embeddings() error = %v", err) + } + + rate := 0.02 + entry := usage.ExtractFromEmbeddingResponse(resp, "req", providerType, "/v1/embeddings", + &core.ModelPricing{Currency: "USD", InputPerMtok: &rate}) + if entry.TotalCost == nil || *entry.TotalCost != 0 { + t.Fatalf("TotalCost = %v, want 0 recorded from Eden's explicit zero", entry.TotalCost) + } + if entry.CostSource != usage.CostSourceEdenAICost { + t.Errorf("CostSource = %q, want %q", entry.CostSource, usage.CostSourceEdenAICost) + } +} + +// TestEmbeddings_UsageLevelCostWins asserts the conventional location stays +// authoritative. If Eden ever also reports cost inside usage, that is the more +// specific reading and the root-level value must not overwrite it. +func TestEmbeddings_UsageLevelCostWins(t *testing.T) { + provider := embeddingsProvider(t, `{ + "object": "list", + "model": "openai/text-embedding-3-small", + "data": [], + "usage": {"prompt_tokens": 9, "total_tokens": 9, "raw_usage": {"cost": 0.5}}, + "cost": 0.0000012 + }`) + + resp, err := provider.Embeddings(context.Background(), embeddingRequest()) + if err != nil { + t.Fatalf("Embeddings() error = %v", err) + } + if got := resp.Usage.RawUsage["cost"]; got != 0.5 { + t.Errorf("RawUsage[cost] = %v, want the pre-existing usage-level 0.5 preserved", got) + } +} + +// TestEmbeddings_NilRequestRejected asserts the guard the shared helper used to +// provide is still in place now that the provider issues the request itself. +func TestEmbeddings_NilRequestRejected(t *testing.T) { + provider := embeddingsProvider(t, edenEmbeddingBody) + + if _, err := provider.Embeddings(context.Background(), nil); err == nil { + t.Fatal("Embeddings(nil) = nil error, want a rejection") + } +} + +// TestEmbeddings_BackfillsModelFromRequest asserts the EnsureModel behavior the +// shared helper applied is preserved: a response that omits the model is +// labelled with the requested one, so usage rows are attributed correctly. +func TestEmbeddings_BackfillsModelFromRequest(t *testing.T) { + provider := embeddingsProvider(t, `{ + "object": "list", + "data": [], + "usage": {"prompt_tokens": 1, "total_tokens": 1} + }`) + + resp, err := provider.Embeddings(context.Background(), embeddingRequest()) + if err != nil { + t.Fatalf("Embeddings() error = %v", err) + } + if resp.Model != slashedModel { + t.Errorf("Model = %q, want the requested %q backfilled", resp.Model, slashedModel) + } +} diff --git a/internal/providers/edenai/models.go b/internal/providers/edenai/models.go index d2a556812..e06c2e314 100644 --- a/internal/providers/edenai/models.go +++ b/internal/providers/edenai/models.go @@ -35,24 +35,100 @@ type modelInfo struct { ContextLength int `json:"context_length"` } -// modelPricing holds one of Eden's per-token USD rate blocks. +// modelPricing holds one of Eden's per-token USD rate blocks. Only the rates +// whose unit and meaning both map onto a core.ModelPricing field are decoded. +// +// Left unread on purpose: +// - Context-length-tiered variants (input_cost_per_token_above_200k_tokens), +// the tiered_pricing list, and per-query search fees (an object, not a +// scalar). The gateway's pricing model has no equivalent, so reading them +// into a flat per-Mtok field would misprice the model. +// - output_cost_per_reasoning_token and input_cost_per_audio_token. The rates +// exist in Eden's catalog, but the gateway would apply them through usage's +// OpenAI-compatible mappings, which treat reasoning and audio counts as +// already included in the base completion/prompt totals and subtract the +// base rate accordingly. Eden's documented usage object carries only +// prompt_tokens, completion_tokens, and total_tokens -- no reasoning or +// audio breakdown -- so there is nothing to confirm that assumption +// against. Mapping them would stake cost math on an unverified split for +// token counts Eden does not currently report. type modelPricing struct { - InputCostPerToken *float64 `json:"input_cost_per_token"` - OutputCostPerToken *float64 `json:"output_cost_per_token"` - CacheReadInputTokenCost *float64 `json:"cache_read_input_token_cost"` + InputCostPerToken *float64 `json:"input_cost_per_token"` + OutputCostPerToken *float64 `json:"output_cost_per_token"` + CacheReadInputTokenCost *float64 `json:"cache_read_input_token_cost"` + CacheCreationInputTokenCost *float64 `json:"cache_creation_input_token_cost"` } -// effectivePricing picks the rate card to publish. Eden's `pricing` is what -// the account is actually charged (any discount already applied), so it wins -// whenever it carries a usable rate. `list_pricing` is the undiscounted card -// and is used only as a fallback: an approximate rate still lets price -// filters and the cost load-balancing strategy rank the model, whereas no -// rate at all drops it from both. +// pricingRate pairs one Eden per-token rate with the core.ModelPricing field +// it feeds, so effectivePricing can resolve every rate the same way instead of +// repeating the fallback per field. The shape follows usage.tokenCostMapping, +// which already expresses the same read-field/write-field pairing. +var pricingRates = []struct { + rate func(*modelPricing) *float64 + assign func(*core.ModelPricing, float64) +}{ + { + func(p *modelPricing) *float64 { return p.InputCostPerToken }, + func(c *core.ModelPricing, v float64) { c.InputPerMtok = &v }, + }, + { + func(p *modelPricing) *float64 { return p.OutputCostPerToken }, + func(c *core.ModelPricing, v float64) { c.OutputPerMtok = &v }, + }, + { + func(p *modelPricing) *float64 { return p.CacheReadInputTokenCost }, + func(c *core.ModelPricing, v float64) { c.CachedInputPerMtok = &v }, + }, + { + func(p *modelPricing) *float64 { return p.CacheCreationInputTokenCost }, + func(c *core.ModelPricing, v float64) { c.CacheWritePerMtok = &v }, + }, +} + +// effectivePricing builds the rate card to publish, resolving each rate +// independently. +// +// Eden's `pricing` is what the account is actually charged — the undiscounted +// `list_pricing` with the account discount already applied — so a usable +// account rate always wins, including an explicit 0 for a genuinely free +// model. `list_pricing` fills in only the individual rates `pricing` leaves +// unusable. The fallback is per field rather than per block so a block that +// prices some token types and not others keeps its account rates instead of +// losing the rates it does not carry. +// +// Eden currently publishes the same key set in both blocks, so in practice +// every rate resolves from `pricing`; the per-field path is what keeps that +// from being load-bearing. A rate perMtok rejects as unusable (negative, +// non-finite, overflowing) is treated the same as an absent one and may fall +// back, since an approximate rate still lets price filters and the cost +// load-balancing strategy rank the model, whereas no rate at all drops it +// from both. func (m modelInfo) effectivePricing() *core.ModelPricing { - if pricing := m.Pricing.toCore(); pricing != nil { - return pricing + pricing := &core.ModelPricing{Currency: "USD"} + priced := false + for _, mapping := range pricingRates { + rate, ok := perMtok(rateFrom(m.Pricing, mapping.rate)) + if !ok { + rate, ok = perMtok(rateFrom(m.ListPricing, mapping.rate)) + } + if !ok { + continue + } + mapping.assign(pricing, rate) + priced = true + } + if !priced { + return nil + } + return pricing +} + +// rateFrom reads one rate out of a rate block Eden may have omitted entirely. +func rateFrom(block *modelPricing, field func(*modelPricing) *float64) *float64 { + if block == nil { + return nil } - return m.ListPricing.toCore() + return field(block) } // ListModels returns Eden's live catalog, retaining the context window, @@ -211,41 +287,17 @@ func stringSlice(value any) []string { return result } -// toCore converts one of Eden's per-token USD rate blocks into the gateway's -// per-million-token pricing. Eden names the unit in each field -// (input_cost_per_token), so the ×1e6 scaling is read off the contract rather -// than assumed. A rate Eden omits stays absent: costing a token type at zero -// because no price was published would understate spend, so only a rate Eden -// explicitly reports as 0 prices at zero. -func (p *modelPricing) toCore() *core.ModelPricing { - if p == nil { - return nil - } - input, hasInput := perMtok(p.InputCostPerToken) - output, hasOutput := perMtok(p.OutputCostPerToken) - cachedInput, hasCachedInput := perMtok(p.CacheReadInputTokenCost) - if !hasInput && !hasOutput && !hasCachedInput { - return nil - } - - pricing := &core.ModelPricing{Currency: "USD"} - if hasInput { - pricing.InputPerMtok = &input - } - if hasOutput { - pricing.OutputPerMtok = &output - } - if hasCachedInput { - pricing.CachedInputPerMtok = &cachedInput - } - return pricing -} - -// perMtok scales one per-token USD rate to per million tokens. Rates that are -// absent, negative, or non-finite report no price rather than a wrong one, and -// so does a rate large enough that scaling overflows to infinity: a corrupt -// number here would propagate into every price comparison, budget total, and -// cost-strategy decision downstream. +// perMtok scales one per-token USD rate to per million tokens. Eden names the +// unit in each field (input_cost_per_token), so the x1e6 scaling is read off +// the contract rather than assumed. +// +// Rates that are absent, negative, or non-finite report no price rather than a +// wrong one, and so does a rate large enough that scaling overflows to +// infinity: a corrupt number here would propagate into every price comparison, +// budget total, and cost-strategy decision downstream. A rate Eden explicitly +// reports as 0 is a real price and is kept -- costing a token type at zero +// because no price was published would understate spend, but a published zero +// means free. func perMtok(perToken *float64) (float64, bool) { if perToken == nil { return 0, false diff --git a/internal/providers/edenai/models_test.go b/internal/providers/edenai/models_test.go index 10c958739..087095cd0 100644 --- a/internal/providers/edenai/models_test.go +++ b/internal/providers/edenai/models_test.go @@ -182,23 +182,181 @@ func TestListModels_FallsBackToListPricing(t *testing.T) { } // TestListModels_ListPricingDoesNotMaskApplicablePricing asserts the fallback -// never overrides a usable applicable rate, including a partial one: a model -// priced only for input keeps that rate rather than swapping in the full list -// card. +// never overrides a usable account rate: a model priced for input keeps its own +// input rate even though the list card also carries one. func TestListModels_ListPricingDoesNotMaskApplicablePricing(t *testing.T) { + model := firstModel(t, `{"object":"list","data":[{ + "id":"openai/gpt-4","object":"model", + "pricing": {"input_cost_per_token": 6e-8, "output_cost_per_token": 1.8e-7}, + "list_pricing": {"input_cost_per_token": 9e-8, "output_cost_per_token": 2.8e-7} + }]}`) + + assertPrice(t, "InputPerMtok", model.Metadata.Pricing.InputPerMtok, 0.06) + assertPrice(t, "OutputPerMtok", model.Metadata.Pricing.OutputPerMtok, 0.18) +} + +// TestListModels_ListPricingFillsOnlyMissingRates asserts the fallback is per +// rate, not per block: an account block that prices input but not output keeps +// its own input rate and takes only the missing output rate from the list card. +// +// Leaving the gap unpriced is not the safer option it looks like. A nil +// OutputPerMtok makes usage.CalculateGranularCost skip the output side +// entirely, so output tokens are billed at $0 and the recorded total is +// silently short, with no caveat attached. The undiscounted list rate +// overstates the output somewhat, which is the same trade the whole-block +// fallback already accepts (see TestListModels_FallsBackToListPricing), and is +// far closer to the real charge than zero. +func TestListModels_ListPricingFillsOnlyMissingRates(t *testing.T) { model := firstModel(t, `{"object":"list","data":[{ "id":"openai/gpt-4","object":"model", "pricing": {"input_cost_per_token": 6e-8}, "list_pricing": {"input_cost_per_token": 9e-8, "output_cost_per_token": 2.8e-7} }]}`) + assertPrice(t, "InputPerMtok", model.Metadata.Pricing.InputPerMtok, 0.06) + assertPrice(t, "OutputPerMtok", model.Metadata.Pricing.OutputPerMtok, 0.28) +} + +// TestListModels_ListPricingFillsMissingInputRate is the mirror case: an +// account block that prices only output keeps that rate and fills input from +// the list card. +func TestListModels_ListPricingFillsMissingInputRate(t *testing.T) { + model := firstModel(t, `{"object":"list","data":[{ + "id":"openai/gpt-4","object":"model", + "pricing": {"output_cost_per_token": 1.8e-7}, + "list_pricing": {"input_cost_per_token": 9e-8, "output_cost_per_token": 2.8e-7} + }]}`) + + assertPrice(t, "InputPerMtok", model.Metadata.Pricing.InputPerMtok, 0.09) + assertPrice(t, "OutputPerMtok", model.Metadata.Pricing.OutputPerMtok, 0.18) +} + +// TestListModels_ListPricingFillsSeveralMissingRates covers an account block +// missing more than one rate, and confirms the rates it does carry -- including +// an explicit zero -- are still preferred field by field. +func TestListModels_ListPricingFillsSeveralMissingRates(t *testing.T) { + model := firstModel(t, `{"object":"list","data":[{ + "id":"openai/gpt-4","object":"model", + "pricing": {"input_cost_per_token": 0}, + "list_pricing": { + "input_cost_per_token": 9e-8, + "output_cost_per_token": 2.8e-7, + "cache_read_input_token_cost": 1.2e-8, + "cache_creation_input_token_cost": 1.5e-7 + } + }]}`) + + pricing := model.Metadata.Pricing + // An explicit account zero is a real price and must not be filled in. + assertPrice(t, "InputPerMtok", pricing.InputPerMtok, 0) + assertPrice(t, "OutputPerMtok", pricing.OutputPerMtok, 0.28) + assertPrice(t, "CachedInputPerMtok", pricing.CachedInputPerMtok, 0.012) + assertPrice(t, "CacheWritePerMtok", pricing.CacheWritePerMtok, 0.15) +} + +// TestListModels_UnusableAccountRateFallsBackToList asserts a rate the account +// block publishes but perMtok rejects (negative, non-finite, overflowing) is +// treated like an absent one, so the list card can still supply a usable number +// instead of the model losing that rate entirely. +func TestListModels_UnusableAccountRateFallsBackToList(t *testing.T) { + for _, tc := range []struct { + name string + rate string + }{ + {"negative", "-1e-8"}, + {"overflowing", "1e308"}, + } { + t.Run(tc.name, func(t *testing.T) { + model := firstModel(t, `{"object":"list","data":[{ + "id":"openai/gpt-4","object":"model", + "pricing": {"input_cost_per_token": `+tc.rate+`, "output_cost_per_token": 1.8e-7}, + "list_pricing": {"input_cost_per_token": 9e-8, "output_cost_per_token": 2.8e-7} + }]}`) + + assertPrice(t, "InputPerMtok", model.Metadata.Pricing.InputPerMtok, 0.09) + assertPrice(t, "OutputPerMtok", model.Metadata.Pricing.OutputPerMtok, 0.18) + }) + } +} + +// TestListModels_UnusableRateInBothBlocksStaysUnpriced asserts a rate no block +// publishes usably is left absent rather than invented. +func TestListModels_UnusableRateInBothBlocksStaysUnpriced(t *testing.T) { + model := firstModel(t, `{"object":"list","data":[{ + "id":"openai/gpt-4","object":"model", + "pricing": {"input_cost_per_token": 6e-8, "output_cost_per_token": -2e-7}, + "list_pricing": {"input_cost_per_token": 9e-8, "output_cost_per_token": -2.8e-7} + }]}`) + assertPrice(t, "InputPerMtok", model.Metadata.Pricing.InputPerMtok, 0.06) if model.Metadata.Pricing.OutputPerMtok != nil { - t.Errorf("OutputPerMtok = %v, want nil: the list card must not fill gaps in applicable pricing", + t.Errorf("OutputPerMtok = %v, want nil: neither block published a usable output rate", *model.Metadata.Pricing.OutputPerMtok) } } +// TestListModels_CachePricingFields asserts both of Eden's cache rates reach +// the matching core.ModelPricing fields, and that the rates whose usage +// semantics are unconfirmed stay unmapped. +// +// The reasoning and audio rates are the ones deliberately skipped: see +// modelPricing. Asserting they stay nil keeps a future change from wiring them +// up without first establishing how Eden reports the matching token counts. +func TestListModels_CachePricingFields(t *testing.T) { + model := firstModel(t, `{"object":"list","data":[{ + "id":"openai/gpt-4","object":"model", + "pricing": { + "input_cost_per_token": 6e-8, + "output_cost_per_token": 1.8e-7, + "cache_read_input_token_cost": 1.2e-8, + "cache_creation_input_token_cost": 7.5e-8, + "output_cost_per_reasoning_token": 3.6e-7, + "input_cost_per_audio_token": 1e-6 + } + }]}`) + + pricing := model.Metadata.Pricing + assertPrice(t, "CachedInputPerMtok", pricing.CachedInputPerMtok, 0.012) + assertPrice(t, "CacheWritePerMtok", pricing.CacheWritePerMtok, 0.075) + if pricing.ReasoningOutputPerMtok != nil { + t.Errorf("ReasoningOutputPerMtok = %v, want nil: Eden reports no reasoning token count to price against", + *pricing.ReasoningOutputPerMtok) + } + if pricing.AudioInputPerMtok != nil { + t.Errorf("AudioInputPerMtok = %v, want nil: Eden reports no audio token count to price against", + *pricing.AudioInputPerMtok) + } +} + +// TestListModels_TieredAndPerQueryPricingIgnored asserts the Eden pricing +// members the gateway has no equivalent for are skipped rather than guessed at. +// Eden publishes context-length-tiered rates, a tiered_pricing list, and +// per-query search fees (an object, not a scalar); reading any of them into a +// flat per-Mtok field would misprice the model. +func TestListModels_TieredAndPerQueryPricingIgnored(t *testing.T) { + model := firstModel(t, `{"object":"list","data":[{ + "id":"openai/gpt-4","object":"model", + "pricing": { + "input_cost_per_token": 6e-8, + "output_cost_per_token": 1.8e-7, + "input_cost_per_token_above_200k_tokens": 1.2e-7, + "output_cost_per_token_above_200k_tokens": 3.6e-7, + "search_context_cost_per_query": {"search_context_size_low": 0.01}, + "tiered_pricing": [{"input_cost_per_token": 5e-8}] + } + }]}`) + + pricing := model.Metadata.Pricing + assertPrice(t, "InputPerMtok", pricing.InputPerMtok, 0.06) + assertPrice(t, "OutputPerMtok", pricing.OutputPerMtok, 0.18) + if len(pricing.Tiers) != 0 { + t.Errorf("Tiers = %v, want empty: Eden's tiered_pricing shape is not mapped", pricing.Tiers) + } + if pricing.PerRequest != nil { + t.Errorf("PerRequest = %v, want nil: per-query search fees are not per-request charges", *pricing.PerRequest) + } +} + // TestListModels_PricingEdgeCases covers partial, zero, negative, and // overflowing rates. A rate Eden omits must stay absent rather than being // costed at zero; a rate Eden reports as zero must be honoured as free. @@ -329,6 +487,29 @@ func TestListModels_ModalityMapping(t *testing.T) { wantModes: []string{"chat", "responses"}, wantCapabilities: []string{"vision", "audio", "video"}, }, + { + // Eden's live catalog never publishes a video output modality + // (video appears only as an input), but the mode is mapped so a + // model Eden adds later is classified rather than silently + // advertised as chat. + name: "video output becomes video_generation", + capabilities: `{"output_modalities":["video"]}`, + wantModes: []string{"video_generation"}, + }, + { + // A model that outputs both text and video keeps its chat modes, + // so it stays advertised and is reached through chat. + name: "video alongside text keeps the chat modes", + capabilities: `{"output_modalities":["text","video"]}`, + wantModes: []string{"chat", "responses", "video_generation"}, + }, + { + // Text maps to two modes, so a repeated modality would duplicate + // them without the dedup in modes(). + name: "repeated modality is deduplicated", + capabilities: `{"output_modalities":["text","text","image","image"]}`, + wantModes: []string{"chat", "responses", "image_generation"}, + }, { name: "unknown modality is ignored", capabilities: `{"output_modalities":["telepathy"]}`, diff --git a/internal/providers/edenai/response.go b/internal/providers/edenai/response.go index 9d87b3100..a378413ed 100644 --- a/internal/providers/edenai/response.go +++ b/internal/providers/edenai/response.go @@ -113,14 +113,50 @@ func normalizeUpstreamProvider(resp *core.ChatResponse) { resp.ExtraFields = merged } -// normalizeEmbeddingResponse applies the same provider-field reasoning to -// embeddings, which Eden also annotates with the upstream provider and which -// feeds the same gateway.ResponseProviderType labeling. core.EmbeddingResponse -// models no unknown-field container, so the upstream value is dropped rather -// than relocated. -func normalizeEmbeddingResponse(resp *core.EmbeddingResponse) { +// embeddingResponse is Eden's embeddings envelope: the OpenAI-compatible shape +// plus the two members Eden adds to it, the exact per-request "cost" in USD and +// the upstream "provider" that served the request. Both sit at the response +// root, which is why the provider decodes into this wrapper instead of straight +// into core.EmbeddingResponse. +type embeddingResponse struct { + core.EmbeddingResponse + Cost json.RawMessage `json:"cost"` +} + +// normalizeEmbeddingResponse reconciles Eden's embeddings extensions with +// GoModel's response semantics, the same way normalizeChatResponse does for +// chat completions. +func normalizeEmbeddingResponse(resp *embeddingResponse) { if resp == nil { return } + liftEmbeddingCost(resp) + // Eden reports the upstream that served the request; Provider means the + // provider GoModel executed against, and the gateway treats a populated + // value as authoritative when labeling telemetry (see + // normalizeUpstreamProvider). Embeddings carry no unknown-field container + // to relocate the value into, so it is dropped rather than re-exposed. resp.Provider = "" } + +// liftEmbeddingCost copies Eden's root-level embeddings "cost" into +// Usage.RawUsage, where usage.ExtractFromEmbeddingResponse picks it up and +// hands it to the cost pipeline as an exact, provider-reported charge. +// +// This is the embeddings twin of liftResponseCost, and screens the value the +// same way: an existing usage-level cost is the more specific reading and wins, +// and a null or otherwise unusable member is ignored so it cannot be recorded +// as a real $0.00 charge. +func liftEmbeddingCost(resp *embeddingResponse) { + cost, ok := decodeCost(resp.Cost) + if !ok { + return + } + if resp.Usage.RawUsage == nil { + resp.Usage.RawUsage = make(map[string]any, 1) + } + if _, exists := resp.Usage.RawUsage[costField]; exists { + return + } + resp.Usage.RawUsage[costField] = cost +} diff --git a/internal/providers/edenai/response_test.go b/internal/providers/edenai/response_test.go index 93c3aa0fe..14efe01ed 100644 --- a/internal/providers/edenai/response_test.go +++ b/internal/providers/edenai/response_test.go @@ -223,3 +223,28 @@ func TestResponses_InheritsCostLifting(t *testing.T) { t.Fatalf("Responses usage cost = %#v, want 0.0002349 carried through the chat translation", resp.Usage.RawUsage["cost"]) } } + +// TestChatCompletion_PropagatesUpstreamErrorWithoutNormalizing asserts an +// upstream failure is returned as-is. Normalization must not run on an error +// path, where there is no response to reconcile. +func TestChatCompletion_PropagatesUpstreamErrorWithoutNormalizing(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(`{"error":{"message":"bad model","type":"invalid_request_error"}}`)) + })) + defer server.Close() + + provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ + Model: slashedModel, + Messages: []core.Message{{Role: "user", Content: "hi"}}, + }) + + if err == nil { + t.Fatal("ChatCompletion() error = nil, want the upstream error propagated") + } + if resp != nil { + t.Errorf("response = %+v, want nil alongside the error", resp) + } +} diff --git a/internal/providers/edenai/transport.go b/internal/providers/edenai/transport.go new file mode 100644 index 000000000..ffcbe88a0 --- /dev/null +++ b/internal/providers/edenai/transport.go @@ -0,0 +1,91 @@ +package edenai + +import ( + "fmt" + "net" + "net/http" + "net/url" + "strings" + + "github.com/enterpilot/gomodel/internal/httpclient" +) + +// maxRedirects matches net/http's own default redirect budget. Installing a +// CheckRedirect replaces that default wholesale, so the cap has to be +// reasserted here or Eden requests would follow redirect chains forever. +const maxRedirects = 10 + +// credentialSafeURL reports whether Eden's bearer token may be put on the wire +// for u. +// +// Eden is a hosted service reached over TLS, so a cleartext destination would +// expose the gateway's Eden key to anyone on the path. Loopback is the one +// exemption: traffic to the local machine never reaches a network, and it is +// how a local Eden-compatible proxy — and this package's own tests — address +// the provider. A URL that cannot be parsed into a host is treated as unsafe. +func credentialSafeURL(u *url.URL) bool { + if u == nil { + return false + } + if strings.EqualFold(u.Scheme, "https") { + return true + } + return isLoopbackHost(u.Hostname()) +} + +// isLoopbackHost reports whether host names the local machine. Both the +// literal addresses (127.0.0.1, ::1) and the conventional name are accepted, +// because httptest servers use the former and local proxies are usually +// configured with the latter. +func isLoopbackHost(host string) bool { + host = strings.ToLower(strings.TrimSpace(host)) + if host == "localhost" { + return true + } + if ip := net.ParseIP(host); ip != nil { + return ip.IsLoopback() + } + return false +} + +// guardedHTTPClient returns the HTTP client Eden requests go out on, with a +// redirect policy that refuses to carry the bearer token into cleartext. +// +// Go's default policy drops Authorization only when the redirect target is a +// different host; it does not look at the scheme, so an HTTPS -> HTTP redirect +// back to the same host forwards the credential in the clear (verified against +// net/http, not assumed). Both of this provider's construction paths route +// through here so the guarantee does not depend on which one a caller used. +// +// base is the caller-supplied client, or nil for the gateway default client +// (the tuned transport and timeouts llmclient would otherwise install). It is +// never mutated: the policy goes on a shallow copy, so a client shared with +// other callers keeps its own redirect behavior and its own transport. +func guardedHTTPClient(base *http.Client) *http.Client { + if base == nil { + base = httpclient.NewDefaultHTTPClient() + } + guarded := *base + guarded.CheckRedirect = checkRedirect + return &guarded +} + +// checkRedirect refuses any redirect that would send Eden's credential to a +// cleartext destination, and otherwise defers to the same budget net/http +// applies by default. +// +// HTTPS -> HTTPS redirects (same host or not) are followed normally. The +// refusal is deliberate rather than merely stripping the header: an +// unauthenticated retry against a downgraded endpoint can only fail, and a +// named error tells the operator why instead of surfacing an opaque 401. The +// message carries scheme and host only — never the URL's userinfo or query, +// which can hold credentials of their own. +func checkRedirect(req *http.Request, via []*http.Request) error { + if len(via) >= maxRedirects { + return fmt.Errorf("stopped after %d redirects", maxRedirects) + } + if credentialSafeURL(req.URL) { + return nil + } + return fmt.Errorf("edenai: refusing redirect to insecure %s://%s: the API credential would be sent in cleartext", req.URL.Scheme, req.URL.Host) +} diff --git a/internal/providers/edenai/transport_test.go b/internal/providers/edenai/transport_test.go new file mode 100644 index 000000000..0549f8de8 --- /dev/null +++ b/internal/providers/edenai/transport_test.go @@ -0,0 +1,501 @@ +package edenai + +import ( + "context" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "testing" + + "github.com/enterpilot/gomodel/internal/core" + "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers" +) + +// embeddingRequest is the smallest request that reaches the wire, used by the +// tests that care about which headers the transport attached rather than about +// the response body. +func embeddingRequest() *core.EmbeddingRequest { + return &core.EmbeddingRequest{Model: slashedModel, Input: "hello"} +} + +// TestCredentialSafeURL covers the single predicate both the header setter and +// the redirect policy consult, so the rule is pinned in one place: TLS is +// always safe, cleartext is safe only when it cannot leave the machine. +func TestCredentialSafeURL(t *testing.T) { + tests := []struct { + name string + raw string + want bool + }{ + {"default eden endpoint", defaultBaseURL, true}, + {"custom https endpoint", "https://eden.example.com/v3", true}, + {"https uppercase scheme", "HTTPS://eden.example.com/v3", true}, + {"cleartext public host", "http://eden.example.com/v3", false}, + {"cleartext ip", "http://203.0.113.10/v3", false}, + {"cleartext loopback name", "http://localhost:8080/v3", true}, + {"cleartext loopback v4", "http://127.0.0.1:8080/v3", true}, + {"cleartext loopback v4 alt", "http://127.10.20.30:8080/v3", true}, + {"cleartext loopback v6", "http://[::1]:8080/v3", true}, + {"no scheme", "//eden.example.com/v3", false}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + parsed, err := url.Parse(tc.raw) + if err != nil { + t.Fatalf("url.Parse(%q) = %v", tc.raw, err) + } + if got := credentialSafeURL(parsed); got != tc.want { + t.Errorf("credentialSafeURL(%q) = %v, want %v", tc.raw, got, tc.want) + } + }) + } +} + +// TestCredentialSafeURL_NilIsUnsafe asserts the predicate fails closed. A +// request with no parsable URL must not be treated as a TLS destination. +func TestCredentialSafeURL_NilIsUnsafe(t *testing.T) { + if credentialSafeURL(nil) { + t.Error("credentialSafeURL(nil) = true, want false: the predicate must fail closed") + } +} + +// TestSetHeaders_SendsCredentialOverHTTPS asserts the ordinary case still +// authenticates: an HTTPS destination gets the bearer token. +func TestSetHeaders_SendsCredentialOverHTTPS(t *testing.T) { + req, err := http.NewRequest(http.MethodPost, defaultBaseURL+"/chat/completions", nil) + if err != nil { + t.Fatalf("http.NewRequest: %v", err) + } + setHeaders(req, "eden-key") + + if got := req.Header.Get("Authorization"); got != "Bearer eden-key" { + t.Errorf("Authorization = %q, want %q", got, "Bearer eden-key") + } +} + +// TestSetHeaders_WithholdsCredentialOverCleartext is the base-URL half of the +// credential guarantee: an operator-supplied http:// endpoint must not put the +// Eden key on the wire in plain text. +func TestSetHeaders_WithholdsCredentialOverCleartext(t *testing.T) { + req, err := http.NewRequest(http.MethodPost, "http://eden.example.com/v3/chat/completions", nil) + if err != nil { + t.Fatalf("http.NewRequest: %v", err) + } + setHeaders(req, "eden-key") + + if got := req.Header.Get("Authorization"); got != "" { + t.Errorf("Authorization = %q, want empty: the credential must not be sent in cleartext", got) + } +} + +// TestSetHeaders_SendsCredentialOverLoopback pins the exemption that keeps a +// local Eden-compatible proxy — and every httptest server in this package — +// working. Cleartext to the local machine never reaches a network. +func TestSetHeaders_SendsCredentialOverLoopback(t *testing.T) { + req, err := http.NewRequest(http.MethodPost, "http://127.0.0.1:9999/v3/chat/completions", nil) + if err != nil { + t.Fatalf("http.NewRequest: %v", err) + } + setHeaders(req, "eden-key") + + if got := req.Header.Get("Authorization"); got != "Bearer eden-key" { + t.Errorf("Authorization = %q, want %q", got, "Bearer eden-key") + } +} + +// TestCleartextBaseURL_RequestCarriesNoCredential proves the guarantee end to +// end rather than only at the header setter: a provider configured with a +// non-loopback cleartext base URL reaches the upstream unauthenticated. +// +// The server here answers on a loopback address but is addressed through a +// public-looking host, which is what a misconfigured or hijacked +// EDENAI_BASE_URL would look like. +func TestCleartextBaseURL_RequestCarriesNoCredential(t *testing.T) { + var gotAuth string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotAuth = r.Header.Get("Authorization") + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"object":"list","model":"openai/text-embedding-3-small","data":[]}`)) + })) + defer server.Close() + + // Rewrite the loopback host to a name that is not loopback, while still + // dialing the test server, by pointing the transport at it explicitly. + target, err := url.Parse(server.URL) + if err != nil { + t.Fatalf("url.Parse: %v", err) + } + client := &http.Client{Transport: &cleartextRouteTransport{cleartext: target.Host}} + + provider := NewWithHTTPClient("eden-key", "http://eden.example.com/v3", client, llmclient.Hooks{}) + if _, err := provider.Embeddings(context.Background(), embeddingRequest()); err != nil { + t.Fatalf("Embeddings: %v", err) + } + + if gotAuth != "" { + t.Errorf("upstream saw Authorization = %q, want empty: the Eden key must never travel in cleartext", gotAuth) + } +} + +// cleartextRouteTransport routes cleartext requests to a fixed address while +// leaving the request URL's host alone, and sends everything else through base +// (http.DefaultTransport when nil). It lets a test address a public-looking +// http:// host — which is what the credential guard and the redirect guard both +// screen on — while still reaching a local test server, with no DNS involved. +type cleartextRouteTransport struct { + base http.RoundTripper + cleartext string +} + +func (t *cleartextRouteTransport) RoundTrip(req *http.Request) (*http.Response, error) { + routed := req.Clone(req.Context()) + if routed.URL.Scheme == "http" { + routed.URL.Host = t.cleartext + } + base := t.base + if base == nil { + base = http.DefaultTransport + } + return base.RoundTrip(routed) +} + +// TestCompatibleConfig_InstallsRedirectGuardOnBothPaths asserts neither +// construction path can reach the network without the redirect policy. Both +// New and NewWithHTTPClient build their transport through compatibleConfig, so +// covering it here covers both. +func TestCompatibleConfig_InstallsRedirectGuardOnBothPaths(t *testing.T) { + tests := []struct { + name string + client *http.Client + }{ + {"gateway default client", nil}, + {"caller-supplied client", &http.Client{}}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + cfg := compatibleConfig(defaultBaseURL, tc.client) + if cfg.HTTPClient == nil { + t.Fatal("HTTPClient = nil, want a client carrying the redirect guard") + } + if cfg.HTTPClient.CheckRedirect == nil { + t.Error("CheckRedirect = nil, want Eden's redirect guard installed") + } + }) + } +} + +// TestGuardedHTTPClient_DoesNotMutateCallerClient asserts the caller's client +// is copied rather than reconfigured. A client shared with another provider +// must keep its own redirect behavior. +func TestGuardedHTTPClient_DoesNotMutateCallerClient(t *testing.T) { + caller := &http.Client{} + guarded := guardedHTTPClient(caller) + + if caller.CheckRedirect != nil { + t.Error("caller's CheckRedirect was set; the guard must be installed on a copy") + } + if guarded == caller { + t.Error("guardedHTTPClient returned the caller's client; want a copy") + } + if guarded.CheckRedirect == nil { + t.Error("guarded client has no CheckRedirect") + } +} + +// TestGuardedHTTPClient_PreservesCallerTransport asserts copying the client +// keeps the transport, so a caller that configured TLS roots or a proxy (and +// httptest's own client) still works. +func TestGuardedHTTPClient_PreservesCallerTransport(t *testing.T) { + transport := &http.Transport{} + caller := &http.Client{Transport: transport} + + if got := guardedHTTPClient(caller).Transport; got != transport { + t.Errorf("Transport = %v, want the caller's transport preserved", got) + } +} + +// TestCheckRedirect_AllowsHTTPSTargets asserts TLS redirects keep working, +// including to a different host: that is an ordinary API redirect and the +// credential stays encrypted. +func TestCheckRedirect_AllowsHTTPSTargets(t *testing.T) { + for _, target := range []string{ + "https://api.edenai.run/v3/chat/completions", + "https://eu.edenai.run/v3/chat/completions", + } { + req, err := http.NewRequest(http.MethodGet, target, nil) + if err != nil { + t.Fatalf("http.NewRequest: %v", err) + } + if err := checkRedirect(req, nil); err != nil { + t.Errorf("checkRedirect(%q) = %v, want nil", target, err) + } + } +} + +// TestCheckRedirect_RefusesSchemeDowngrade is the core of the redirect +// guarantee. net/http drops Authorization only when the redirect target is a +// different host; it does not consider the scheme, so a same-host HTTPS -> HTTP +// redirect would otherwise forward the bearer token in the clear. +func TestCheckRedirect_RefusesSchemeDowngrade(t *testing.T) { + req, err := http.NewRequest(http.MethodGet, "http://api.edenai.run/v3/chat/completions", nil) + if err != nil { + t.Fatalf("http.NewRequest: %v", err) + } + + err = checkRedirect(req, nil) + if err == nil { + t.Fatal("checkRedirect = nil, want a refusal for a cleartext redirect target") + } + if !strings.Contains(err.Error(), "api.edenai.run") { + t.Errorf("error %q should name the refused host", err) + } +} + +// TestCheckRedirect_ErrorOmitsCredentials asserts the refusal message cannot +// leak a secret carried in the redirect URL's userinfo or query string. +func TestCheckRedirect_ErrorOmitsCredentials(t *testing.T) { + req, err := http.NewRequest(http.MethodGet, "http://user:s3cret@api.edenai.run/v3?api_key=leaked", nil) + if err != nil { + t.Fatalf("http.NewRequest: %v", err) + } + + err = checkRedirect(req, nil) + if err == nil { + t.Fatal("checkRedirect = nil, want a refusal") + } + for _, secret := range []string{"s3cret", "leaked"} { + if strings.Contains(err.Error(), secret) { + t.Errorf("error %q leaked %q from the redirect URL", err, secret) + } + } +} + +// TestCheckRedirect_EnforcesRedirectBudget asserts installing a policy did not +// silently remove net/http's own protection against endless redirect chains. +func TestCheckRedirect_EnforcesRedirectBudget(t *testing.T) { + req, err := http.NewRequest(http.MethodGet, defaultBaseURL, nil) + if err != nil { + t.Fatalf("http.NewRequest: %v", err) + } + via := make([]*http.Request, maxRedirects) + + if err := checkRedirect(req, via); err == nil { + t.Fatalf("checkRedirect with %d prior hops = nil, want the redirect budget enforced", maxRedirects) + } +} + +// TestRedirect_HTTPSToHTTPDoesNotForwardCredential exercises the guard through +// a real transport: an HTTPS Eden endpoint that redirects to a cleartext host +// must not deliver the bearer token there. +// +// The redirect names a non-loopback host on purpose. Cleartext to loopback is +// deliberately allowed (see TestRedirect_CleartextLoopbackStillFollowed), so +// pointing this at the test server's own 127.0.0.1 address would exercise the +// exemption instead of the guard. +func TestRedirect_HTTPSToHTTPDoesNotForwardCredential(t *testing.T) { + var cleartextAuth string + var cleartextHits int + cleartext := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + cleartextHits++ + cleartextAuth = r.Header.Get("Authorization") + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"object":"list","data":[]}`)) + })) + defer cleartext.Close() + + secure := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, "http://eden.example.com/v3/embeddings", http.StatusFound) + })) + defer secure.Close() + + cleartextTarget, err := url.Parse(cleartext.URL) + if err != nil { + t.Fatalf("url.Parse: %v", err) + } + client := secure.Client() + client.Transport = &cleartextRouteTransport{ + base: client.Transport, + cleartext: cleartextTarget.Host, + } + + provider := NewWithHTTPClient("eden-key", secure.URL+"/v3", client, llmclient.Hooks{}) + _, err = provider.Embeddings(context.Background(), embeddingRequest()) + if err == nil { + t.Fatal("Embeddings succeeded through an HTTPS -> HTTP redirect, want the redirect refused") + } + + if cleartextHits != 0 { + t.Errorf("cleartext endpoint received %d request(s), want 0", cleartextHits) + } + if cleartextAuth != "" { + t.Errorf("cleartext endpoint saw Authorization = %q, want empty", cleartextAuth) + } +} + +// TestRedirect_CleartextLoopbackStillFollowed pins the exemption's scope: a +// cleartext redirect that stays on the local machine is followed and still +// authenticates, which is what keeps a local Eden-compatible proxy usable. +func TestRedirect_CleartextLoopbackStillFollowed(t *testing.T) { + var finalAuth string + var served bool + + mux := http.NewServeMux() + mux.HandleFunc("/v3/embeddings", func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, "/v3/embeddings-moved", http.StatusTemporaryRedirect) + }) + mux.HandleFunc("/v3/embeddings-moved", func(w http.ResponseWriter, r *http.Request) { + served = true + finalAuth = r.Header.Get("Authorization") + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"object":"list","model":"openai/text-embedding-3-small","data":[]}`)) + }) + server := httptest.NewServer(mux) + defer server.Close() + + provider := NewWithHTTPClient("eden-key", server.URL+"/v3", server.Client(), llmclient.Hooks{}) + if _, err := provider.Embeddings(context.Background(), embeddingRequest()); err != nil { + t.Fatalf("Embeddings through a loopback redirect: %v", err) + } + if !served { + t.Fatal("redirect target was never reached") + } + if finalAuth != "Bearer eden-key" { + t.Errorf("redirect target saw Authorization = %q, want %q", finalAuth, "Bearer eden-key") + } +} + +// TestRedirect_HTTPSToHTTPSStillFollowed asserts the guard is narrow: a +// same-host TLS redirect is followed and still authenticates, so a legitimate +// custom HTTPS endpoint that redirects keeps working. +func TestRedirect_HTTPSToHTTPSStillFollowed(t *testing.T) { + var finalAuth string + var served bool + + mux := http.NewServeMux() + mux.HandleFunc("/v3/embeddings", func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, "/v3/embeddings-moved", http.StatusTemporaryRedirect) + }) + mux.HandleFunc("/v3/embeddings-moved", func(w http.ResponseWriter, r *http.Request) { + served = true + finalAuth = r.Header.Get("Authorization") + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"object":"list","model":"openai/text-embedding-3-small","data":[],"usage":{"prompt_tokens":9,"total_tokens":9}}`)) + }) + secure := httptest.NewTLSServer(mux) + defer secure.Close() + + provider := NewWithHTTPClient("eden-key", secure.URL+"/v3", secure.Client(), llmclient.Hooks{}) + resp, err := provider.Embeddings(context.Background(), embeddingRequest()) + if err != nil { + t.Fatalf("Embeddings through an HTTPS -> HTTPS redirect: %v", err) + } + if !served { + t.Fatal("redirect target was never reached") + } + if resp == nil { + t.Fatal("response = nil") + } + if finalAuth != "Bearer eden-key" { + t.Errorf("redirect target saw Authorization = %q, want %q", finalAuth, "Bearer eden-key") + } +} + +// TestDefaultBaseURLIsHTTPS guards the default every deployment uses. The +// credential guard keys off the request scheme, so a default that regressed to +// http:// would silently withhold the API key on ordinary traffic. +// TestNew_DefaultsBaseURL already covers New resolving to this value. +func TestDefaultBaseURLIsHTTPS(t *testing.T) { + if !strings.HasPrefix(defaultBaseURL, "https://") { + t.Fatalf("defaultBaseURL = %q, want an https:// endpoint", defaultBaseURL) + } +} + +// TestSetBaseURL_OverridesResolvedEndpoint covers the public base-URL override, +// which the credential guard then reads at request time rather than using the +// endpoint the provider was constructed with. +func TestSetBaseURL_OverridesResolvedEndpoint(t *testing.T) { + provider, ok := New(providers.ProviderConfig{APIKey: "eden-key"}, providers.ProviderOptions{}).(*Provider) + if !ok { + t.Fatal("New did not return *Provider") + } + + const override = "https://eden.eu.example.com/v3" + provider.SetBaseURL(override) + if got := provider.GetBaseURL(); got != override { + t.Errorf("GetBaseURL() = %q, want %q after SetBaseURL", got, override) + } +} + +// TestGuardedHTTPClient_PreservesDefaultClientSemantics pins what +// NewWithHTTPClient's nil path produces. Every other chat-compatible provider +// documents "if httpClient is nil, http.DefaultClient is used", and gets that +// from NewCompatibleProviderWithHTTPClient; Eden hands that helper an +// already-guarded client, so it substitutes http.DefaultClient itself. Guarding +// that client must leave its transport and its absent timeout alone, or the +// constructor would quietly diverge from every peer. +// +// The guard also has to go on a copy: writing CheckRedirect onto +// http.DefaultClient would change redirect behavior for every other user of +// that global in the process. That is the assertion that matters most here. +// +// The substitution itself is not observable from outside the provider (the +// transport lives on an unexported field of openai.CompatibleProvider), so it +// is pinned by this test together with +// TestGuardedHTTPClient_NilBuildsGatewayDefault, which shows the two callers +// deliberately get different clients. +func TestGuardedHTTPClient_PreservesDefaultClientSemantics(t *testing.T) { + guarded := guardedHTTPClient(http.DefaultClient) + + if guarded == http.DefaultClient { + t.Fatal("guardedHTTPClient returned http.DefaultClient itself; want a copy") + } + if http.DefaultClient.CheckRedirect != nil { + t.Error("http.DefaultClient.CheckRedirect was set; the process-wide client must not be modified") + } + if guarded.CheckRedirect == nil { + t.Error("CheckRedirect = nil, want Eden's redirect guard on the copy") + } + if guarded.Transport != http.DefaultClient.Transport { + t.Errorf("Transport = %v, want http.DefaultClient's transport preserved", guarded.Transport) + } + if guarded.Timeout != http.DefaultClient.Timeout { + t.Errorf("Timeout = %v, want http.DefaultClient's %v, not the gateway client's", guarded.Timeout, http.DefaultClient.Timeout) + } +} + +// TestGuardedHTTPClient_NilBuildsGatewayDefault covers the other caller: New +// passes nil because the factory path has no client of its own, and must get +// the tuned transport and timeouts llmclient would otherwise have installed — +// not http.DefaultClient's absent timeout. +func TestGuardedHTTPClient_NilBuildsGatewayDefault(t *testing.T) { + guarded := guardedHTTPClient(nil) + + if guarded.CheckRedirect == nil { + t.Error("CheckRedirect = nil, want Eden's redirect guard") + } + if guarded.Timeout <= 0 { + t.Errorf("Timeout = %v, want the gateway default client's positive timeout", guarded.Timeout) + } +} + +// TestNewWithHTTPClient_NilClientStillGuardsRedirects proves the nil path is +// wired end to end: a provider built with no client still refuses a +// credential-leaking redirect rather than following it. +func TestNewWithHTTPClient_NilClientStillGuardsRedirects(t *testing.T) { + provider := NewWithHTTPClient("eden-key", "http://eden.example.com/v3", nil, llmclient.Hooks{}) + if provider == nil { + t.Fatal("NewWithHTTPClient(..., nil, ...) returned nil") + } + + req, err := http.NewRequest(http.MethodGet, "http://api.edenai.run/v3/chat/completions", nil) + if err != nil { + t.Fatalf("http.NewRequest: %v", err) + } + if err := checkRedirect(req, nil); err == nil { + t.Error("checkRedirect = nil for a cleartext target; the guard must apply on the nil-client path too") + } +} diff --git a/internal/usage/cost.go b/internal/usage/cost.go index a640acb8b..9b0694b07 100644 --- a/internal/usage/cost.go +++ b/internal/usage/cost.go @@ -600,6 +600,20 @@ func isEdenAIProvider(providerType string) bool { return strings.EqualFold(strings.TrimSpace(providerType), "edenai") } +// isProviderReportedCostSource reports whether a CostSource names a charge the +// provider itself returned, rather than one reconstructed from token counts and +// a rate card. Callers use it to suppress caveats that only make sense for a +// rate-card reconstruction: a provider-reported total is authoritative no +// matter what token counts came with it. +func isProviderReportedCostSource(source string) bool { + switch strings.TrimSpace(source) { + case CostSourceOpenRouterCredits, CostSourceXAITicks, CostSourceEdenAICost: + return true + default: + return false + } +} + func openRouterCreditCostSplit(rawData map[string]any, total float64) (float64, float64, bool) { details, ok := nestedUsageMap(rawData["cost_details"]) if !ok { diff --git a/internal/usage/extractor.go b/internal/usage/extractor.go index 916728148..78a6480df 100644 --- a/internal/usage/extractor.go +++ b/internal/usage/extractor.go @@ -193,6 +193,11 @@ func ExtractFromEmbeddingResponse(resp *core.EmbeddingResponse, requestID, provi Endpoint: endpoint, InputTokens: resp.Usage.PromptTokens, TotalTokens: resp.Usage.TotalTokens, + // Carry the provider's extra usage members into the cost pipeline, so + // an embeddings provider that returns an exact per-request charge is + // costed from that figure rather than from token rates. Nil for + // providers that report nothing beyond the token counts. + RawData: cloneRawData(resp.Usage.RawUsage), } applyUsageCosts(entry, provider, endpoint, pricing...) @@ -201,8 +206,11 @@ func ExtractFromEmbeddingResponse(resp *core.EmbeddingResponse, requestID, provi // rates and understates the real call. Flag that — but not when the // configured pricing determines the cost without token counts (a // per-request price, or an explicit zero rate), where the recorded cost is - // correct and calling it uncalculated would be false. + // correct and calling it uncalculated would be false, and not when the + // cost came from a figure the provider itself reported, which is + // authoritative however many tokens it counted. if resp.Usage.PromptTokens == 0 && resp.Usage.TotalTokens == 0 && + !isProviderReportedCostSource(entry.CostSource) && tokenRatesAffectCost(effectiveEndpointPricing(endpoint, pricing...)) && entry.CostsCalculationCaveat == "" { entry.CostsCalculationCaveat = caveatEmbeddingMissingUsage diff --git a/internal/usage/extractor_test.go b/internal/usage/extractor_test.go index 210df6666..e5cf921bc 100644 --- a/internal/usage/extractor_test.go +++ b/internal/usage/extractor_test.go @@ -1083,3 +1083,117 @@ func TestExtractFromChatResponse_EdenAIWithoutCostFallsBackToPricing(t *testing. t.Fatalf("TotalCost = %v, want 0.15 from discovered per-model pricing", entry.TotalCost) } } + +// TestExtractFromEmbeddingResponse_ForwardsRawUsage asserts the provider's +// extra usage members reach the cost pipeline. Without this the embeddings +// path has no channel for a provider-reported exact charge at all, because +// ExtractFromEmbeddingResponse builds the entry itself. +func TestExtractFromEmbeddingResponse_ForwardsRawUsage(t *testing.T) { + resp := &core.EmbeddingResponse{ + Model: "openai/text-embedding-3-small", + Usage: core.EmbeddingUsage{ + PromptTokens: 9, + TotalTokens: 9, + RawUsage: map[string]any{"cost": 0.0000012}, + }, + } + + entry := ExtractFromEmbeddingResponse(resp, "req", "edenai", "/v1/embeddings") + if entry == nil { + t.Fatal("ExtractFromEmbeddingResponse returned nil") + } + if got := entry.RawData["cost"]; got != 0.0000012 { + t.Errorf("RawData[cost] = %v, want 0.0000012 forwarded from RawUsage", got) + } + if entry.TotalCost == nil || *entry.TotalCost != 0.0000012 { + t.Errorf("TotalCost = %v, want the provider-reported 0.0000012", entry.TotalCost) + } + if entry.CostSource != CostSourceEdenAICost { + t.Errorf("CostSource = %q, want %q", entry.CostSource, CostSourceEdenAICost) + } +} + +// TestExtractFromEmbeddingResponse_NilRawUsageStaysNil asserts a provider that +// reports nothing beyond the token counts is unaffected: RawData stays nil, so +// cost is calculated exactly as it was before embeddings gained the carrier. +func TestExtractFromEmbeddingResponse_NilRawUsageStaysNil(t *testing.T) { + rate := 0.02 + resp := &core.EmbeddingResponse{ + Model: "text-embedding-3-small", + Usage: core.EmbeddingUsage{PromptTokens: 1000, TotalTokens: 1000}, + } + + entry := ExtractFromEmbeddingResponse(resp, "req", "openai", "/v1/embeddings", + &core.ModelPricing{Currency: "USD", InputPerMtok: &rate}) + if entry.RawData != nil { + t.Errorf("RawData = %v, want nil when the provider reported no extra usage", entry.RawData) + } + if entry.TotalCost == nil || *entry.TotalCost != 0.00002 { + t.Errorf("TotalCost = %v, want 0.00002 from the token rate", entry.TotalCost) + } + if entry.CostSource != CostSourceModelPricing { + t.Errorf("CostSource = %q, want %q", entry.CostSource, CostSourceModelPricing) + } +} + +// TestExtractFromEmbeddingResponse_ExactCostSuppressesMissingUsageCaveat +// asserts a zero-token row is not flagged as uncalculated when the cost came +// from a figure the provider itself reported. +// +// The caveat exists because zero tokens priced from token rates understate the +// call. A provider-reported total is authoritative however many tokens came +// with it, so flagging it would tell operators the cost is unreliable when it +// is the most reliable number available. +func TestExtractFromEmbeddingResponse_ExactCostSuppressesMissingUsageCaveat(t *testing.T) { + rate := 0.02 + resp := &core.EmbeddingResponse{ + Model: "openai/text-embedding-3-small", + Usage: core.EmbeddingUsage{RawUsage: map[string]any{"cost": 0.0000012}}, + } + + entry := ExtractFromEmbeddingResponse(resp, "req", "edenai", "/v1/embeddings", + &core.ModelPricing{Currency: "USD", InputPerMtok: &rate}) + if entry.CostsCalculationCaveat != "" { + t.Errorf("CostsCalculationCaveat = %q, want empty: the cost was reported by the provider", entry.CostsCalculationCaveat) + } + if entry.TotalCost == nil || *entry.TotalCost != 0.0000012 { + t.Errorf("TotalCost = %v, want the provider-reported 0.0000012", entry.TotalCost) + } +} + +// TestExtractFromEmbeddingResponse_ZeroTokenCaveatStillApplies asserts the +// guard above did not disable the caveat generally: a zero-token row with no +// provider-reported cost is still flagged. +func TestExtractFromEmbeddingResponse_ZeroTokenCaveatStillApplies(t *testing.T) { + rate := 0.02 + resp := &core.EmbeddingResponse{Model: "gemini-embedding-001"} + + entry := ExtractFromEmbeddingResponse(resp, "req", "gemini", "/v1/embeddings", + &core.ModelPricing{Currency: "USD", InputPerMtok: &rate}) + if entry.CostsCalculationCaveat == "" { + t.Error("CostsCalculationCaveat = empty, want the zero-token row flagged") + } +} + +// TestisProviderReportedCostSource pins which cost sources count as figures the +// provider returned rather than rate-card reconstructions. +func TestIsProviderReportedCostSource(t *testing.T) { + tests := []struct { + source string + want bool + }{ + {CostSourceOpenRouterCredits, true}, + {CostSourceXAITicks, true}, + {CostSourceEdenAICost, true}, + {" " + CostSourceEdenAICost + " ", true}, + {CostSourceModelPricing, false}, + {"", false}, + {"something_else", false}, + } + + for _, tc := range tests { + if got := isProviderReportedCostSource(tc.source); got != tc.want { + t.Errorf("isProviderReportedCostSource(%q) = %v, want %v", tc.source, got, tc.want) + } + } +} From 3e162558d003d2ebbd0ae5f9de72d73096b2067d Mon Sep 17 00:00:00 2001 From: NaDdjg Date: Fri, 11 Sep 2026 12:31:06 +0100 Subject: [PATCH 3/6] fix: harden redirect and credential handling --- docs/providers/edenai.mdx | 23 +- internal/providers/edenai/edenai.go | 13 +- internal/providers/edenai/edenai_test.go | 8 + internal/providers/edenai/transport.go | 138 +++++-- internal/providers/edenai/transport_test.go | 437 ++++++++++++++++++-- 5 files changed, 545 insertions(+), 74 deletions(-) diff --git a/docs/providers/edenai.mdx b/docs/providers/edenai.mdx index 80fafdb4d..d490bb9c4 100644 --- a/docs/providers/edenai.mdx +++ b/docs/providers/edenai.mdx @@ -28,12 +28,23 @@ EDENAI_BASE_URL=https://api.edenai.run/v3 ``` - Use an `https://` endpoint. GoModel withholds the Eden API key from any - cleartext destination, so an `http://` base URL reaches Eden - unauthenticated and fails with a 401 rather than sending the key in the - clear. The same applies to redirects: an HTTPS endpoint that redirects to - `http://` is refused instead of followed. Cleartext to `localhost` is - allowed, which is what keeps a local Eden-compatible proxy usable. + Use an `https://` endpoint. GoModel **refuses** any Eden request bound for a + cleartext destination rather than sending it, because the request body + carries the prompt or the embedding input, not just the API key. A + non-loopback `http://` base URL therefore fails with an explicit error + instead of transmitting anything. + + Redirects are held to the same standard, and only a redirect that stays on + the host you configured is followed. A redirect to `http://`, or to any + other host — including a subdomain, which Go would otherwise let carry the + `Authorization` header — is refused rather than followed, so an + upstream-chosen `Location` cannot move your key or your prompt somewhere you + did not configure. + + Cleartext to `localhost` (or `127.0.0.1`) is allowed, since that traffic + never reaches a network -- this is what keeps a local Eden-compatible proxy + usable. If you front Eden with a proxy on another host, terminate TLS on it + and point `EDENAI_BASE_URL` at its `https://` address. Or in `config.yaml`: diff --git a/internal/providers/edenai/edenai.go b/internal/providers/edenai/edenai.go index 81655c4c0..99829bf58 100644 --- a/internal/providers/edenai/edenai.go +++ b/internal/providers/edenai/edenai.go @@ -118,14 +118,13 @@ func compatibleConfig(baseURL string, httpClient *http.Client) openai.Compatible // sends no credential when SetHeaders is nil (unlike ChatCompatible, which // defaults to bearer), so this must stay wired up. // -// The credential is withheld from a destination that would carry it in -// cleartext. Eden's base URL is operator-supplied, so an http:// override — -// whether set by mistake or by a downgrade attempt — would otherwise put the -// gateway's Eden key on the wire in plain text. Loopback is exempt, which is -// what keeps local proxies and this package's httptest servers working. A -// withheld credential yields an Eden 401 rather than a leaked key. +// The credential is also withheld from a destination that would carry it in +// cleartext. secureTransport already refuses such a request outright, so this +// is defense in depth: it keeps the key out of the request even if the +// transport guard is ever bypassed or removed. Loopback is exempt, which is +// what keeps local proxies and this package's httptest servers working. func setHeaders(req *http.Request, apiKey string) { - if !credentialSafeURL(req.URL) { + if !secureDestination(req.URL) { return } providers.SetAuthHeaders(req, apiKey, providers.AuthHeaderConfig{AuthScheme: "Bearer "}) diff --git a/internal/providers/edenai/edenai_test.go b/internal/providers/edenai/edenai_test.go index 052737c66..548ebc9f7 100644 --- a/internal/providers/edenai/edenai_test.go +++ b/internal/providers/edenai/edenai_test.go @@ -316,6 +316,14 @@ func TestEmbeddings_ForwardsToEmbeddingsEndpoint(t *testing.T) { if gotBody["model"] != "openai/text-embedding-3-small" { t.Fatalf("request model = %#v, want provider/model ID forwarded unchanged", gotBody["model"]) } + // core.EmbeddingRequest.Provider is a gateway routing hint that the router + // clears on the forwarded clone (providers.forwardEmbeddingRequest) before + // any provider sees it, so it must not appear on the wire. Eden dispatches + // the request it is handed, exactly as the shared + // CompatibleProvider.Embeddings helper does. + if _, leaked := gotBody["provider"]; leaked { + t.Errorf("request body carried the gateway-only provider field: %#v", gotBody) + } if len(resp.Data) != 1 || resp.Data[0].Index != 0 || len(resp.Data[0].Embedding) == 0 { t.Fatalf("embedding data = %+v, want one populated vector", resp.Data) } diff --git a/internal/providers/edenai/transport.go b/internal/providers/edenai/transport.go index ffcbe88a0..b0ad95ce0 100644 --- a/internal/providers/edenai/transport.go +++ b/internal/providers/edenai/transport.go @@ -15,15 +15,20 @@ import ( // reasserted here or Eden requests would follow redirect chains forever. const maxRedirects = 10 -// credentialSafeURL reports whether Eden's bearer token may be put on the wire -// for u. +// secureDestination reports whether an Eden request may be sent to u at all. // // Eden is a hosted service reached over TLS, so a cleartext destination would -// expose the gateway's Eden key to anyone on the path. Loopback is the one -// exemption: traffic to the local machine never reaches a network, and it is -// how a local Eden-compatible proxy — and this package's own tests — address -// the provider. A URL that cannot be parsed into a host is treated as unsafe. -func credentialSafeURL(u *url.URL) bool { +// expose everything on the wire to anyone on the path: not only the gateway's +// Eden key, but the prompt, the embedding input, and the completion coming +// back. Withholding the credential alone would still ship the payload, so this +// predicate gates the request itself (see secureTransport) as well as the +// Authorization header. +// +// Loopback is the one exemption: traffic to the local machine never reaches a +// network, and it is how a local Eden-compatible proxy — and this package's own +// tests — address the provider. A URL that cannot be parsed into a host is +// treated as unsafe. +func secureDestination(u *url.URL) bool { if u == nil { return false } @@ -48,44 +53,125 @@ func isLoopbackHost(host string) bool { return false } -// guardedHTTPClient returns the HTTP client Eden requests go out on, with a -// redirect policy that refuses to carry the bearer token into cleartext. +// guardedHTTPClient returns the HTTP client Eden requests go out on, carrying +// two policies that keep an Eden request off a cleartext connection: the +// transport refuses an insecure destination outright, and the redirect policy +// refuses to be steered onto one. // -// Go's default policy drops Authorization only when the redirect target is a -// different host; it does not look at the scheme, so an HTTPS -> HTTP redirect -// back to the same host forwards the credential in the clear (verified against +// The redirect half is not redundant with Go's own behavior. Go's default +// policy drops Authorization only when the redirect target is a different +// host; it does not look at the scheme, so an HTTPS -> HTTP redirect back to +// the same host forwards the credential in the clear (verified against // net/http, not assumed). Both of this provider's construction paths route // through here so the guarantee does not depend on which one a caller used. // // base is the caller-supplied client, or nil for the gateway default client // (the tuned transport and timeouts llmclient would otherwise install). It is -// never mutated: the policy goes on a shallow copy, so a client shared with +// never mutated: both policies go on a shallow copy, so a client shared with // other callers keeps its own redirect behavior and its own transport. func guardedHTTPClient(base *http.Client) *http.Client { if base == nil { base = httpclient.NewDefaultHTTPClient() } guarded := *base - guarded.CheckRedirect = checkRedirect + guarded.Transport = &secureTransport{base: guarded.Transport} + guarded.CheckRedirect = redirectPolicy(base.CheckRedirect) return &guarded } -// checkRedirect refuses any redirect that would send Eden's credential to a -// cleartext destination, and otherwise defers to the same budget net/http -// applies by default. +// redirectPolicy composes Eden's redirect rules with whatever policy the +// caller's client already carried. // -// HTTPS -> HTTPS redirects (same host or not) are followed normally. The -// refusal is deliberate rather than merely stripping the header: an -// unauthenticated retry against a downgraded endpoint can only fail, and a -// named error tells the operator why instead of surfacing an opaque 401. The -// message carries scheme and host only — never the URL's userinfo or query, -// which can hold credentials of their own. +// Eden's rules run first and a refusal is final, so a permissive caller policy +// cannot waive the cleartext or cross-host guarantees; a caller can only +// restrict further. When Eden allows the hop, the caller's callback decides, +// and its result is returned verbatim -- including http.ErrUseLastResponse, +// the sentinel that means "stop here and hand back the redirect response". +// +// Replacing the field outright would silently discard that policy (verified: +// an overwritten CheckRedirect never runs and the redirect is followed anyway), +// which is why the guard composes here the same way secureTransport wraps the +// caller's transport rather than replacing it. +func redirectPolicy(caller func(*http.Request, []*http.Request) error) func(*http.Request, []*http.Request) error { + return func(req *http.Request, via []*http.Request) error { + if err := checkRedirect(req, via); err != nil { + return err + } + if caller == nil { + return nil + } + return caller(req, via) + } +} + +// checkRedirect applies Eden's own redirect rules: stay encrypted, stay on the +// host the request was addressed to, and respect net/http's redirect budget. +// +// Only same-host redirects are followed. Go's default policy forwards +// Authorization to a subdomain of the original host (verified: a redirect from +// api.edenai.run to evil.api.edenai.run arrives carrying the bearer token; an +// unrelated host does not), so an upstream-controlled Location header would +// otherwise be enough to hand the Eden key to a neighbouring name. Refusing +// rather than merely stripping the header is deliberate for the same reason +// secureTransport refuses: following the hop would still ship the prompt or +// embedding input to a destination the operator never configured, which is the +// server-side request forgery half of the problem. +// +// A path-only redirect on the configured host -- the one shape a REST endpoint +// plausibly returns -- keeps working. func checkRedirect(req *http.Request, via []*http.Request) error { if len(via) >= maxRedirects { return fmt.Errorf("stopped after %d redirects", maxRedirects) } - if credentialSafeURL(req.URL) { - return nil + if !secureDestination(req.URL) { + return insecureDestinationError("follow redirect to", req.URL) + } + // via is ordered oldest first, so via[0] is the request the caller made. + // Compare against that rather than the previous hop: a chain of same-step + // redirects must not be able to walk away from the original host one label + // at a time. + if len(via) > 0 && via[0].URL != nil && !strings.EqualFold(via[0].URL.Host, req.URL.Host) { + return fmt.Errorf("edenai: refusing cross-host redirect from %s to %s: the credential and the request body must not follow an upstream-chosen destination", + via[0].URL.Host, req.URL.Host) } - return fmt.Errorf("edenai: refusing redirect to insecure %s://%s: the API credential would be sent in cleartext", req.URL.Scheme, req.URL.Host) + return nil +} + +// secureTransport stops a request bound for a cleartext destination before any +// of it reaches the network. +// +// Withholding the Authorization header protects the credential but nothing +// else: the request body still carries the prompt or the embedding input, and +// it is written to the socket before the upstream ever gets to reject the call +// (verified: the payload arrives in full). Refusing the request is what +// actually prevents the disclosure, and it matches how checkRedirect already +// handles a downgraded redirect target. +// +// The check lives on the transport rather than in each provider method so it +// covers every path uniformly — chat, streaming, embeddings, model listing, and +// passthrough, which forwards an opaque caller-supplied body — and so it reads +// the endpoint in force at request time, which SetBaseURL can change after +// construction. +type secureTransport struct { + base http.RoundTripper +} + +func (t *secureTransport) RoundTrip(req *http.Request) (*http.Response, error) { + if !secureDestination(req.URL) { + return nil, insecureDestinationError("send request to", req.URL) + } + base := t.base + if base == nil { + // A nil Transport means http.DefaultTransport, the same default + // net/http applies (http.DefaultClient leaves it nil). + base = http.DefaultTransport + } + return base.RoundTrip(req) +} + +// insecureDestinationError describes a refused destination without quoting the +// URL wholesale: scheme and host only, never the userinfo or query string, +// which can hold credentials of their own. +func insecureDestinationError(action string, u *url.URL) error { + return fmt.Errorf("edenai: refusing to %s insecure %s://%s: the request would travel in cleartext", action, u.Scheme, u.Host) } diff --git a/internal/providers/edenai/transport_test.go b/internal/providers/edenai/transport_test.go index 0549f8de8..3de5fd55a 100644 --- a/internal/providers/edenai/transport_test.go +++ b/internal/providers/edenai/transport_test.go @@ -2,6 +2,8 @@ package edenai import ( "context" + "errors" + "io" "net/http" "net/http/httptest" "net/url" @@ -47,8 +49,8 @@ func TestCredentialSafeURL(t *testing.T) { if err != nil { t.Fatalf("url.Parse(%q) = %v", tc.raw, err) } - if got := credentialSafeURL(parsed); got != tc.want { - t.Errorf("credentialSafeURL(%q) = %v, want %v", tc.raw, got, tc.want) + if got := secureDestination(parsed); got != tc.want { + t.Errorf("secureDestination(%q) = %v, want %v", tc.raw, got, tc.want) } }) } @@ -57,8 +59,8 @@ func TestCredentialSafeURL(t *testing.T) { // TestCredentialSafeURL_NilIsUnsafe asserts the predicate fails closed. A // request with no parsable URL must not be treated as a TLS destination. func TestCredentialSafeURL_NilIsUnsafe(t *testing.T) { - if credentialSafeURL(nil) { - t.Error("credentialSafeURL(nil) = true, want false: the predicate must fail closed") + if secureDestination(nil) { + t.Error("secureDestination(nil) = true, want false: the predicate must fail closed") } } @@ -106,24 +108,31 @@ func TestSetHeaders_SendsCredentialOverLoopback(t *testing.T) { } } -// TestCleartextBaseURL_RequestCarriesNoCredential proves the guarantee end to -// end rather than only at the header setter: a provider configured with a -// non-loopback cleartext base URL reaches the upstream unauthenticated. +// TestCleartextBaseURL_RequestRefusedBeforeSending is the payload half of the +// cleartext guarantee, and the reason withholding the credential is not enough +// on its own. // -// The server here answers on a loopback address but is addressed through a +// A request to a non-loopback http:// endpoint is written to the socket in +// full before the upstream can reject it, so an unauthenticated send still +// discloses the prompt — and on the embeddings surface, the input text — to +// anyone on the path. Nothing may reach the endpoint at all. +// +// The server answers on a loopback address but is addressed through a // public-looking host, which is what a misconfigured or hijacked // EDENAI_BASE_URL would look like. -func TestCleartextBaseURL_RequestCarriesNoCredential(t *testing.T) { - var gotAuth string +func TestCleartextBaseURL_RequestRefusedBeforeSending(t *testing.T) { + var hits int + var gotAuth, gotBody string server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hits++ gotAuth = r.Header.Get("Authorization") + body, _ := io.ReadAll(r.Body) + gotBody = string(body) w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"object":"list","model":"openai/text-embedding-3-small","data":[]}`)) + _, _ = w.Write([]byte(`{"id":"x","choices":[]}`)) })) defer server.Close() - // Rewrite the loopback host to a name that is not loopback, while still - // dialing the test server, by pointing the transport at it explicitly. target, err := url.Parse(server.URL) if err != nil { t.Fatalf("url.Parse: %v", err) @@ -131,15 +140,149 @@ func TestCleartextBaseURL_RequestCarriesNoCredential(t *testing.T) { client := &http.Client{Transport: &cleartextRouteTransport{cleartext: target.Host}} provider := NewWithHTTPClient("eden-key", "http://eden.example.com/v3", client, llmclient.Hooks{}) - if _, err := provider.Embeddings(context.Background(), embeddingRequest()); err != nil { - t.Fatalf("Embeddings: %v", err) + _, err = provider.ChatCompletion(context.Background(), &core.ChatRequest{ + Model: slashedModel, + Messages: []core.Message{{Role: "user", Content: "secret-prompt"}}, + }) + if err == nil { + t.Fatal("ChatCompletion succeeded against a cleartext endpoint, want the request refused") } + if hits != 0 { + t.Errorf("cleartext endpoint received %d request(s), want 0", hits) + } if gotAuth != "" { - t.Errorf("upstream saw Authorization = %q, want empty: the Eden key must never travel in cleartext", gotAuth) + t.Errorf("cleartext endpoint saw Authorization = %q, want empty", gotAuth) + } + if strings.Contains(gotBody, "secret-prompt") { + t.Errorf("cleartext endpoint received the prompt body %q; the payload must never be sent", gotBody) + } + // The refusal must name the destination without quoting the whole URL. + if !strings.Contains(err.Error(), "eden.example.com") { + t.Errorf("error %q should name the refused host", err) + } +} + +// TestCleartextLoopback_RequestStillSent pins the other side of the exemption: +// a cleartext endpoint on the local machine is allowed through and still +// authenticates, which is what keeps a local Eden-compatible proxy usable and +// what every httptest-backed test in this package relies on. +func TestCleartextLoopback_RequestStillSent(t *testing.T) { + var hits int + var gotAuth string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hits++ + gotAuth = r.Header.Get("Authorization") + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"object":"list","model":"openai/text-embedding-3-small","data":[]}`)) + })) + defer server.Close() + + provider := NewWithHTTPClient("eden-key", server.URL, server.Client(), llmclient.Hooks{}) + if _, err := provider.Embeddings(context.Background(), embeddingRequest()); err != nil { + t.Fatalf("Embeddings against a loopback endpoint: %v", err) + } + if hits != 1 { + t.Errorf("loopback endpoint received %d request(s), want 1", hits) + } + if gotAuth != "Bearer eden-key" { + t.Errorf("loopback endpoint saw Authorization = %q, want %q", gotAuth, "Bearer eden-key") + } +} + +// TestSecureTransport_RoundTrip covers the guard directly, including that an +// allowed request is handed to the underlying transport unchanged and that a +// nil base delegates to http.DefaultTransport the way net/http does. +func TestSecureTransport_RoundTrip(t *testing.T) { + tests := []struct { + name string + target string + wantCalled bool + }{ + {"https", "https://api.edenai.run/v3/models", true}, + {"loopback http", "http://127.0.0.1:9999/v3/models", true}, + {"loopback name", "http://localhost:9999/v3/models", true}, + {"public cleartext", "http://eden.example.com/v3/models", false}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + stub := &recordingRoundTripper{} + transport := &secureTransport{base: stub} + + req, err := http.NewRequest(http.MethodGet, tc.target, nil) + if err != nil { + t.Fatalf("http.NewRequest: %v", err) + } + resp, err := transport.RoundTrip(req) + if resp != nil && resp.Body != nil { + _ = resp.Body.Close() + } + + if stub.calls != 0 != tc.wantCalled { + t.Errorf("underlying transport calls = %d, wantCalled = %v", stub.calls, tc.wantCalled) + } + if tc.wantCalled && err != nil { + t.Errorf("RoundTrip(%q) = %v, want the request passed through", tc.target, err) + } + if !tc.wantCalled && err == nil { + t.Errorf("RoundTrip(%q) = nil error, want a refusal", tc.target) + } + }) } } +// TestSecureTransport_NilBaseUsesDefaultTransport asserts the nil-base fallback +// is wired, since http.DefaultClient carries a nil Transport and Eden's guard +// wraps exactly that on the NewWithHTTPClient nil path. +func TestSecureTransport_NilBaseUsesDefaultTransport(t *testing.T) { + transport := &secureTransport{} + + // A refused destination never reaches the base, so it proves the guard runs + // without needing a live server for the delegating case. + req, err := http.NewRequest(http.MethodGet, "http://eden.example.com/v3", nil) + if err != nil { + t.Fatalf("http.NewRequest: %v", err) + } + if _, err := transport.RoundTrip(req); err == nil { + t.Error("RoundTrip = nil error for a cleartext target, want a refusal") + } + + // An allowed loopback destination must reach the network through + // http.DefaultTransport rather than panicking on the nil base. + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusNoContent) + })) + defer server.Close() + + req, err = http.NewRequest(http.MethodGet, server.URL, nil) + if err != nil { + t.Fatalf("http.NewRequest: %v", err) + } + resp, err := transport.RoundTrip(req) + if err != nil { + t.Fatalf("RoundTrip through the nil base: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusNoContent { + t.Errorf("status = %d, want %d", resp.StatusCode, http.StatusNoContent) + } +} + +// recordingRoundTripper counts the requests that made it past the guard. +type recordingRoundTripper struct { + calls int +} + +func (r *recordingRoundTripper) RoundTrip(*http.Request) (*http.Response, error) { + r.calls++ + return &http.Response{ + StatusCode: http.StatusOK, + Body: http.NoBody, + Header: make(http.Header), + }, nil +} + // cleartextRouteTransport routes cleartext requests to a fixed address while // leaving the request URL's host alone, and sends everything else through base // (http.DefaultTransport when nil). It lets a test address a public-looking @@ -206,34 +349,186 @@ func TestGuardedHTTPClient_DoesNotMutateCallerClient(t *testing.T) { } } -// TestGuardedHTTPClient_PreservesCallerTransport asserts copying the client -// keeps the transport, so a caller that configured TLS roots or a proxy (and -// httptest's own client) still works. -func TestGuardedHTTPClient_PreservesCallerTransport(t *testing.T) { +// TestGuardedHTTPClient_WrapsCallerTransport asserts the guard is layered over +// the caller's transport rather than replacing it, so a caller that configured +// TLS roots or a proxy (and httptest's own client) still reaches its server. +func TestGuardedHTTPClient_WrapsCallerTransport(t *testing.T) { transport := &http.Transport{} caller := &http.Client{Transport: transport} - if got := guardedHTTPClient(caller).Transport; got != transport { - t.Errorf("Transport = %v, want the caller's transport preserved", got) + guarded, ok := guardedHTTPClient(caller).Transport.(*secureTransport) + if !ok { + t.Fatalf("Transport = %T, want *secureTransport wrapping the caller's", guardedHTTPClient(caller).Transport) + } + if guarded.base != transport { + t.Errorf("wrapped base = %v, want the caller's transport preserved", guarded.base) + } + if caller.Transport != transport { + t.Error("the caller's client was modified; the guard must go on a copy") + } +} + +// TestCheckRedirect_AllowsSameHostHTTPS asserts the one redirect shape a REST +// endpoint plausibly returns keeps working: a path change on the host the +// request was already addressed to. +func TestCheckRedirect_AllowsSameHostHTTPS(t *testing.T) { + origin := mustRequest(t, "https://api.edenai.run/v3/chat/completions") + target := mustRequest(t, "https://api.edenai.run/v3/chat/completions-moved") + + if err := checkRedirect(target, []*http.Request{origin}); err != nil { + t.Errorf("checkRedirect same-host = %v, want nil", err) + } + // Host comparison is case-insensitive, as hostnames are. + upper := mustRequest(t, "https://API.EdenAI.run/v3/chat/completions-moved") + if err := checkRedirect(upper, []*http.Request{origin}); err != nil { + t.Errorf("checkRedirect differing-case host = %v, want nil", err) } } -// TestCheckRedirect_AllowsHTTPSTargets asserts TLS redirects keep working, -// including to a different host: that is an ordinary API redirect and the -// credential stays encrypted. -func TestCheckRedirect_AllowsHTTPSTargets(t *testing.T) { +// TestCheckRedirect_RefusesCrossHostHTTPS covers the credential-exposure path +// that TLS alone does not close. Go's default policy forwards Authorization to +// a subdomain of the original host, so an upstream-chosen Location is otherwise +// enough to hand the Eden key to a neighbouring name — and following the hop +// would ship the request body there too. +func TestCheckRedirect_RefusesCrossHostHTTPS(t *testing.T) { + origin := mustRequest(t, "https://api.edenai.run/v3/chat/completions") + for _, target := range []string{ - "https://api.edenai.run/v3/chat/completions", - "https://eu.edenai.run/v3/chat/completions", + "https://evil.api.edenai.run/v3/chat/completions", // subdomain: Go would forward the token + "https://eu.edenai.run/v3/chat/completions", // sibling host + "https://attacker.example.com/v3/chat/completions", } { - req, err := http.NewRequest(http.MethodGet, target, nil) - if err != nil { - t.Fatalf("http.NewRequest: %v", err) + req := mustRequest(t, target) + err := checkRedirect(req, []*http.Request{origin}) + if err == nil { + t.Errorf("checkRedirect(%q) = nil, want a refusal", target) + continue + } + if !strings.Contains(err.Error(), "cross-host") { + t.Errorf("checkRedirect(%q) = %v, want a cross-host refusal", target, err) } - if err := checkRedirect(req, nil); err != nil { - t.Errorf("checkRedirect(%q) = %v, want nil", target, err) + } +} + +// TestCheckRedirect_ComparesAgainstOriginalHost asserts a chain cannot walk off +// the configured host one hop at a time: the check is against the request the +// caller made, not the previous hop. +func TestCheckRedirect_ComparesAgainstOriginalHost(t *testing.T) { + origin := mustRequest(t, "https://api.edenai.run/v3/models") + hop := mustRequest(t, "https://api.edenai.run/v3/models-moved") + target := mustRequest(t, "https://elsewhere.edenai.run/v3/models") + + if err := checkRedirect(target, []*http.Request{origin, hop}); err == nil { + t.Error("checkRedirect = nil for a second hop leaving the original host, want a refusal") + } +} + +// TestRedirectPolicy_PreservesCallerPolicy asserts the guard composes with the +// policy a caller's client already carried instead of replacing it. Overwriting +// the field would silently drop the caller's rules, letting an upstream Location +// reach a destination the caller had rejected. +func TestRedirectPolicy_PreservesCallerPolicy(t *testing.T) { + origin := mustRequest(t, "https://api.edenai.run/v3/models") + target := mustRequest(t, "https://api.edenai.run/v3/models-moved") + callerErr := errors.New("caller rejected this destination") + + var called int + policy := redirectPolicy(func(*http.Request, []*http.Request) error { + called++ + return callerErr + }) + + // Eden allows this same-host hop, so the caller's policy decides. + if err := policy(target, []*http.Request{origin}); !errors.Is(err, callerErr) { + t.Errorf("policy = %v, want the caller's error returned verbatim", err) + } + if called != 1 { + t.Errorf("caller policy invoked %d times, want 1", called) + } +} + +// TestRedirectPolicy_PropagatesErrUseLastResponse asserts the sentinel a caller +// uses to stop following redirects survives composition. Swallowing it would +// turn "hand me the redirect response" into "follow the redirect". +func TestRedirectPolicy_PropagatesErrUseLastResponse(t *testing.T) { + origin := mustRequest(t, "https://api.edenai.run/v3/models") + target := mustRequest(t, "https://api.edenai.run/v3/models-moved") + + policy := redirectPolicy(func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + }) + + if err := policy(target, []*http.Request{origin}); !errors.Is(err, http.ErrUseLastResponse) { + t.Errorf("policy = %v, want http.ErrUseLastResponse propagated", err) + } +} + +// TestRedirectPolicy_EdenRefusalWinsOverPermissiveCaller asserts a caller +// cannot waive Eden's guarantees: Eden's rules run first and a refusal is +// final, so a policy that approves everything still cannot allow a cleartext or +// cross-host hop. +func TestRedirectPolicy_EdenRefusalWinsOverPermissiveCaller(t *testing.T) { + origin := mustRequest(t, "https://api.edenai.run/v3/models") + + var called int + policy := redirectPolicy(func(*http.Request, []*http.Request) error { + called++ + return nil + }) + + for _, target := range []string{ + "http://api.edenai.run/v3/models", // cleartext downgrade + "https://evil.api.edenai.run/v3/models", // cross-host + } { + if err := policy(mustRequest(t, target), []*http.Request{origin}); err == nil { + t.Errorf("policy(%q) = nil, want Eden's refusal to stand", target) } } + if called != 0 { + t.Errorf("caller policy invoked %d times, want 0: Eden refuses before delegating", called) + } +} + +// TestRedirectPolicy_NilCallerAllowsEdenApprovedHop asserts the common case — +// a client with no policy of its own — still follows an Eden-approved redirect. +func TestRedirectPolicy_NilCallerAllowsEdenApprovedHop(t *testing.T) { + origin := mustRequest(t, "https://api.edenai.run/v3/models") + target := mustRequest(t, "https://api.edenai.run/v3/models-moved") + + if err := redirectPolicy(nil)(target, []*http.Request{origin}); err != nil { + t.Errorf("redirectPolicy(nil) = %v, want nil for a same-host TLS hop", err) + } +} + +// TestGuardedHTTPClient_ComposesCallerRedirectPolicy asserts the composition is +// actually wired by the constructor, not only available as a helper. +func TestGuardedHTTPClient_ComposesCallerRedirectPolicy(t *testing.T) { + var called bool + caller := &http.Client{CheckRedirect: func(*http.Request, []*http.Request) error { + called = true + return nil + }} + + guarded := guardedHTTPClient(caller) + origin := mustRequest(t, "https://api.edenai.run/v3/models") + target := mustRequest(t, "https://api.edenai.run/v3/models-moved") + + if err := guarded.CheckRedirect(target, []*http.Request{origin}); err != nil { + t.Fatalf("CheckRedirect = %v, want nil", err) + } + if !called { + t.Error("the caller's redirect policy was not invoked; the guard must compose, not replace") + } +} + +// mustRequest builds a GET request for a URL a test controls. +func mustRequest(t *testing.T, rawURL string) *http.Request { + t.Helper() + req, err := http.NewRequest(http.MethodGet, rawURL, nil) + if err != nil { + t.Fatalf("http.NewRequest(%q): %v", rawURL, err) + } + return req } // TestCheckRedirect_RefusesSchemeDowngrade is the core of the redirect @@ -404,6 +699,72 @@ func TestRedirect_HTTPSToHTTPSStillFollowed(t *testing.T) { } } +// TestRedirect_HTTPSSubdomainDoesNotForwardCredential is the end-to-end form of +// the cross-host rule, against a real TLS server and a real redirect. +// +// Both hops are served by the same httptest instance, addressed through two +// different hostnames so Go sees a genuine subdomain redirect — the case where +// its default policy forwards Authorization. Neither the credential nor the +// request may reach the second name. +func TestRedirect_HTTPSSubdomainDoesNotForwardCredential(t *testing.T) { + var secondHopHits int + var secondHopAuth string + + mux := http.NewServeMux() + mux.HandleFunc("/v3/embeddings", func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, "https://evil.eden.test/v3/stolen", http.StatusFound) + }) + mux.HandleFunc("/v3/stolen", func(w http.ResponseWriter, r *http.Request) { + secondHopHits++ + secondHopAuth = r.Header.Get("Authorization") + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"object":"list","data":[]}`)) + }) + server := httptest.NewTLSServer(mux) + defer server.Close() + + target, err := url.Parse(server.URL) + if err != nil { + t.Fatalf("url.Parse: %v", err) + } + client := server.Client() + // Route both hostnames to the one test server; the TLS config from + // server.Client() already trusts its certificate. + client.Transport = &hostPinnedTransport{base: client.Transport, addr: target.Host} + + provider := NewWithHTTPClient("eden-key", "https://eden.test/v3", client, llmclient.Hooks{}) + _, err = provider.Embeddings(context.Background(), embeddingRequest()) + if err == nil { + t.Fatal("Embeddings followed an HTTPS subdomain redirect, want it refused") + } + + if secondHopHits != 0 { + t.Errorf("subdomain endpoint received %d request(s), want 0", secondHopHits) + } + if secondHopAuth != "" { + t.Errorf("subdomain endpoint saw Authorization = %q, want empty: the Eden key must not follow an upstream-chosen host", secondHopAuth) + } +} + +// hostPinnedTransport dials one fixed address whatever hostname the request +// carries, so a test can exercise multi-host redirect rules against a single +// server without DNS. The request URL's host is left intact, which is what the +// redirect policy inspects. +type hostPinnedTransport struct { + base http.RoundTripper + addr string +} + +func (t *hostPinnedTransport) RoundTrip(req *http.Request) (*http.Response, error) { + routed := req.Clone(req.Context()) + routed.URL.Host = t.addr + base := t.base + if base == nil { + base = http.DefaultTransport + } + return base.RoundTrip(routed) +} + // TestDefaultBaseURLIsHTTPS guards the default every deployment uses. The // credential guard keys off the request scheme, so a default that regressed to // http:// would silently withhold the API key on ordinary traffic. @@ -459,8 +820,14 @@ func TestGuardedHTTPClient_PreservesDefaultClientSemantics(t *testing.T) { if guarded.CheckRedirect == nil { t.Error("CheckRedirect = nil, want Eden's redirect guard on the copy") } - if guarded.Transport != http.DefaultClient.Transport { - t.Errorf("Transport = %v, want http.DefaultClient's transport preserved", guarded.Transport) + wrapped, ok := guarded.Transport.(*secureTransport) + if !ok { + t.Fatalf("Transport = %T, want *secureTransport", guarded.Transport) + } + // http.DefaultClient leaves Transport nil, meaning http.DefaultTransport; + // the wrapper preserves that by delegating to it when its base is nil. + if wrapped.base != http.DefaultClient.Transport { + t.Errorf("wrapped base = %v, want http.DefaultClient's transport (nil)", wrapped.base) } if guarded.Timeout != http.DefaultClient.Timeout { t.Errorf("Timeout = %v, want http.DefaultClient's %v, not the gateway client's", guarded.Timeout, http.DefaultClient.Timeout) From d02e72520b90f684bf4359ab442bef074e169794 Mon Sep 17 00:00:00 2001 From: NaDdjg Date: Fri, 11 Sep 2026 12:50:08 +0100 Subject: [PATCH 4/6] fix: re-validate redirect target after caller policy --- internal/providers/edenai/transport.go | 16 ++- internal/providers/edenai/transport_test.go | 123 ++++++++++++++++++++ 2 files changed, 138 insertions(+), 1 deletion(-) diff --git a/internal/providers/edenai/transport.go b/internal/providers/edenai/transport.go index b0ad95ce0..2f4ee4616 100644 --- a/internal/providers/edenai/transport.go +++ b/internal/providers/edenai/transport.go @@ -92,6 +92,14 @@ func guardedHTTPClient(base *http.Client) *http.Client { // an overwritten CheckRedirect never runs and the redirect is followed anyway), // which is why the guard composes here the same way secureTransport wraps the // caller's transport rather than replacing it. +// +// Eden's rules are checked twice, before and after the callback, because +// net/http hands CheckRedirect the very *http.Request it is about to send and +// honors any change made to it (verified: a callback that rewrites req.URL +// redirects the request to the rewritten target). Validating only up front +// would leave the checks describing a URL that is no longer the one going out, +// so a callback that rewrites the host -- a region or proxy rewrite as much as +// anything hostile -- would carry the request past them. func redirectPolicy(caller func(*http.Request, []*http.Request) error) func(*http.Request, []*http.Request) error { return func(req *http.Request, via []*http.Request) error { if err := checkRedirect(req, via); err != nil { @@ -100,7 +108,13 @@ func redirectPolicy(caller func(*http.Request, []*http.Request) error) func(*htt if caller == nil { return nil } - return caller(req, via) + if err := caller(req, via); err != nil { + // Returned unchanged: the caller may be signalling + // http.ErrUseLastResponse, which net/http reads as "stop here and + // hand back the redirect response" rather than as a failure. + return err + } + return checkRedirect(req, via) } } diff --git a/internal/providers/edenai/transport_test.go b/internal/providers/edenai/transport_test.go index 3de5fd55a..cdd2fa7d2 100644 --- a/internal/providers/edenai/transport_test.go +++ b/internal/providers/edenai/transport_test.go @@ -463,6 +463,129 @@ func TestRedirectPolicy_PropagatesErrUseLastResponse(t *testing.T) { } } +// TestRedirectPolicy_RecheckAfterCallerMutation is the check-then-mutate case. +// +// net/http passes CheckRedirect the very request it is about to send and +// honors edits to it, so a callback that approves a hop and rewrites req.URL +// on the way out would move the request past checks that already ran. Eden's +// rules are therefore re-applied to whatever the callback leaves behind. +func TestRedirectPolicy_RecheckAfterCallerMutation(t *testing.T) { + tests := []struct { + name string + mutate func(*http.Request) + wantErr string + }{ + { + name: "rewritten to another https host", + mutate: func(req *http.Request) { + req.URL.Host = "attacker.example.com" + }, + wantErr: "cross-host", + }, + { + name: "rewritten to a subdomain of the original host", + mutate: func(req *http.Request) { + req.URL.Host = "evil.api.edenai.run" + }, + wantErr: "cross-host", + }, + { + name: "downgraded to cleartext", + mutate: func(req *http.Request) { + req.URL.Scheme = "http" + }, + wantErr: "cleartext", + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + origin := mustRequest(t, "https://api.edenai.run/v3/models") + // A target Eden approves on its own, so only the mutation can + // make it fail. + target := mustRequest(t, "https://api.edenai.run/v3/models-moved") + + policy := redirectPolicy(func(req *http.Request, _ []*http.Request) error { + tc.mutate(req) + return nil + }) + + err := policy(target, []*http.Request{origin}) + if err == nil { + t.Fatalf("policy = nil, want a refusal after the callback rewrote the target to %s", target.URL) + } + if !strings.Contains(err.Error(), tc.wantErr) { + t.Errorf("policy = %v, want an error mentioning %q", err, tc.wantErr) + } + }) + } +} + +// TestRedirectPolicy_AllowsHarmlessCallerMutation asserts the re-check is not +// blanket paranoia: a callback that rewrites only the path, staying on the +// configured host over TLS, is still allowed through. +func TestRedirectPolicy_AllowsHarmlessCallerMutation(t *testing.T) { + origin := mustRequest(t, "https://api.edenai.run/v3/models") + target := mustRequest(t, "https://api.edenai.run/v3/models-moved") + + policy := redirectPolicy(func(req *http.Request, _ []*http.Request) error { + req.URL.Path = "/v3/models-rewritten" + return nil + }) + + if err := policy(target, []*http.Request{origin}); err != nil { + t.Errorf("policy = %v, want nil: a same-host path rewrite is fine", err) + } +} + +// TestRedirectPolicy_CallerErrorsPropagateUnchanged asserts every non-nil +// result from the callback reaches net/http exactly as returned, so the +// re-check cannot convert a caller's decision into a different outcome. +// http.ErrUseLastResponse matters most: net/http reads it as "stop here and +// return the redirect response", not as a failure. +func TestRedirectPolicy_CallerErrorsPropagateUnchanged(t *testing.T) { + sentinel := errors.New("caller rejected this destination") + + tests := []struct { + name string + err error + }{ + {"use last response", http.ErrUseLastResponse}, + {"caller error", sentinel}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + origin := mustRequest(t, "https://api.edenai.run/v3/models") + target := mustRequest(t, "https://api.edenai.run/v3/models-moved") + + policy := redirectPolicy(func(*http.Request, []*http.Request) error { + return tc.err + }) + if err := policy(target, []*http.Request{origin}); !errors.Is(err, tc.err) { + t.Errorf("policy = %v, want %v returned unchanged", err, tc.err) + } + }) + } +} + +// TestRedirectPolicy_CallerErrorSurvivesAMutation asserts the error path wins +// over the re-check: a callback that both rewrites the target and returns a +// sentinel must have its sentinel propagated, not replaced by Eden's refusal. +func TestRedirectPolicy_CallerErrorSurvivesAMutation(t *testing.T) { + origin := mustRequest(t, "https://api.edenai.run/v3/models") + target := mustRequest(t, "https://api.edenai.run/v3/models-moved") + + policy := redirectPolicy(func(req *http.Request, _ []*http.Request) error { + req.URL.Host = "attacker.example.com" + return http.ErrUseLastResponse + }) + + if err := policy(target, []*http.Request{origin}); !errors.Is(err, http.ErrUseLastResponse) { + t.Errorf("policy = %v, want http.ErrUseLastResponse propagated unchanged", err) + } +} + // TestRedirectPolicy_EdenRefusalWinsOverPermissiveCaller asserts a caller // cannot waive Eden's guarantees: Eden's rules run first and a refusal is // final, so a policy that approves everything still cannot allow a cleartext or From de22882e2212994f5df328b6d42bbdde2b604275 Mon Sep 17 00:00:00 2001 From: Nada Farah Djedjig Date: Wed, 16 Sep 2026 09:56:30 +0100 Subject: [PATCH 5/6] Updated the description for the EDENAI_API_KEY. Updated the description for the EDENAI_API_KEY. --- docs/advanced/configuration.mdx | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/advanced/configuration.mdx b/docs/advanced/configuration.mdx index 521eb2cbe..8cb7fb646 100644 --- a/docs/advanced/configuration.mdx +++ b/docs/advanced/configuration.mdx @@ -321,7 +321,7 @@ Set these to automatically register providers. No YAML configuration required. | `DEEPSEEK_API_KEY` | DeepSeek | | `OPENROUTER_API_KEY` | OpenRouter | | `KILO_API_KEY` | Kilo AI Gateway | -| `EDENAI_API_KEY` | Eden AI (`EDENAI_BASE_URL` optional) | +| `EDENAI_API_KEY` | Eden AI | | `ZAI_API_KEY` | Z.ai | | `XAI_API_KEY` | xAI (Grok) | | `GROQ_API_KEY` | Groq | From 4ecc42a616fa004368602a214a22cb9ddaafac74 Mon Sep 17 00:00:00 2001 From: NaDdjg Date: Wed, 30 Sep 2026 12:40:43 +0100 Subject: [PATCH 6/6] fix: resolve CI issues for Eden AI provider --- .env.template | 7 +- config/config.go | 2 +- docs/features/passthrough-api.mdx | 2 +- internal/providers/config_test.go | 24 +- .../providers/edenai/capabilities_test.go | 59 +-- internal/providers/edenai/edenai.go | 52 +- internal/providers/edenai/edenai_test.go | 298 ++++------- .../providers/edenai/embeddings_cost_test.go | 128 ++--- internal/providers/edenai/models_test.go | 179 +++---- .../providers/edenai/newtestprovider_test.go | 16 + .../edenai/passthrough_semantics_test.go | 34 +- internal/providers/edenai/response_test.go | 102 ++-- internal/providers/edenai/transport.go | 5 +- internal/providers/edenai/transport_test.go | 470 ++++++------------ .../registry_provider_pricing_test.go | 59 +-- internal/server/passthrough_support_test.go | 37 +- internal/usage/cost_test.go | 29 +- internal/usage/extractor_test.go | 82 +-- internal/usage/stream_observer_test.go | 69 +-- run/lifecycle_test.go | 13 +- 20 files changed, 531 insertions(+), 1136 deletions(-) create mode 100644 internal/providers/edenai/newtestprovider_test.go diff --git a/.env.template b/.env.template index a131c5cb6..375f6373b 100644 --- a/.env.template +++ b/.env.template @@ -106,12 +106,11 @@ # Allow optional /p/{provider}/v1/... passthrough aliases while keeping /p/{provider}/... canonical (default: true) # ALLOW_PASSTHROUGH_V1_ALIAS=true - -# Comma-separated list of provider types enabled for /p/{provider}/... passthrough (default: openai,anthropic,openrouter,kilo,zai,sglang,vllm,llamacpp,llmd,deepseek,jev,edenai) +# Comma-separated list of provider types enabled for /p/{provider}/... passthrough (default: openai,anthropic,openrouter,kilo,zai,sglang,vllm,llamacpp,llmd,deepseek,edenai,jev) # Cohere and audio.cpp native passthrough are opt-in; add cohere or audiocpp when # those routes are needed. audio.cpp's native surface includes model management # and server-local file paths, so enable it only for trusted callers. -# ENABLED_PASSTHROUGH_PROVIDERS=openai,anthropic,cohere,openrouter,kilo,zai,sglang,vllm,llamacpp,llmd,audiocpp,deepseek,hetzner,jev,edenai +# ENABLED_PASSTHROUGH_PROVIDERS=openai,anthropic,cohere,openrouter,kilo,zai,sglang,vllm,llamacpp,llmd,audiocpp,deepseek,hetzner,edenai,jev # Enable the realtime (speech-to-speech) endpoints (default: true): the /v1/realtime # websocket (and /p/{provider}/v1/realtime passthrough upgrade), the WebRTC SDP @@ -730,7 +729,7 @@ # Eden AI (default base URL: https://api.edenai.run/v3) # Multi-provider gateway. Model IDs use provider/model and pass through unchanged. # EDENAI_BASE_URL is optional: the default above is used when it is unset. -# EDENAI_API_KEY= +# EDENAI_API_KEY= # EDENAI_BASE_URL=https://api.edenai.run/v3 # Optional configured model list; see CONFIGURED_PROVIDER_MODELS_MODE below # EDENAI_MODELS=openai/gpt-4,anthropic/claude-sonnet-latest diff --git a/config/config.go b/config/config.go index bd52ca1ee..52e0b8fb2 100644 --- a/config/config.go +++ b/config/config.go @@ -128,7 +128,7 @@ func buildDefaultConfig() *Config { "llamacpp", "llmd", "deepseek", - "edenai", + "edenai", "jev", }, }, diff --git a/docs/features/passthrough-api.mdx b/docs/features/passthrough-api.mdx index 71e9d3b41..485d8ce16 100644 --- a/docs/features/passthrough-api.mdx +++ b/docs/features/passthrough-api.mdx @@ -136,7 +136,7 @@ from passthrough requests before forwarding them upstream. Passthrough is intentionally narrow while the API is in beta. -- `openai`, `anthropic`, `openrouter`, `kilo`, `zai`, `sglang`, `vllm`, `llamacpp`, `llmd`, `deepseek`, `edenai` and `jev` +- `openai`, `anthropic`, `openrouter`, `kilo`, `zai`, `sglang`, `vllm`, `llamacpp`, `llmd`, `deepseek`, `edenai`, and `jev` are enabled by default. - Chutes supports passthrough but requires explicit operator opt-in because passthrough can forward provider-native routes that do not identify a model. diff --git a/internal/providers/config_test.go b/internal/providers/config_test.go index 436c5fc19..bd6ff7e3c 100644 --- a/internal/providers/config_test.go +++ b/internal/providers/config_test.go @@ -1481,18 +1481,10 @@ func TestBuildProviderConfig_EdenAI_ResolvesBaseURL(t *testing.T) { got := applyProviderEnvVars(map[string]config.RawProviderConfig{}, testDiscoveryConfigs) p, exists := got["edenai"] - if !exists { - t.Fatal("edenai not discovered by config parser") - } - if p.Type != "edenai" { - t.Errorf("Type = %q, want edenai", p.Type) - } - if p.APIKey != "edenai-test-key" { - t.Errorf("APIKey = %q, want edenai-test-key", p.APIKey) - } - if p.BaseURL != "https://api.edenai.run/v3" { - t.Errorf("BaseURL = %q, want https://api.edenai.run/v3", p.BaseURL) - } + require.True(t, exists, "edenai not discovered by config parser") + assert.Equal(t, "edenai", p.Type) + assert.Equal(t, "edenai-test-key", p.APIKey) + assert.Equal(t, "https://api.edenai.run/v3", p.BaseURL) } // TestBuildProviderConfig_EdenAI_BaseURLOverride asserts EDENAI_BASE_URL wins @@ -1505,12 +1497,8 @@ func TestBuildProviderConfig_EdenAI_BaseURLOverride(t *testing.T) { got := applyProviderEnvVars(map[string]config.RawProviderConfig{}, testDiscoveryConfigs) p, exists := got["edenai"] - if !exists { - t.Fatal("edenai not discovered by config parser") - } - if p.BaseURL != "https://eden.internal.example/v3" { - t.Errorf("BaseURL = %q, want https://eden.internal.example/v3", p.BaseURL) - } + require.True(t, exists, "edenai not discovered by config parser") + assert.Equal(t, "https://eden.internal.example/v3", p.BaseURL) } func TestApplyProviderEnvVars_ModelFilter(t *testing.T) { diff --git a/internal/providers/edenai/capabilities_test.go b/internal/providers/edenai/capabilities_test.go index a96311d4f..f7b3fb708 100644 --- a/internal/providers/edenai/capabilities_test.go +++ b/internal/providers/edenai/capabilities_test.go @@ -2,6 +2,8 @@ package edenai import ( "testing" + + "github.com/stretchr/testify/assert" ) // TestCapabilities_LiveSchemaOnly pins the one capability schema Eden actually @@ -34,25 +36,17 @@ func TestCapabilities_LiveSchemaOnly(t *testing.T) { // The supports_ prefix is stripped so the names read the way other // providers report them. for _, name := range []string{"reasoning", "web_search", "function_calling", "prompt_caching"} { - if !capabilities[name] { - t.Errorf("capability %q = false, want true", name) - } + assert.True(t, capabilities[name], "capability %q = false, want true", name) } // An image input modality is reported as vision, the gateway's name for it. - if !capabilities["vision"] { - t.Error(`capability "vision" = false, want true for an image input modality`) - } + assert.True(t, capabilities["vision"], `capability "vision" = false, want true for an image input modality`) // A flag Eden reports as false must not be advertised at all. for _, name := range []string{"tool_choice", "computer_use"} { - if _, present := capabilities[name]; present { - t.Errorf("capability %q is present, want absent: Eden reported it false", name) - } + assert.NotContains(t, capabilities, name, "capability %q is present, want absent: Eden reported it false", name) } // The prefixed keys must not leak through under their raw Eden names. for _, name := range []string{"supports_reasoning", "supports_web_search", "input_modalities", "output_modalities"} { - if _, present := capabilities[name]; present { - t.Errorf("capability %q is present, want the Eden key name not to leak", name) - } + assert.NotContains(t, capabilities, name, "capability %q is present, want the Eden key name not to leak", name) } } @@ -86,9 +80,7 @@ func TestCapabilities_RareSupportsFlagsFlowThrough(t *testing.T) { "structured_output", "responses_api", } { - if !model.Metadata.Capabilities[name] { - t.Errorf("capability %q = false, want true", name) - } + assert.True(t, model.Metadata.Capabilities[name], "capability %q = false, want true", name) } } @@ -113,12 +105,8 @@ func TestCapabilities_ObjectValuedReasoningIsNotAFlag(t *testing.T) { }]}`) capabilities := model.Metadata.Capabilities - if _, present := capabilities["reasoning"]; present { - t.Error(`capability "reasoning" is present, but Eden reported supports_reasoning: false; the object-valued "reasoning" member must not be read as a flag`) - } - if !capabilities["web_search"] { - t.Error(`capability "web_search" = false, want true: a sibling object member must not stop the real flags being read`) - } + assert.NotContains(t, capabilities, "reasoning", `capability "reasoning" is present, but Eden reported supports_reasoning: false; the object-valued "reasoning" member must not be read as a flag`) + assert.True(t, capabilities["web_search"], `capability "web_search" = false, want true: a sibling object member must not stop the real flags being read`) } // TestCapabilities_NonBooleanSupportsValuesIgnored asserts a supports_* member @@ -138,13 +126,9 @@ func TestCapabilities_NonBooleanSupportsValuesIgnored(t *testing.T) { capabilities := model.Metadata.Capabilities for _, name := range []string{"web_search", "tool_choice", "prompt_caching"} { - if _, present := capabilities[name]; present { - t.Errorf("capability %q is present, want absent: Eden did not report a boolean", name) - } - } - if !capabilities["reasoning"] { - t.Error(`capability "reasoning" = false, want true`) + assert.NotContains(t, capabilities, name, "capability %q is present, want absent: Eden did not report a boolean", name) } + assert.True(t, capabilities["reasoning"], `capability "reasoning" = false, want true`) } // TestCapabilities_BareSupportsPrefixIgnored asserts a key that is exactly the @@ -155,9 +139,7 @@ func TestCapabilities_BareSupportsPrefixIgnored(t *testing.T) { "capabilities": {"output_modalities": ["text"], "supports_": true} }]}`) - if _, present := model.Metadata.Capabilities[""]; present { - t.Error(`an empty-named capability was recorded for the bare "supports_" key`) - } + assert.NotContains(t, model.Metadata.Capabilities, "", `an empty-named capability was recorded for the bare "supports_" key`) } // TestCapabilities_InputModalityMapping pins the input modalities the live @@ -175,14 +157,10 @@ func TestCapabilities_InputModalityMapping(t *testing.T) { capabilities := model.Metadata.Capabilities for _, name := range []string{"vision", "video", "audio"} { - if !capabilities[name] { - t.Errorf("capability %q = false, want true", name) - } + assert.True(t, capabilities[name], "capability %q = false, want true", name) } for _, name := range []string{"file", "text"} { - if _, present := capabilities[name]; present { - t.Errorf("capability %q is present; modalities with no gateway capability must be skipped", name) - } + assert.NotContains(t, capabilities, name, "capability %q is present; modalities with no gateway capability must be skipped", name) } } @@ -202,11 +180,6 @@ func TestCapabilities_VideoInputIsNotVideoOutput(t *testing.T) { } }]}`) - if !model.Metadata.Capabilities["video"] { - t.Error(`capability "video" = false, want true for a video input modality`) - } - modes := model.Metadata.Modes - if len(modes) != 2 || modes[0] != "chat" || modes[1] != "responses" { - t.Errorf("modes = %v, want [chat responses]: video input must not produce a video mode", modes) - } + assert.True(t, model.Metadata.Capabilities["video"], `capability "video" = false, want true for a video input modality`) + assert.Equal(t, []string{"chat", "responses"}, model.Metadata.Modes, "video input must not produce a video mode") } diff --git a/internal/providers/edenai/edenai.go b/internal/providers/edenai/edenai.go index 99829bf58..7d4e1b309 100644 --- a/internal/providers/edenai/edenai.go +++ b/internal/providers/edenai/edenai.go @@ -62,49 +62,27 @@ var ( _ core.PassthroughProvider = (*Provider)(nil) ) -// New creates a new Eden AI provider. NewCompatibleProvider takes its -// transport from CompatibleProviderConfig.HTTPClient, so the guarded client -// compatibleConfig installs is the one that reaches the network. +// New creates a new Eden AI provider. +// +// The transport comes from opts.HTTPClient — the factory sets it when an +// outbound proxy applies, and tests point it at their own server — or, when +// that is nil, from the gateway default client. Either way it is wrapped by +// guardedHTTPClient (see compatibleConfig), which is what keeps every Eden +// request off a cleartext connection. NewCompatibleProvider prefers +// CompatibleProviderConfig.HTTPClient over opts.HTTPClient, so the guarded +// client has to be handed over through the config; leaving the caller's client +// on opts alone would send Eden requests through it unguarded. func New(cfg providers.ProviderConfig, opts providers.ProviderOptions) core.Provider { return &Provider{compat: openai.NewCompatibleProvider(cfg.APIKey, opts, compatibleConfig( providers.ResolveBaseURL(cfg.BaseURL, defaultBaseURL), - nil, + opts.HTTPClient, ))} } -// NewWithHTTPClient creates a new Eden AI provider with a custom HTTP client. -// If httpClient is nil, http.DefaultClient is used. -// -// The nil default is applied here rather than left to -// NewCompatibleProviderWithHTTPClient so it stays the same default every other -// chat-compatible provider gets from that helper. Eden always hands it a -// non-nil client, so the helper's own nil branch is never reached. -// -// Either way the client carries Eden's redirect guard (see guardedHTTPClient), -// installed on a copy so a shared client -- http.DefaultClient above all -- is -// never modified. -// -// Unlike NewCompatibleProvider, NewCompatibleProviderWithHTTPClient takes its -// transport from the positional argument and ignores -// CompatibleProviderConfig.HTTPClient entirely, so the guarded client has to be -// handed over there as well. Passing the raw caller client here would leave -// this construction path — and only this one — following credential-leaking -// redirects. -// -// The signature matches every other chat-compatible provider on main: -// (apiKey, baseURL, httpClient, hooks). -func NewWithHTTPClient(apiKey string, baseURL string, httpClient *http.Client, hooks llmclient.Hooks) *Provider { - if httpClient == nil { - httpClient = http.DefaultClient - } - cfg := compatibleConfig(providers.ResolveBaseURL(baseURL, defaultBaseURL), httpClient) - return &Provider{compat: openai.NewCompatibleProviderWithHTTPClient(apiKey, cfg.HTTPClient, hooks, cfg)} -} - // compatibleConfig returns the shared OpenAI-compatible transport settings for -// Eden AI. httpClient is the caller-supplied client, or nil to build the -// gateway default; either way it is wrapped by guardedHTTPClient, so both -// constructors get the same redirect policy from one place. +// Eden AI. httpClient is the caller-supplied client (opts.HTTPClient), or nil +// to build the gateway default; either way it is wrapped by guardedHTTPClient, +// so every provider instance gets the same redirect policy from one place. func compatibleConfig(baseURL string, httpClient *http.Client) openai.CompatibleProviderConfig { return openai.CompatibleProviderConfig{ ProviderName: providerType, @@ -169,7 +147,7 @@ func (p *Provider) StreamChatCompletion(ctx context.Context, req *core.ChatReque // completions. Eden's native /responses route is a different API and is // deliberately never called. func (p *Provider) Responses(ctx context.Context, req *core.ResponsesRequest) (*core.ResponsesResponse, error) { - return providers.ResponsesViaChat(ctx, p, req) + return providers.ResponsesViaChat(ctx, p, req, providerType) } // StreamResponses translates a streaming Responses request through Eden chat diff --git a/internal/providers/edenai/edenai_test.go b/internal/providers/edenai/edenai_test.go index 548ebc9f7..4f7a081ec 100644 --- a/internal/providers/edenai/edenai_test.go +++ b/internal/providers/edenai/edenai_test.go @@ -12,6 +12,9 @@ import ( "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) // slashedModel is an Eden AI model ID in provider/model notation. Every @@ -24,29 +27,19 @@ const slashedModel = "openai/gpt-4" func TestNew_ReturnsProvider(t *testing.T) { provider := New(providers.ProviderConfig{APIKey: "test-api-key"}, providers.ProviderOptions{}) - if provider == nil { - t.Fatal("provider should not be nil") - } + require.NotNil(t, provider, "provider should not be nil") concrete, ok := provider.(*Provider) - if !ok { - t.Fatalf("New() returned %T, want *edenai.Provider", provider) - } - if concrete.compat == nil { - t.Error("composed CompatibleProvider should not be nil") - } + require.True(t, ok, "New() returned %T, want *edenai.Provider", provider) + assert.NotNil(t, concrete.compat, "composed CompatibleProvider should not be nil") } // TestNew_DefaultsBaseURL asserts that a config without an explicit base URL // falls back to Eden's public endpoint rather than an empty target. func TestNew_DefaultsBaseURL(t *testing.T) { provider, ok := New(providers.ProviderConfig{APIKey: "test-api-key"}, providers.ProviderOptions{}).(*Provider) - if !ok { - t.Fatal("New() did not return *edenai.Provider") - } - if got := provider.GetBaseURL(); got != defaultBaseURL { - t.Errorf("GetBaseURL() = %q, want %q", got, defaultBaseURL) - } + require.True(t, ok, "New() did not return *edenai.Provider") + assert.Equal(t, defaultBaseURL, provider.GetBaseURL()) } // TestNew_HonoursConfiguredBaseURL asserts that EDENAI_BASE_URL (surfaced here @@ -54,53 +47,27 @@ func TestNew_DefaultsBaseURL(t *testing.T) { func TestNew_HonoursConfiguredBaseURL(t *testing.T) { const custom = "https://eden.internal.example/v3" provider, ok := New(providers.ProviderConfig{APIKey: "k", BaseURL: custom}, providers.ProviderOptions{}).(*Provider) - if !ok { - t.Fatal("New() did not return *edenai.Provider") - } - if got := provider.GetBaseURL(); got != custom { - t.Errorf("GetBaseURL() = %q, want %q", got, custom) - } + require.True(t, ok, "New() did not return *edenai.Provider") + assert.Equal(t, custom, provider.GetBaseURL()) } -// TestNewWithHTTPClient_ReturnsProvider asserts the explicit HTTP-client constructor -// returns a valid Provider. -func TestNewWithHTTPClient_ReturnsProvider(t *testing.T) { - provider := NewWithHTTPClient("test-api-key", "http://example.invalid", &http.Client{}, llmclient.Hooks{}) - - if provider == nil { - t.Fatal("provider should not be nil") - } - if provider.compat == nil { - t.Error("composed CompatibleProvider should not be nil") - } -} - -// TestNewWithHTTPClient_NilHTTPClientDoesNotPanic asserts that passing nil for the -// HTTP client falls back to http.DefaultClient without panicking. -func TestNewWithHTTPClient_NilHTTPClientDoesNotPanic(t *testing.T) { - defer func() { - if r := recover(); r != nil { - t.Fatalf("NewWithHTTPClient(nil, ...) panicked: %v", r) - } - }() - provider := NewWithHTTPClient("test-api-key", "http://example.invalid", nil, llmclient.Hooks{}) - if provider == nil { - t.Fatal("provider should not be nil") - } -} - -// TestNewWithHTTPClient_ZeroHooksDoesNotPanic asserts that the hooks argument can be -// an empty struct (no hooks registered) without panicking. -func TestNewWithHTTPClient_ZeroHooksDoesNotPanic(t *testing.T) { - defer func() { - if r := recover(); r != nil { - t.Fatalf("NewWithHTTPClient(..., llmclient.Hooks{}) panicked: %v", r) - } - }() - provider := NewWithHTTPClient("test-api-key", "http://example.invalid", &http.Client{}, llmclient.Hooks{}) - if provider == nil { - t.Fatal("provider should not be nil") - } +// TestChatCompatibleContract checks the contract every provider built on the +// shared OpenAI-compatible adapter shares: registration metadata, constructor +// safety with a nil client and zero hooks, the injected transport being the +// one requests travel through, and chat, streaming, model listing, Responses, +// and embeddings reaching the expected upstream paths with the bearer token +// attached. Responses are expected translated through chat completions, since +// Eden's own /responses route is a different API (see the package comment). +func TestChatCompatibleContract(t *testing.T) { + providertest.AssertChatCompatible(t, providertest.ChatCompatible{ + Registration: Registration, + Type: providerType, + DefaultBaseURL: defaultBaseURL, + New: func(apiKey, baseURL string, client *http.Client, hooks llmclient.Hooks) core.Provider { + return newTestProvider(apiKey, baseURL, client, hooks) + }, + Embeddings: true, + }) } // TestRegistration_TypeAndDiscovery asserts the Registration struct exposes the @@ -109,25 +76,12 @@ func TestNewWithHTTPClient_ZeroHooksDoesNotPanic(t *testing.T) { // EDENAI_BASE_URL) and the provider gate in internal/usage, so it is asserted // exactly. func TestRegistration_TypeAndDiscovery(t *testing.T) { - if Registration.Type != "edenai" { - t.Errorf("Registration.Type = %q, want %q", Registration.Type, "edenai") - } - if Registration.New == nil { - t.Error("Registration.New should not be nil") - } - want := "https://api.edenai.run/v3" - if Registration.Discovery.DefaultBaseURL != want { - t.Errorf("Registration.Discovery.DefaultBaseURL = %q, want %q", Registration.Discovery.DefaultBaseURL, want) - } - if Registration.PassthroughSemanticEnricher == nil { - t.Error("Registration.PassthroughSemanticEnricher should not be nil") - } - if Registration.Discovery.RequireBaseURL { - t.Error("Discovery.RequireBaseURL should be false: Eden has a public default endpoint") - } - if Registration.Discovery.AllowAPIKeyless { - t.Error("Discovery.AllowAPIKeyless should be false: Eden always requires an API key") - } + assert.Equal(t, "edenai", Registration.Type) + assert.NotNil(t, Registration.New, "Registration.New should not be nil") + assert.Equal(t, "https://api.edenai.run/v3", Registration.Discovery.DefaultBaseURL) + assert.NotNil(t, Registration.PassthroughSemanticEnricher, "Registration.PassthroughSemanticEnricher should not be nil") + assert.False(t, Registration.Discovery.RequireBaseURL, "Discovery.RequireBaseURL should be false: Eden has a public default endpoint") + assert.False(t, Registration.Discovery.AllowAPIKeyless, "Discovery.AllowAPIKeyless should be false: Eden always requires an API key") } // TestProvider_ImplementsCoreProvider is a compile-time check that *Provider @@ -163,29 +117,18 @@ func TestChatCompletion_UsesBearerAuthAndForwardsModel(t *testing.T) { })) defer server.Close() - provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + provider := newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ Model: slashedModel, Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotAuth != "Bearer edenai-key" { - t.Fatalf("authorization = %q, want Bearer edenai-key", gotAuth) - } - if gotBody["model"] != slashedModel { - t.Fatalf("request model = %#v, want %q (provider/model IDs must pass through unchanged)", gotBody["model"], slashedModel) - } - if resp.Model != slashedModel { - t.Fatalf("response model = %q, want %q", resp.Model, slashedModel) - } - if len(resp.Choices) != 1 || resp.Choices[0].Message.Content != "hello" { - t.Fatalf("unexpected response: %+v", resp) - } + require.NoError(t, err) + require.Equal(t, "/chat/completions", gotPath) + require.Equal(t, "Bearer edenai-key", gotAuth) + require.Equal(t, slashedModel, gotBody["model"], "provider/model IDs must pass through unchanged") + require.Equal(t, slashedModel, resp.Model) + require.Len(t, resp.Choices, 1, "unexpected response: %+v", resp) + require.Equal(t, "hello", resp.Choices[0].Message.Content, "unexpected response: %+v", resp) } // TestChatCompletion_ForwardsEdenExtraFields asserts that Eden-only request @@ -208,23 +151,18 @@ func TestChatCompletion_ForwardsEdenExtraFields(t *testing.T) { var req core.ChatRequest raw := `{"model":"openai/gpt-4","messages":[{"role":"user","content":"hi"}],` + `"fallbacks":["anthropic/claude-sonnet-latest"],"routing":{"strategy":"cost"}}` - if err := json.Unmarshal([]byte(raw), &req); err != nil { - t.Fatalf("Unmarshal() error = %v", err) - } + require.NoError(t, json.Unmarshal([]byte(raw), &req)) - provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) - if _, err := provider.ChatCompletion(context.Background(), &req); err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } + provider := newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + _, err := provider.ChatCompletion(context.Background(), &req) + require.NoError(t, err) fallbacks, ok := gotBody["fallbacks"].([]any) - if !ok || len(fallbacks) != 1 || fallbacks[0] != "anthropic/claude-sonnet-latest" { - t.Fatalf("fallbacks = %#v, want Eden fallback list forwarded unchanged", gotBody["fallbacks"]) - } + require.True(t, ok, "fallbacks = %#v, want Eden fallback list forwarded unchanged", gotBody["fallbacks"]) + require.Equal(t, []any{"anthropic/claude-sonnet-latest"}, fallbacks, "want Eden fallback list forwarded unchanged") routing, ok := gotBody["routing"].(map[string]any) - if !ok || routing["strategy"] != "cost" { - t.Fatalf("routing = %#v, want Eden routing object forwarded unchanged", gotBody["routing"]) - } + require.True(t, ok, "routing = %#v, want Eden routing object forwarded unchanged", gotBody["routing"]) + require.Equal(t, "cost", routing["strategy"], "want Eden routing object forwarded unchanged") } // TestStreamChatCompletion_UsesSSE asserts that streaming requests go to @@ -247,31 +185,21 @@ func TestStreamChatCompletion_UsesSSE(t *testing.T) { })) defer server.Close() - provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + provider := newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) stream, err := provider.StreamChatCompletion(context.Background(), &core.ChatRequest{ Model: slashedModel, Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err != nil { - t.Fatalf("StreamChatCompletion() error = %v", err) - } + require.NoError(t, err) defer stream.Close() body, err := io.ReadAll(stream) - if err != nil { - t.Fatalf("ReadAll() error = %v", err) - } - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotAuth != "Bearer edenai-key" { - t.Fatalf("authorization = %q, want Bearer edenai-key", gotAuth) - } - if gotBody["model"] != slashedModel || gotBody["stream"] != true { - t.Fatalf("stream request body = %#v", gotBody) - } - if !strings.Contains(string(body), "data: [DONE]") { - t.Fatalf("stream body = %q, want SSE terminator", body) - } + require.NoError(t, err) + require.Equal(t, "/chat/completions", gotPath) + require.Equal(t, "Bearer edenai-key", gotAuth) + require.Equal(t, slashedModel, gotBody["model"], "stream request body = %#v", gotBody) + streaming, _ := gotBody["stream"].(bool) + require.True(t, streaming, "stream request body = %#v, want stream=true", gotBody) + require.Contains(t, string(body), "data: [DONE]", "stream body = %q, want SSE terminator", body) } // TestEmbeddings_ForwardsToEmbeddingsEndpoint asserts that embeddings reach @@ -299,40 +227,27 @@ func TestEmbeddings_ForwardsToEmbeddingsEndpoint(t *testing.T) { })) defer server.Close() - provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + provider := newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ Model: "openai/text-embedding-3-small", Input: "hello", }) - if err != nil { - t.Fatalf("Embeddings() error = %v", err) - } - if gotPath != "/embeddings" { - t.Fatalf("path = %q, want /embeddings", gotPath) - } - if gotAuth != "Bearer edenai-key" { - t.Fatalf("authorization = %q, want Bearer edenai-key", gotAuth) - } - if gotBody["model"] != "openai/text-embedding-3-small" { - t.Fatalf("request model = %#v, want provider/model ID forwarded unchanged", gotBody["model"]) - } + require.NoError(t, err) + require.Equal(t, "/embeddings", gotPath) + require.Equal(t, "Bearer edenai-key", gotAuth) + require.Equal(t, "openai/text-embedding-3-small", gotBody["model"], "want provider/model ID forwarded unchanged") // core.EmbeddingRequest.Provider is a gateway routing hint that the router // clears on the forwarded clone (providers.forwardEmbeddingRequest) before // any provider sees it, so it must not appear on the wire. Eden dispatches // the request it is handed, exactly as the shared // CompatibleProvider.Embeddings helper does. - if _, leaked := gotBody["provider"]; leaked { - t.Errorf("request body carried the gateway-only provider field: %#v", gotBody) - } - if len(resp.Data) != 1 || resp.Data[0].Index != 0 || len(resp.Data[0].Embedding) == 0 { - t.Fatalf("embedding data = %+v, want one populated vector", resp.Data) - } - if resp.Model != "openai/text-embedding-3-small" { - t.Errorf("response model = %q, want openai/text-embedding-3-small", resp.Model) - } - if resp.Usage.PromptTokens != 4 || resp.Usage.TotalTokens != 4 { - t.Errorf("usage = %+v, want prompt 4 / total 4", resp.Usage) - } + assert.NotContains(t, gotBody, "provider", "request body carried the gateway-only provider field: %#v", gotBody) + require.Len(t, resp.Data, 1, "embedding data = %+v, want one populated vector", resp.Data) + require.Equal(t, 0, resp.Data[0].Index, "embedding data = %+v, want one populated vector", resp.Data) + require.NotEmpty(t, resp.Data[0].Embedding, "embedding data = %+v, want one populated vector", resp.Data) + assert.Equal(t, "openai/text-embedding-3-small", resp.Model) + assert.Equal(t, 4, resp.Usage.PromptTokens, "usage = %+v, want prompt 4 / total 4", resp.Usage) + assert.Equal(t, 4, resp.Usage.TotalTokens, "usage = %+v, want prompt 4 / total 4", resp.Usage) } // TestResponses_TranslatesToChatCompletions is the load-bearing test for this @@ -362,28 +277,19 @@ func TestResponses_TranslatesToChatCompletions(t *testing.T) { })) defer server.Close() - provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + provider := newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{ Model: slashedModel, Input: "hi", }) - if err != nil { - t.Fatalf("Responses() error = %v", err) - } - if len(paths) != 1 || paths[0] != "/chat/completions" { - t.Fatalf("upstream paths = %v, want exactly [/chat/completions]", paths) - } + require.NoError(t, err) + require.Equal(t, []string{"/chat/completions"}, paths, "upstream paths = %v, want exactly [/chat/completions]", paths) for _, path := range paths { - if strings.Contains(path, "/responses") { - t.Fatalf("request reached %q; Eden's native /responses is not the OpenAI Responses API and must never be used", path) - } - } - if gotBody.Model != slashedModel { - t.Fatalf("request model = %q, want %q", gotBody.Model, slashedModel) - } - if resp.Object != "response" || resp.Status != "completed" { - t.Fatalf("response metadata = object %q status %q, want response/completed", resp.Object, resp.Status) + require.NotContains(t, path, "/responses", "request reached %q; Eden's native /responses is not the OpenAI Responses API and must never be used", path) } + require.Equal(t, slashedModel, gotBody.Model) + require.Equal(t, "response", resp.Object, "response metadata = object %q status %q, want response/completed", resp.Object, resp.Status) + require.Equal(t, "completed", resp.Status, "response metadata = object %q status %q, want response/completed", resp.Object, resp.Status) } // TestStreamResponses_TranslatesToChatCompletions asserts the streaming @@ -399,21 +305,16 @@ func TestStreamResponses_TranslatesToChatCompletions(t *testing.T) { })) defer server.Close() - provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + provider := newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) stream, err := provider.StreamResponses(context.Background(), &core.ResponsesRequest{ Model: slashedModel, Input: "hi", }) - if err != nil { - t.Fatalf("StreamResponses() error = %v", err) - } + require.NoError(t, err) defer stream.Close() - if _, err := io.ReadAll(stream); err != nil { - t.Fatalf("ReadAll() error = %v", err) - } - if len(paths) != 1 || paths[0] != "/chat/completions" { - t.Fatalf("upstream paths = %v, want exactly [/chat/completions]", paths) - } + _, err = io.ReadAll(stream) + require.NoError(t, err) + require.Equal(t, []string{"/chat/completions"}, paths, "upstream paths = %v, want exactly [/chat/completions]", paths) } // TestPassthrough_ForwardsOpaqueRequest asserts the provider forwards an opaque @@ -430,25 +331,17 @@ func TestPassthrough_ForwardsOpaqueRequest(t *testing.T) { })) defer server.Close() - provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + provider := newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.Passthrough(context.Background(), &core.PassthroughRequest{ Method: http.MethodPost, Endpoint: "chat/completions", Body: io.NopCloser(strings.NewReader(`{"model":"openai/gpt-4"}`)), }) - if err != nil { - t.Fatalf("Passthrough() error = %v", err) - } + require.NoError(t, err) defer resp.Body.Close() - if gotPath != "/chat/completions" { - t.Fatalf("path = %q, want /chat/completions", gotPath) - } - if gotAuth != "Bearer edenai-key" { - t.Fatalf("authorization = %q, want Bearer edenai-key", gotAuth) - } - if resp.StatusCode != http.StatusOK { - t.Fatalf("status = %d, want 200", resp.StatusCode) - } + require.Equal(t, "/chat/completions", gotPath) + require.Equal(t, "Bearer edenai-key", gotAuth) + require.Equal(t, http.StatusOK, resp.StatusCode) } // TestProvider_DoesNotExposeOptionalOpenAICompatibleInterfaces guards the @@ -459,18 +352,9 @@ func TestPassthrough_ForwardsOpaqueRequest(t *testing.T) { // upstream cannot honour. This matters more under composition than it did // under embedding: adding a delegation by mistake is all it would take. func TestProvider_DoesNotExposeOptionalOpenAICompatibleInterfaces(t *testing.T) { - provider := NewWithHTTPClient("edenai-key", "", nil, llmclient.Hooks{}) + provider := newTestProvider("edenai-key", "", nil, llmclient.Hooks{}) - if _, ok := any(provider).(core.NativeBatchProvider); ok { - t.Fatal("edenai provider should not implement native batch provider") - } - if _, ok := any(provider).(core.NativeFileProvider); ok { - t.Fatal("edenai provider should not implement native file provider") - } - if _, ok := any(provider).(core.AudioProvider); ok { - t.Fatal("edenai provider should not implement audio provider") - } - if _, ok := any(provider).(core.ImageProvider); ok { - t.Fatal("edenai provider should not implement image provider") - } + providertest.AssertNoNativeSurfaces(t, provider) + _, ok := any(provider).(core.ImageProvider) + assert.False(t, ok, "edenai provider should not implement image provider") } diff --git a/internal/providers/edenai/embeddings_cost_test.go b/internal/providers/edenai/embeddings_cost_test.go index fb350312b..d6cb9faef 100644 --- a/internal/providers/edenai/embeddings_cost_test.go +++ b/internal/providers/edenai/embeddings_cost_test.go @@ -9,6 +9,8 @@ import ( "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" "github.com/enterpilot/gomodel/internal/usage" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) // edenEmbeddingBody is the shape Eden documents for /v3/embeddings: the @@ -31,7 +33,7 @@ func embeddingsProvider(t *testing.T, payload string) *Provider { _, _ = w.Write([]byte(payload)) })) t.Cleanup(server.Close) - return NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + return newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) } // TestEmbeddings_LiftsRootLevelCostIntoUsage asserts Eden's root-level @@ -45,31 +47,19 @@ func TestEmbeddings_LiftsRootLevelCostIntoUsage(t *testing.T) { provider := embeddingsProvider(t, edenEmbeddingBody) resp, err := provider.Embeddings(context.Background(), embeddingRequest()) - if err != nil { - t.Fatalf("Embeddings() error = %v", err) - } + require.NoError(t, err) cost, ok := resp.Usage.RawUsage["cost"] - if !ok { - t.Fatalf("Usage.RawUsage = %v, want Eden's root-level cost lifted into it", resp.Usage.RawUsage) - } - if cost != 0.0000012 { - t.Errorf("RawUsage[cost] = %v, want 0.0000012 exactly", cost) - } + require.True(t, ok, "Usage.RawUsage = %v, want Eden's root-level cost lifted into it", resp.Usage.RawUsage) + assert.Equal(t, 0.0000012, cost, "RawUsage[cost] = %v, want 0.0000012 exactly", cost) // The rest of the envelope must still decode normally. - if resp.Usage.PromptTokens != 9 || resp.Usage.TotalTokens != 9 { - t.Errorf("usage tokens = %d/%d, want 9/9", resp.Usage.PromptTokens, resp.Usage.TotalTokens) - } - if len(resp.Data) != 1 || len(resp.Data[0].Embedding) == 0 { - t.Errorf("data = %+v, want one embedding", resp.Data) - } - if resp.Model != "openai/text-embedding-3-small" { - t.Errorf("model = %q, want the Eden model ID forwarded verbatim", resp.Model) - } + assert.Equal(t, 9, resp.Usage.PromptTokens, "usage tokens = %d/%d, want 9/9", resp.Usage.PromptTokens, resp.Usage.TotalTokens) + assert.Equal(t, 9, resp.Usage.TotalTokens, "usage tokens = %d/%d, want 9/9", resp.Usage.PromptTokens, resp.Usage.TotalTokens) + require.Len(t, resp.Data, 1, "data = %+v, want one embedding", resp.Data) + assert.NotEmpty(t, resp.Data[0].Embedding, "data = %+v, want one embedding", resp.Data) + assert.Equal(t, "openai/text-embedding-3-small", resp.Model, "want the Eden model ID forwarded verbatim") // Eden's upstream name must not be reported as the executing provider. - if resp.Provider != "" { - t.Errorf("Provider = %q, want empty so the gateway reports edenai", resp.Provider) - } + assert.Empty(t, resp.Provider, "Provider = %q, want empty so the gateway reports edenai", resp.Provider) } // TestEmbeddings_ExactCostReachesRecordedTotal is the end-to-end assertion for @@ -84,29 +74,17 @@ func TestEmbeddings_ExactCostReachesRecordedTotal(t *testing.T) { provider := embeddingsProvider(t, edenEmbeddingBody) resp, err := provider.Embeddings(context.Background(), embeddingRequest()) - if err != nil { - t.Fatalf("Embeddings() error = %v", err) - } + require.NoError(t, err) wrongRate := 100.0 // $100/MTok would price 9 tokens at $0.0009, not $0.0000012. pricing := &core.ModelPricing{Currency: "USD", InputPerMtok: &wrongRate} entry := usage.ExtractFromEmbeddingResponse(resp, "req-1", providerType, "/v1/embeddings", pricing) - if entry == nil { - t.Fatal("ExtractFromEmbeddingResponse returned nil") - } - if entry.TotalCost == nil { - t.Fatal("TotalCost = nil, want Eden's exact charge recorded") - } - if *entry.TotalCost != 0.0000012 { - t.Errorf("TotalCost = %v, want 0.0000012 exactly (Eden's reported charge, not a token-rate estimate)", *entry.TotalCost) - } - if entry.CostSource != usage.CostSourceEdenAICost { - t.Errorf("CostSource = %q, want %q", entry.CostSource, usage.CostSourceEdenAICost) - } - if entry.CostsCalculationCaveat != "" { - t.Errorf("CostsCalculationCaveat = %q, want empty for an exact provider-reported cost", entry.CostsCalculationCaveat) - } + require.NotNil(t, entry, "ExtractFromEmbeddingResponse returned nil") + require.NotNil(t, entry.TotalCost, "TotalCost = nil, want Eden's exact charge recorded") + assert.Equal(t, 0.0000012, *entry.TotalCost, "TotalCost = %v, want 0.0000012 exactly (Eden's reported charge, not a token-rate estimate)", *entry.TotalCost) + assert.Equal(t, usage.CostSourceEdenAICost, entry.CostSource) + assert.Empty(t, entry.CostsCalculationCaveat, "CostsCalculationCaveat = %q, want empty for an exact provider-reported cost", entry.CostsCalculationCaveat) } // TestEmbeddings_FallsBackToTokenPricingWithoutCost asserts the fallback still @@ -121,29 +99,17 @@ func TestEmbeddings_FallsBackToTokenPricingWithoutCost(t *testing.T) { }`) resp, err := provider.Embeddings(context.Background(), embeddingRequest()) - if err != nil { - t.Fatalf("Embeddings() error = %v", err) - } - if len(resp.Usage.RawUsage) != 0 { - t.Errorf("Usage.RawUsage = %v, want empty when Eden reports no cost", resp.Usage.RawUsage) - } + require.NoError(t, err) + assert.Empty(t, resp.Usage.RawUsage, "Usage.RawUsage = %v, want empty when Eden reports no cost", resp.Usage.RawUsage) rate := 0.02 // $0.02/MTok * 1000 tokens = $0.00002 pricing := &core.ModelPricing{Currency: "USD", InputPerMtok: &rate} entry := usage.ExtractFromEmbeddingResponse(resp, "req-2", providerType, "/v1/embeddings", pricing) - if entry == nil { - t.Fatal("ExtractFromEmbeddingResponse returned nil") - } - if entry.TotalCost == nil { - t.Fatal("TotalCost = nil, want the token-rate fallback to apply") - } - if *entry.TotalCost != 0.00002 { - t.Errorf("TotalCost = %v, want 0.00002 from the catalog rate", *entry.TotalCost) - } - if entry.CostSource != usage.CostSourceModelPricing { - t.Errorf("CostSource = %q, want %q", entry.CostSource, usage.CostSourceModelPricing) - } + require.NotNil(t, entry, "ExtractFromEmbeddingResponse returned nil") + require.NotNil(t, entry.TotalCost, "TotalCost = nil, want the token-rate fallback to apply") + assert.Equal(t, 0.00002, *entry.TotalCost, "TotalCost = %v, want 0.00002 from the catalog rate", *entry.TotalCost) + assert.Equal(t, usage.CostSourceModelPricing, entry.CostSource) } // TestEmbeddings_RejectsUnusableCost asserts a cost member that would corrupt @@ -175,19 +141,13 @@ func TestEmbeddings_RejectsUnusableCost(t *testing.T) { }`) resp, err := provider.Embeddings(context.Background(), embeddingRequest()) - if err != nil { - t.Fatalf("Embeddings() error = %v", err) - } - if _, present := resp.Usage.RawUsage["cost"]; present { - t.Fatalf("RawUsage[cost] = %v, want absent for an unusable cost member", resp.Usage.RawUsage["cost"]) - } + require.NoError(t, err) + require.NotContains(t, resp.Usage.RawUsage, "cost", "RawUsage[cost] = %v, want absent for an unusable cost member", resp.Usage.RawUsage["cost"]) rate := 0.02 entry := usage.ExtractFromEmbeddingResponse(resp, "req", providerType, "/v1/embeddings", &core.ModelPricing{Currency: "USD", InputPerMtok: &rate}) - if entry.CostSource != usage.CostSourceModelPricing { - t.Errorf("CostSource = %q, want the token-rate fallback %q", entry.CostSource, usage.CostSourceModelPricing) - } + assert.Equal(t, usage.CostSourceModelPricing, entry.CostSource, "CostSource = %q, want the token-rate fallback %q", entry.CostSource, usage.CostSourceModelPricing) }) } } @@ -205,19 +165,14 @@ func TestEmbeddings_ZeroCostIsRecordedAsFree(t *testing.T) { }`) resp, err := provider.Embeddings(context.Background(), embeddingRequest()) - if err != nil { - t.Fatalf("Embeddings() error = %v", err) - } + require.NoError(t, err) rate := 0.02 entry := usage.ExtractFromEmbeddingResponse(resp, "req", providerType, "/v1/embeddings", &core.ModelPricing{Currency: "USD", InputPerMtok: &rate}) - if entry.TotalCost == nil || *entry.TotalCost != 0 { - t.Fatalf("TotalCost = %v, want 0 recorded from Eden's explicit zero", entry.TotalCost) - } - if entry.CostSource != usage.CostSourceEdenAICost { - t.Errorf("CostSource = %q, want %q", entry.CostSource, usage.CostSourceEdenAICost) - } + require.NotNil(t, entry.TotalCost, "TotalCost = nil, want 0 recorded from Eden's explicit zero") + require.Zero(t, *entry.TotalCost, "TotalCost = %v, want 0 recorded from Eden's explicit zero", *entry.TotalCost) + assert.Equal(t, usage.CostSourceEdenAICost, entry.CostSource) } // TestEmbeddings_UsageLevelCostWins asserts the conventional location stays @@ -233,12 +188,8 @@ func TestEmbeddings_UsageLevelCostWins(t *testing.T) { }`) resp, err := provider.Embeddings(context.Background(), embeddingRequest()) - if err != nil { - t.Fatalf("Embeddings() error = %v", err) - } - if got := resp.Usage.RawUsage["cost"]; got != 0.5 { - t.Errorf("RawUsage[cost] = %v, want the pre-existing usage-level 0.5 preserved", got) - } + require.NoError(t, err) + assert.Equal(t, 0.5, resp.Usage.RawUsage["cost"], "RawUsage[cost] = %v, want the pre-existing usage-level 0.5 preserved", resp.Usage.RawUsage["cost"]) } // TestEmbeddings_NilRequestRejected asserts the guard the shared helper used to @@ -246,9 +197,8 @@ func TestEmbeddings_UsageLevelCostWins(t *testing.T) { func TestEmbeddings_NilRequestRejected(t *testing.T) { provider := embeddingsProvider(t, edenEmbeddingBody) - if _, err := provider.Embeddings(context.Background(), nil); err == nil { - t.Fatal("Embeddings(nil) = nil error, want a rejection") - } + _, err := provider.Embeddings(context.Background(), nil) + require.Error(t, err, "Embeddings(nil) = nil error, want a rejection") } // TestEmbeddings_BackfillsModelFromRequest asserts the EnsureModel behavior the @@ -262,10 +212,6 @@ func TestEmbeddings_BackfillsModelFromRequest(t *testing.T) { }`) resp, err := provider.Embeddings(context.Background(), embeddingRequest()) - if err != nil { - t.Fatalf("Embeddings() error = %v", err) - } - if resp.Model != slashedModel { - t.Errorf("Model = %q, want the requested %q backfilled", resp.Model, slashedModel) - } + require.NoError(t, err) + assert.Equal(t, slashedModel, resp.Model, "Model = %q, want the requested %q backfilled", resp.Model, slashedModel) } diff --git a/internal/providers/edenai/models_test.go b/internal/providers/edenai/models_test.go index 087095cd0..8de01d0be 100644 --- a/internal/providers/edenai/models_test.go +++ b/internal/providers/edenai/models_test.go @@ -5,11 +5,12 @@ import ( "math" "net/http" "net/http/httptest" - "slices" "testing" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) // edenCatalogEntry is a verbatim entry from a live Eden /v3/models response. @@ -67,19 +68,15 @@ func modelsServer(t *testing.T, payload string) (*Provider, *string, *string) { _, _ = w.Write([]byte(payload)) })) t.Cleanup(server.Close) - return NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}), gotPath, gotAuth + return newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}), gotPath, gotAuth } func firstModel(t *testing.T, payload string) core.Model { t.Helper() provider, _, _ := modelsServer(t, payload) resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if len(resp.Data) != 1 { - t.Fatalf("models = %+v, want exactly one entry", resp.Data) - } + require.NoError(t, err) + require.Len(t, resp.Data, 1, "models = %+v, want exactly one entry", resp.Data) return resp.Data[0] } @@ -90,67 +87,44 @@ func TestListModels_MapsLiveCatalogEntry(t *testing.T) { provider, gotPath, gotAuth := modelsServer(t, `{"object":"list","data":[`+edenCatalogEntry+`]}`) resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if *gotPath != "/models" { - t.Fatalf("path = %q, want /models", *gotPath) - } - if *gotAuth != "Bearer edenai-key" { - t.Fatalf("authorization = %q, want Bearer edenai-key", *gotAuth) - } - if resp.Object != "list" || len(resp.Data) != 1 { - t.Fatalf("response = %+v, want a one-entry list", resp) - } + require.NoError(t, err) + require.Equal(t, "/models", *gotPath) + require.Equal(t, "Bearer edenai-key", *gotAuth) + require.Equal(t, "list", resp.Object, "response = %+v, want a one-entry list", resp) + require.Len(t, resp.Data, 1, "response = %+v, want a one-entry list", resp) model := resp.Data[0] - if model.ID != "deepinfra/inclusionAI/Ling-3.0-flash-VL" { - t.Errorf("ID = %q, want the Eden ID forwarded unchanged", model.ID) - } - if model.Object != "model" || model.OwnedBy != "deepinfra" || model.Created != 1788880306 { - t.Errorf("identity = %+v, want object/owned_by/created preserved", model) - } + assert.Equal(t, "deepinfra/inclusionAI/Ling-3.0-flash-VL", model.ID, "want the Eden ID forwarded unchanged") + assert.Equal(t, "model", model.Object, "identity = %+v, want object/owned_by/created preserved", model) + assert.Equal(t, "deepinfra", model.OwnedBy, "identity = %+v, want object/owned_by/created preserved", model) + assert.Equal(t, int64(1788880306), model.Created, "identity = %+v, want object/owned_by/created preserved", model) meta := model.Metadata - if meta == nil { - t.Fatal("Metadata = nil, want Eden catalog metadata") - } - if meta.ContextWindow == nil || *meta.ContextWindow != 131072 { - t.Errorf("ContextWindow = %v, want 131072", meta.ContextWindow) - } + require.NotNil(t, meta, "Metadata = nil, want Eden catalog metadata") + require.NotNil(t, meta.ContextWindow, "ContextWindow = nil, want 131072") + assert.Equal(t, 131072, *meta.ContextWindow) // output_modalities ["text"] -> chat + responses (Responses is served by // translating through chat completions). - if want := []string{"chat", "responses"}; !slices.Equal(meta.Modes, want) { - t.Errorf("Modes = %v, want %v", meta.Modes, want) - } - if len(meta.Categories) != 1 || meta.Categories[0] != core.CategoryTextGeneration { - t.Errorf("Categories = %v, want [%v]", meta.Categories, core.CategoryTextGeneration) - } + assert.Equal(t, []string{"chat", "responses"}, meta.Modes) + assert.Equal(t, []core.ModelCategory{core.CategoryTextGeneration}, meta.Categories) // Only the true supports_* flags become capabilities, with the prefix // stripped; input_modalities ["text","image"] adds vision. wantCapabilities := map[string]bool{"reasoning": true, "prompt_caching": true, "vision": true} - if len(meta.Capabilities) != len(wantCapabilities) { - t.Errorf("Capabilities = %v, want %v", meta.Capabilities, wantCapabilities) - } + assert.Len(t, meta.Capabilities, len(wantCapabilities), "Capabilities = %v, want %v", meta.Capabilities, wantCapabilities) for name := range wantCapabilities { - if !meta.Capabilities[name] { - t.Errorf("Capabilities[%q] = false, want true", name) - } - } - if meta.Capabilities["web_search"] || meta.Capabilities["function_calling"] { - t.Errorf("Capabilities = %v, want false flags omitted", meta.Capabilities) + assert.True(t, meta.Capabilities[name], "Capabilities[%q] = false, want true", name) } + assert.False(t, meta.Capabilities["web_search"], "Capabilities = %v, want false flags omitted", meta.Capabilities) + assert.False(t, meta.Capabilities["function_calling"], "Capabilities = %v, want false flags omitted", meta.Capabilities) // 6e-8 USD/token -> $0.06/MTok, 1.8e-7 -> $0.18, 1.2e-8 -> $0.012. // The values come from `pricing`, not the higher `list_pricing` block. assertPrice(t, "InputPerMtok", meta.Pricing.InputPerMtok, 0.06) assertPrice(t, "OutputPerMtok", meta.Pricing.OutputPerMtok, 0.18) assertPrice(t, "CachedInputPerMtok", meta.Pricing.CachedInputPerMtok, 0.012) - if meta.Pricing.Currency != "USD" { - t.Errorf("Currency = %q, want USD", meta.Pricing.Currency) - } + assert.Equal(t, "USD", meta.Pricing.Currency) } // TestListModels_UsesDiscountedPricingNotListPricing pins the choice of block: @@ -158,9 +132,7 @@ func TestListModels_MapsLiveCatalogEntry(t *testing.T) { func TestListModels_UsesDiscountedPricingNotListPricing(t *testing.T) { model := firstModel(t, `{"object":"list","data":[`+edenCatalogEntry+`]}`) assertPrice(t, "InputPerMtok", model.Metadata.Pricing.InputPerMtok, 0.06) - if got := *model.Metadata.Pricing.InputPerMtok; got == 0.09 { - t.Fatal("InputPerMtok took the list_pricing rate; want the discounted pricing block") - } + require.NotEqual(t, 0.09, *model.Metadata.Pricing.InputPerMtok, "InputPerMtok took the list_pricing rate; want the discounted pricing block") } // TestListModels_FallsBackToListPricing asserts the undiscounted rate card is @@ -174,9 +146,8 @@ func TestListModels_FallsBackToListPricing(t *testing.T) { "list_pricing": {"input_cost_per_token": 9e-8, "output_cost_per_token": 2.8e-7} }]}`) - if model.Metadata == nil || model.Metadata.Pricing == nil { - t.Fatal("Pricing = nil, want the list_pricing fallback") - } + require.NotNil(t, model.Metadata, "Pricing = nil, want the list_pricing fallback") + require.NotNil(t, model.Metadata.Pricing, "Pricing = nil, want the list_pricing fallback") assertPrice(t, "InputPerMtok", model.Metadata.Pricing.InputPerMtok, 0.09) assertPrice(t, "OutputPerMtok", model.Metadata.Pricing.OutputPerMtok, 0.28) } @@ -289,10 +260,7 @@ func TestListModels_UnusableRateInBothBlocksStaysUnpriced(t *testing.T) { }]}`) assertPrice(t, "InputPerMtok", model.Metadata.Pricing.InputPerMtok, 0.06) - if model.Metadata.Pricing.OutputPerMtok != nil { - t.Errorf("OutputPerMtok = %v, want nil: neither block published a usable output rate", - *model.Metadata.Pricing.OutputPerMtok) - } + assert.Nil(t, model.Metadata.Pricing.OutputPerMtok, "OutputPerMtok want nil: neither block published a usable output rate") } // TestListModels_CachePricingFields asserts both of Eden's cache rates reach @@ -318,14 +286,8 @@ func TestListModels_CachePricingFields(t *testing.T) { pricing := model.Metadata.Pricing assertPrice(t, "CachedInputPerMtok", pricing.CachedInputPerMtok, 0.012) assertPrice(t, "CacheWritePerMtok", pricing.CacheWritePerMtok, 0.075) - if pricing.ReasoningOutputPerMtok != nil { - t.Errorf("ReasoningOutputPerMtok = %v, want nil: Eden reports no reasoning token count to price against", - *pricing.ReasoningOutputPerMtok) - } - if pricing.AudioInputPerMtok != nil { - t.Errorf("AudioInputPerMtok = %v, want nil: Eden reports no audio token count to price against", - *pricing.AudioInputPerMtok) - } + assert.Nil(t, pricing.ReasoningOutputPerMtok, "ReasoningOutputPerMtok want nil: Eden reports no reasoning token count to price against") + assert.Nil(t, pricing.AudioInputPerMtok, "AudioInputPerMtok want nil: Eden reports no audio token count to price against") } // TestListModels_TieredAndPerQueryPricingIgnored asserts the Eden pricing @@ -349,12 +311,8 @@ func TestListModels_TieredAndPerQueryPricingIgnored(t *testing.T) { pricing := model.Metadata.Pricing assertPrice(t, "InputPerMtok", pricing.InputPerMtok, 0.06) assertPrice(t, "OutputPerMtok", pricing.OutputPerMtok, 0.18) - if len(pricing.Tiers) != 0 { - t.Errorf("Tiers = %v, want empty: Eden's tiered_pricing shape is not mapped", pricing.Tiers) - } - if pricing.PerRequest != nil { - t.Errorf("PerRequest = %v, want nil: per-query search fees are not per-request charges", *pricing.PerRequest) - } + assert.Empty(t, pricing.Tiers, "Tiers = %v, want empty: Eden's tiered_pricing shape is not mapped", pricing.Tiers) + assert.Nil(t, pricing.PerRequest, "PerRequest want nil: per-query search fees are not per-request charges") } // TestListModels_PricingEdgeCases covers partial, zero, negative, and @@ -418,14 +376,13 @@ func TestListModels_PricingEdgeCases(t *testing.T) { model := firstModel(t, payload) if tt.wantNil { - if model.Metadata != nil && model.Metadata.Pricing != nil { - t.Fatalf("Pricing = %+v, want nil", model.Metadata.Pricing) + if model.Metadata != nil { + require.Nil(t, model.Metadata.Pricing, "Pricing = %+v, want nil", model.Metadata.Pricing) } return } - if model.Metadata == nil || model.Metadata.Pricing == nil { - t.Fatal("Pricing = nil, want a partial pricing block") - } + require.NotNil(t, model.Metadata, "Pricing = nil, want a partial pricing block") + require.NotNil(t, model.Metadata.Pricing, "Pricing = nil, want a partial pricing block") assertOptionalPrice(t, "InputPerMtok", model.Metadata.Pricing.InputPerMtok, tt.wantInput) assertOptionalPrice(t, "OutputPerMtok", model.Metadata.Pricing.OutputPerMtok, tt.wantOutput) assertOptionalPrice(t, "CachedInputPerMtok", model.Metadata.Pricing.CachedInputPerMtok, tt.wantCached) @@ -438,17 +395,14 @@ func TestListModels_PricingEdgeCases(t *testing.T) { // that tolerates them, and both must be refused rather than scaled. func TestPerMtok_RejectsNonFiniteRates(t *testing.T) { for _, rate := range []float64{math.NaN(), math.Inf(1), math.Inf(-1), -1} { - if _, ok := perMtok(&rate); ok { - t.Errorf("perMtok(%v) reported a usable price, want rejected", rate) - } - } - if _, ok := perMtok(nil); ok { - t.Error("perMtok(nil) reported a usable price, want rejected") + _, ok := perMtok(&rate) + assert.False(t, ok, "perMtok(%v) reported a usable price, want rejected", rate) } + _, ok := perMtok(nil) + assert.False(t, ok, "perMtok(nil) reported a usable price, want rejected") value, ok := perMtok(new(6e-8)) - if !ok || math.Abs(value-0.06) > 1e-9 { - t.Errorf("perMtok(6e-8) = %v, %v; want 0.06, true", value, ok) - } + assert.True(t, ok, "perMtok(6e-8) = %v, %v; want 0.06, true", value, ok) + assert.InDelta(t, 0.06, value, 1e-9, "perMtok(6e-8) = %v, %v; want 0.06, true", value, ok) } // TestListModels_ModalityMapping asserts output modalities become modes (so @@ -526,13 +480,10 @@ func TestListModels_ModalityMapping(t *testing.T) { if model.Metadata != nil { modes = model.Metadata.Modes } - if !slices.Equal(modes, tt.wantModes) { - t.Errorf("Modes = %v, want %v", modes, tt.wantModes) - } + assert.Equal(t, tt.wantModes, modes) for _, capability := range tt.wantCapabilities { - if model.Metadata == nil || !model.Metadata.Capabilities[capability] { - t.Errorf("Capabilities missing %q", capability) - } + require.NotNil(t, model.Metadata, "Capabilities missing %q", capability) + assert.True(t, model.Metadata.Capabilities[capability], "Capabilities missing %q", capability) } }) } @@ -549,18 +500,11 @@ func TestListModels_SkipsInvalidEntriesAndKeepsBareOnes(t *testing.T) { ]}`) resp, err := provider.ListModels(context.Background()) - if err != nil { - t.Fatalf("ListModels() error = %v", err) - } - if len(resp.Data) != 2 { - t.Fatalf("models = %+v, want the blank ID dropped", resp.Data) - } - if resp.Data[0].ID != "openai/gpt-4" || resp.Data[0].Metadata != nil { - t.Errorf("bare entry = %+v, want nil Metadata", resp.Data[0]) - } - if resp.Data[1].Object != "model" { - t.Errorf("Object = %q, want the default \"model\" applied", resp.Data[1].Object) - } + require.NoError(t, err) + require.Len(t, resp.Data, 2, "models = %+v, want the blank ID dropped", resp.Data) + assert.Equal(t, "openai/gpt-4", resp.Data[0].ID, "bare entry = %+v", resp.Data[0]) + assert.Nil(t, resp.Data[0].Metadata, "bare entry = %+v, want nil Metadata", resp.Data[0]) + assert.Equal(t, "model", resp.Data[1].Object, `want the default "model" applied`) } // TestListModels_PropagatesUpstreamError asserts a failed catalog fetch @@ -573,33 +517,22 @@ func TestListModels_PropagatesUpstreamError(t *testing.T) { })) defer server.Close() - provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + provider := newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.ListModels(context.Background()) - if err == nil { - t.Fatalf("ListModels() error = nil, want the upstream failure; resp = %+v", resp) - } - if resp != nil { - t.Errorf("ListModels() resp = %+v, want nil on error", resp) - } + require.Error(t, err, "ListModels() error = nil, want the upstream failure; resp = %+v", resp) + assert.Nil(t, resp, "ListModels() resp = %+v, want nil on error", resp) } func assertPrice(t *testing.T, name string, got *float64, want float64) { t.Helper() - if got == nil { - t.Errorf("%s = nil, want %v", name, want) - return - } - if math.Abs(*got-want) > 1e-9 { - t.Errorf("%s = %v, want %v", name, *got, want) - } + require.NotNil(t, got, "%s = nil, want %v", name, want) + assert.InDelta(t, want, *got, 1e-9, "%s = %v, want %v", name, *got, want) } func assertOptionalPrice(t *testing.T, name string, got, want *float64) { t.Helper() if want == nil { - if got != nil { - t.Errorf("%s = %v, want nil (an unreported rate must not be costed)", name, *got) - } + assert.Nil(t, got, "%s want nil (an unreported rate must not be costed)", name) return } assertPrice(t, name, got, *want) diff --git a/internal/providers/edenai/newtestprovider_test.go b/internal/providers/edenai/newtestprovider_test.go new file mode 100644 index 000000000..7ad40c439 --- /dev/null +++ b/internal/providers/edenai/newtestprovider_test.go @@ -0,0 +1,16 @@ +package edenai + +import ( + "net/http" + + "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" +) + +// newTestProvider builds the provider through its own constructor on a test transport. +func newTestProvider(apiKey, baseURL string, httpClient *http.Client, hooks llmclient.Hooks) *Provider { + opts := providertest.Options(hooks) + opts.HTTPClient = httpClient + return New(providers.ProviderConfig{APIKey: apiKey, BaseURL: baseURL}, opts).(*Provider) +} diff --git a/internal/providers/edenai/passthrough_semantics_test.go b/internal/providers/edenai/passthrough_semantics_test.go index 698b4652f..40ced510b 100644 --- a/internal/providers/edenai/passthrough_semantics_test.go +++ b/internal/providers/edenai/passthrough_semantics_test.go @@ -4,12 +4,12 @@ import ( "testing" "github.com/enterpilot/gomodel/internal/core" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestPassthroughSemanticEnricher(t *testing.T) { - if got := passthroughSemanticEnricher.ProviderType(); got != "edenai" { - t.Fatalf("ProviderType() = %q, want edenai", got) - } + require.Equal(t, "edenai", passthroughSemanticEnricher.ProviderType()) tests := []struct { name string @@ -43,18 +43,10 @@ func TestPassthroughSemanticEnricher(t *testing.T) { RawEndpoint: tt.rawEndpoint, NormalizedEndpoint: tt.normalizedEndpoint, }) - if got == nil { - t.Fatal("Enrich() returned nil") - } - if got.SemanticOperation != tt.wantOperation { - t.Errorf("SemanticOperation = %q, want %q", got.SemanticOperation, tt.wantOperation) - } - if got.GenAIOperation != tt.wantGenAIOperation { - t.Errorf("GenAIOperation = %q, want %q", got.GenAIOperation, tt.wantGenAIOperation) - } - if got.AuditPath != tt.wantAuditPath { - t.Errorf("AuditPath = %q, want %q", got.AuditPath, tt.wantAuditPath) - } + require.NotNil(t, got, "Enrich() returned nil") + assert.Equal(t, tt.wantOperation, got.SemanticOperation) + assert.Equal(t, tt.wantGenAIOperation, got.GenAIOperation) + assert.Equal(t, tt.wantAuditPath, got.AuditPath) }) } } @@ -69,13 +61,7 @@ func TestPassthroughSemanticEnricher_ResponsesIsNotOpenAIShaped(t *testing.T) { RawEndpoint: "v1/responses", NormalizedEndpoint: "responses", }) - if got == nil { - t.Fatal("Enrich() returned nil") - } - if got.SemanticOperation != "" { - t.Errorf("SemanticOperation = %q, want empty: Eden /responses must not be advertised as OpenAI Responses", got.SemanticOperation) - } - if got.AuditPath != "/p/edenai/responses" { - t.Errorf("AuditPath = %q, want /p/edenai/responses", got.AuditPath) - } + require.NotNil(t, got, "Enrich() returned nil") + assert.Empty(t, got.SemanticOperation, "SemanticOperation = %q, want empty: Eden /responses must not be advertised as OpenAI Responses", got.SemanticOperation) + assert.Equal(t, "/p/edenai/responses", got.AuditPath) } diff --git a/internal/providers/edenai/response_test.go b/internal/providers/edenai/response_test.go index 14efe01ed..976cdc25c 100644 --- a/internal/providers/edenai/response_test.go +++ b/internal/providers/edenai/response_test.go @@ -2,7 +2,6 @@ package edenai import ( "context" - "math" "net/http" "net/http/httptest" "testing" @@ -11,6 +10,8 @@ import ( "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) // edenChatResponse is a verbatim non-streaming Eden /v3/chat/completions body. @@ -41,14 +42,12 @@ func chatResponseFrom(t *testing.T, payload string) *core.ChatResponse { })) t.Cleanup(server.Close) - provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + provider := newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ Model: slashedModel, Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err != nil { - t.Fatalf("ChatCompletion() error = %v", err) - } + require.NoError(t, err) return resp } @@ -60,16 +59,12 @@ func TestChatCompletion_LiftsRootCostIntoRawUsage(t *testing.T) { resp := chatResponseFrom(t, edenChatResponse) cost, ok := resp.Usage.RawUsage["cost"] - if !ok { - t.Fatalf("Usage.RawUsage = %v, want Eden's root-level cost lifted in", resp.Usage.RawUsage) - } + require.True(t, ok, "Usage.RawUsage = %v, want Eden's root-level cost lifted in", resp.Usage.RawUsage) value, ok := cost.(float64) - if !ok || math.Abs(value-0.0002349) > 1e-12 { - t.Fatalf("Usage.RawUsage[\"cost\"] = %#v, want 0.0002349", cost) - } - if resp.Usage.PromptTokens != 1170 || resp.Usage.CompletionTokens != 99 { - t.Errorf("token counts = %+v, want the reported usage preserved", resp.Usage) - } + require.True(t, ok, "Usage.RawUsage[\"cost\"] = %#v, want 0.0002349", cost) + require.InDelta(t, 0.0002349, value, 1e-12, "Usage.RawUsage[\"cost\"] = %#v, want 0.0002349", cost) + assert.Equal(t, 1170, resp.Usage.PromptTokens, "token counts = %+v, want the reported usage preserved", resp.Usage) + assert.Equal(t, 99, resp.Usage.CompletionTokens, "token counts = %+v, want the reported usage preserved", resp.Usage) } // TestChatCompletion_KeepsCostVisibleToClients asserts lifting the value into @@ -78,16 +73,12 @@ func TestChatCompletion_KeepsCostVisibleToClients(t *testing.T) { resp := chatResponseFrom(t, edenChatResponse) encoded, err := json.Marshal(resp) - if err != nil { - t.Fatalf("Marshal() error = %v", err) - } + require.NoError(t, err) var decoded map[string]any - if err := json.Unmarshal(encoded, &decoded); err != nil { - t.Fatalf("Unmarshal() error = %v", err) - } - if cost, ok := decoded["cost"].(float64); !ok || math.Abs(cost-0.0002349) > 1e-12 { - t.Errorf("serialized cost = %#v, want Eden's cost preserved for the client", decoded["cost"]) - } + require.NoError(t, json.Unmarshal(encoded, &decoded)) + cost, ok := decoded["cost"].(float64) + require.True(t, ok, "serialized cost = %#v, want Eden's cost preserved for the client", decoded["cost"]) + assert.InDelta(t, 0.0002349, cost, 1e-12, "serialized cost = %#v, want Eden's cost preserved for the client", decoded["cost"]) } // TestChatCompletion_RejectsUnusableCost asserts a cost that would corrupt @@ -109,9 +100,7 @@ func TestChatCompletion_RejectsUnusableCost(t *testing.T) { payload := `{"id":"c","created":1,"model":"m","choices":[],` + tt.cost + `"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}` resp := chatResponseFrom(t, payload) - if _, ok := resp.Usage.RawUsage["cost"]; ok { - t.Fatalf("Usage.RawUsage = %v, want no cost lifted for an unusable value", resp.Usage.RawUsage) - } + require.NotContains(t, resp.Usage.RawUsage, "cost", "Usage.RawUsage = %v, want no cost lifted for an unusable value", resp.Usage.RawUsage) }) } } @@ -124,9 +113,8 @@ func TestChatCompletion_UsageLevelCostWins(t *testing.T) { resp := chatResponseFrom(t, payload) value, ok := resp.Usage.RawUsage["cost"].(float64) - if !ok || value != 0.5 { - t.Fatalf("Usage.RawUsage[\"cost\"] = %#v, want the usage-level 0.5", resp.Usage.RawUsage["cost"]) - } + require.True(t, ok, "Usage.RawUsage[\"cost\"] = %#v, want the usage-level 0.5", resp.Usage.RawUsage["cost"]) + require.Equal(t, 0.5, value, "Usage.RawUsage[\"cost\"] = %#v, want the usage-level 0.5", resp.Usage.RawUsage["cost"]) } // TestChatCompletion_DoesNotReportEdenUpstreamAsExecutingProvider asserts @@ -139,9 +127,7 @@ func TestChatCompletion_UsageLevelCostWins(t *testing.T) { func TestChatCompletion_DoesNotReportEdenUpstreamAsExecutingProvider(t *testing.T) { resp := chatResponseFrom(t, edenChatResponse) - if resp.Provider != "" { - t.Fatalf("Provider = %q, want empty so the gateway labels the request edenai", resp.Provider) - } + require.Empty(t, resp.Provider, "Provider = %q, want empty so the gateway labels the request edenai", resp.Provider) } // TestChatCompletion_PreservesEdenUpstreamProvider asserts the upstream is not @@ -151,16 +137,10 @@ func TestChatCompletion_PreservesEdenUpstreamProvider(t *testing.T) { resp := chatResponseFrom(t, edenChatResponse) raw := resp.ExtraFields.Lookup(upstreamProviderField) - if len(raw) == 0 { - t.Fatalf("ExtraFields missing %q", upstreamProviderField) - } + require.NotEmpty(t, raw, "ExtraFields missing %q", upstreamProviderField) var upstream string - if err := json.Unmarshal(raw, &upstream); err != nil { - t.Fatalf("Unmarshal(%s) error = %v", raw, err) - } - if upstream != "openai" { - t.Errorf("%s = %q, want openai", upstreamProviderField, upstream) - } + require.NoError(t, json.Unmarshal(raw, &upstream), "Unmarshal(%s)", raw) + assert.Equal(t, "openai", upstream, "%s = %q, want openai", upstreamProviderField, upstream) } // TestChatCompletion_NoUpstreamProviderLeavesNoMarker asserts a response @@ -169,12 +149,9 @@ func TestChatCompletion_NoUpstreamProviderLeavesNoMarker(t *testing.T) { resp := chatResponseFrom(t, `{"id":"c","created":1,"model":"m","choices":[], "usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`) - if resp.Provider != "" { - t.Errorf("Provider = %q, want empty", resp.Provider) - } - if raw := resp.ExtraFields.Lookup(upstreamProviderField); len(raw) != 0 { - t.Errorf("ExtraFields[%q] = %s, want absent", upstreamProviderField, raw) - } + assert.Empty(t, resp.Provider, "Provider = %q, want empty", resp.Provider) + raw := resp.ExtraFields.Lookup(upstreamProviderField) + assert.Empty(t, raw, "ExtraFields[%q] = %s, want absent", upstreamProviderField, raw) } // TestEmbeddings_DoesNotReportEdenUpstreamAsExecutingProvider applies the same @@ -188,17 +165,13 @@ func TestEmbeddings_DoesNotReportEdenUpstreamAsExecutingProvider(t *testing.T) { })) defer server.Close() - provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + provider := newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.Embeddings(context.Background(), &core.EmbeddingRequest{ Model: "openai/text-embedding-3-small", Input: "hello", }) - if err != nil { - t.Fatalf("Embeddings() error = %v", err) - } - if resp.Provider != "" { - t.Fatalf("Provider = %q, want empty so the gateway labels the request edenai", resp.Provider) - } + require.NoError(t, err) + require.Empty(t, resp.Provider, "Provider = %q, want empty so the gateway labels the request edenai", resp.Provider) } // TestResponses_InheritsCostLifting asserts the Responses surface picks up the @@ -210,18 +183,15 @@ func TestResponses_InheritsCostLifting(t *testing.T) { })) defer server.Close() - provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + provider := newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{ Model: slashedModel, Input: "hi", }) - if err != nil { - t.Fatalf("Responses() error = %v", err) - } + require.NoError(t, err) value, ok := resp.Usage.RawUsage["cost"].(float64) - if !ok || math.Abs(value-0.0002349) > 1e-12 { - t.Fatalf("Responses usage cost = %#v, want 0.0002349 carried through the chat translation", resp.Usage.RawUsage["cost"]) - } + require.True(t, ok, "Responses usage cost = %#v, want 0.0002349 carried through the chat translation", resp.Usage.RawUsage["cost"]) + require.InDelta(t, 0.0002349, value, 1e-12, "Responses usage cost = %#v, want 0.0002349 carried through the chat translation", resp.Usage.RawUsage["cost"]) } // TestChatCompletion_PropagatesUpstreamErrorWithoutNormalizing asserts an @@ -235,16 +205,12 @@ func TestChatCompletion_PropagatesUpstreamErrorWithoutNormalizing(t *testing.T) })) defer server.Close() - provider := NewWithHTTPClient("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + provider := newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) resp, err := provider.ChatCompletion(context.Background(), &core.ChatRequest{ Model: slashedModel, Messages: []core.Message{{Role: "user", Content: "hi"}}, }) - if err == nil { - t.Fatal("ChatCompletion() error = nil, want the upstream error propagated") - } - if resp != nil { - t.Errorf("response = %+v, want nil alongside the error", resp) - } + require.Error(t, err, "ChatCompletion() error = nil, want the upstream error propagated") + assert.Nil(t, resp, "response = %+v, want nil alongside the error", resp) } diff --git a/internal/providers/edenai/transport.go b/internal/providers/edenai/transport.go index 2f4ee4616..bb52e7ca0 100644 --- a/internal/providers/edenai/transport.go +++ b/internal/providers/edenai/transport.go @@ -62,8 +62,9 @@ func isLoopbackHost(host string) bool { // policy drops Authorization only when the redirect target is a different // host; it does not look at the scheme, so an HTTPS -> HTTP redirect back to // the same host forwards the credential in the clear (verified against -// net/http, not assumed). Both of this provider's construction paths route -// through here so the guarantee does not depend on which one a caller used. +// net/http, not assumed). Every client the provider is built on — the gateway +// default and a caller-supplied override alike — routes through here, so the +// guarantee does not depend on how the provider was constructed. // // base is the caller-supplied client, or nil for the gateway default client // (the tuned transport and timeouts llmclient would otherwise install). It is diff --git a/internal/providers/edenai/transport_test.go b/internal/providers/edenai/transport_test.go index cdd2fa7d2..fec701088 100644 --- a/internal/providers/edenai/transport_test.go +++ b/internal/providers/edenai/transport_test.go @@ -13,6 +13,9 @@ import ( "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/llmclient" "github.com/enterpilot/gomodel/internal/providers" + "github.com/enterpilot/gomodel/internal/providers/providertest" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) // embeddingRequest is the smallest request that reaches the wire, used by the @@ -46,12 +49,8 @@ func TestCredentialSafeURL(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { parsed, err := url.Parse(tc.raw) - if err != nil { - t.Fatalf("url.Parse(%q) = %v", tc.raw, err) - } - if got := secureDestination(parsed); got != tc.want { - t.Errorf("secureDestination(%q) = %v, want %v", tc.raw, got, tc.want) - } + require.NoError(t, err, "url.Parse(%q)", tc.raw) + assert.Equal(t, tc.want, secureDestination(parsed), "secureDestination(%q)", tc.raw) }) } } @@ -59,23 +58,17 @@ func TestCredentialSafeURL(t *testing.T) { // TestCredentialSafeURL_NilIsUnsafe asserts the predicate fails closed. A // request with no parsable URL must not be treated as a TLS destination. func TestCredentialSafeURL_NilIsUnsafe(t *testing.T) { - if secureDestination(nil) { - t.Error("secureDestination(nil) = true, want false: the predicate must fail closed") - } + assert.False(t, secureDestination(nil), "secureDestination(nil) = true, want false: the predicate must fail closed") } // TestSetHeaders_SendsCredentialOverHTTPS asserts the ordinary case still // authenticates: an HTTPS destination gets the bearer token. func TestSetHeaders_SendsCredentialOverHTTPS(t *testing.T) { req, err := http.NewRequest(http.MethodPost, defaultBaseURL+"/chat/completions", nil) - if err != nil { - t.Fatalf("http.NewRequest: %v", err) - } + require.NoError(t, err, "http.NewRequest") setHeaders(req, "eden-key") - if got := req.Header.Get("Authorization"); got != "Bearer eden-key" { - t.Errorf("Authorization = %q, want %q", got, "Bearer eden-key") - } + assert.Equal(t, "Bearer eden-key", req.Header.Get("Authorization")) } // TestSetHeaders_WithholdsCredentialOverCleartext is the base-URL half of the @@ -83,14 +76,10 @@ func TestSetHeaders_SendsCredentialOverHTTPS(t *testing.T) { // Eden key on the wire in plain text. func TestSetHeaders_WithholdsCredentialOverCleartext(t *testing.T) { req, err := http.NewRequest(http.MethodPost, "http://eden.example.com/v3/chat/completions", nil) - if err != nil { - t.Fatalf("http.NewRequest: %v", err) - } + require.NoError(t, err, "http.NewRequest") setHeaders(req, "eden-key") - if got := req.Header.Get("Authorization"); got != "" { - t.Errorf("Authorization = %q, want empty: the credential must not be sent in cleartext", got) - } + assert.Empty(t, req.Header.Get("Authorization"), "the credential must not be sent in cleartext") } // TestSetHeaders_SendsCredentialOverLoopback pins the exemption that keeps a @@ -98,14 +87,10 @@ func TestSetHeaders_WithholdsCredentialOverCleartext(t *testing.T) { // working. Cleartext to the local machine never reaches a network. func TestSetHeaders_SendsCredentialOverLoopback(t *testing.T) { req, err := http.NewRequest(http.MethodPost, "http://127.0.0.1:9999/v3/chat/completions", nil) - if err != nil { - t.Fatalf("http.NewRequest: %v", err) - } + require.NoError(t, err, "http.NewRequest") setHeaders(req, "eden-key") - if got := req.Header.Get("Authorization"); got != "Bearer eden-key" { - t.Errorf("Authorization = %q, want %q", got, "Bearer eden-key") - } + assert.Equal(t, "Bearer eden-key", req.Header.Get("Authorization")) } // TestCleartextBaseURL_RequestRefusedBeforeSending is the payload half of the @@ -134,33 +119,26 @@ func TestCleartextBaseURL_RequestRefusedBeforeSending(t *testing.T) { defer server.Close() target, err := url.Parse(server.URL) - if err != nil { - t.Fatalf("url.Parse: %v", err) - } + require.NoError(t, err, "url.Parse") client := &http.Client{Transport: &cleartextRouteTransport{cleartext: target.Host}} - provider := NewWithHTTPClient("eden-key", "http://eden.example.com/v3", client, llmclient.Hooks{}) + provider := newTestProvider("eden-key", "http://eden.example.com/v3", client, llmclient.Hooks{}) _, err = provider.ChatCompletion(context.Background(), &core.ChatRequest{ Model: slashedModel, Messages: []core.Message{{Role: "user", Content: "secret-prompt"}}, }) - if err == nil { - t.Fatal("ChatCompletion succeeded against a cleartext endpoint, want the request refused") - } + require.Error(t, err, "ChatCompletion succeeded against a cleartext endpoint, want the request refused") - if hits != 0 { - t.Errorf("cleartext endpoint received %d request(s), want 0", hits) - } - if gotAuth != "" { - t.Errorf("cleartext endpoint saw Authorization = %q, want empty", gotAuth) - } - if strings.Contains(gotBody, "secret-prompt") { - t.Errorf("cleartext endpoint received the prompt body %q; the payload must never be sent", gotBody) - } - // The refusal must name the destination without quoting the whole URL. - if !strings.Contains(err.Error(), "eden.example.com") { - t.Errorf("error %q should name the refused host", err) - } + assert.Zero(t, hits, "cleartext endpoint received requests, want 0") + assert.Empty(t, gotAuth, "cleartext endpoint saw an Authorization header, want empty") + assert.NotContains(t, gotBody, "secret-prompt", "cleartext endpoint received the prompt body; the payload must never be sent") + // The refusal must name the destination without quoting the whole URL. It + // is the guard's own error, reached through the *url.Error net/http wraps + // it in: llmclient deliberately keeps upstream details out of the + // client-facing message and retains the cause only through Unwrap. + var urlErr *url.Error + require.ErrorAs(t, err, &urlErr, "want the transport refusal retained as the cause") + assert.ErrorContains(t, urlErr.Err, "eden.example.com", "the refusal should name the refused host") } // TestCleartextLoopback_RequestStillSent pins the other side of the exemption: @@ -178,16 +156,11 @@ func TestCleartextLoopback_RequestStillSent(t *testing.T) { })) defer server.Close() - provider := NewWithHTTPClient("eden-key", server.URL, server.Client(), llmclient.Hooks{}) - if _, err := provider.Embeddings(context.Background(), embeddingRequest()); err != nil { - t.Fatalf("Embeddings against a loopback endpoint: %v", err) - } - if hits != 1 { - t.Errorf("loopback endpoint received %d request(s), want 1", hits) - } - if gotAuth != "Bearer eden-key" { - t.Errorf("loopback endpoint saw Authorization = %q, want %q", gotAuth, "Bearer eden-key") - } + provider := newTestProvider("eden-key", server.URL, server.Client(), llmclient.Hooks{}) + _, err := provider.Embeddings(context.Background(), embeddingRequest()) + require.NoError(t, err, "Embeddings against a loopback endpoint") + assert.Equal(t, 1, hits, "loopback endpoint request count") + assert.Equal(t, "Bearer eden-key", gotAuth, "loopback endpoint Authorization header") } // TestSecureTransport_RoundTrip covers the guard directly, including that an @@ -211,42 +184,35 @@ func TestSecureTransport_RoundTrip(t *testing.T) { transport := &secureTransport{base: stub} req, err := http.NewRequest(http.MethodGet, tc.target, nil) - if err != nil { - t.Fatalf("http.NewRequest: %v", err) - } + require.NoError(t, err, "http.NewRequest") resp, err := transport.RoundTrip(req) if resp != nil && resp.Body != nil { _ = resp.Body.Close() } - if stub.calls != 0 != tc.wantCalled { - t.Errorf("underlying transport calls = %d, wantCalled = %v", stub.calls, tc.wantCalled) - } - if tc.wantCalled && err != nil { - t.Errorf("RoundTrip(%q) = %v, want the request passed through", tc.target, err) - } - if !tc.wantCalled && err == nil { - t.Errorf("RoundTrip(%q) = nil error, want a refusal", tc.target) + if tc.wantCalled { + assert.NotZero(t, stub.calls, "underlying transport was never called, want the request passed through") + require.NoError(t, err, "RoundTrip(%q), want the request passed through", tc.target) + } else { + assert.Zero(t, stub.calls, "underlying transport was called, want the request refused before it") + require.Error(t, err, "RoundTrip(%q) = nil error, want a refusal", tc.target) } }) } } // TestSecureTransport_NilBaseUsesDefaultTransport asserts the nil-base fallback -// is wired, since http.DefaultClient carries a nil Transport and Eden's guard -// wraps exactly that on the NewWithHTTPClient nil path. +// is wired: http.DefaultClient carries a nil Transport, and Eden's guard wraps +// exactly that when a caller hands it in through ProviderOptions.HTTPClient. func TestSecureTransport_NilBaseUsesDefaultTransport(t *testing.T) { transport := &secureTransport{} // A refused destination never reaches the base, so it proves the guard runs // without needing a live server for the delegating case. req, err := http.NewRequest(http.MethodGet, "http://eden.example.com/v3", nil) - if err != nil { - t.Fatalf("http.NewRequest: %v", err) - } - if _, err := transport.RoundTrip(req); err == nil { - t.Error("RoundTrip = nil error for a cleartext target, want a refusal") - } + require.NoError(t, err, "http.NewRequest") + _, err = transport.RoundTrip(req) + require.Error(t, err, "RoundTrip = nil error for a cleartext target, want a refusal") // An allowed loopback destination must reach the network through // http.DefaultTransport rather than panicking on the nil base. @@ -256,17 +222,11 @@ func TestSecureTransport_NilBaseUsesDefaultTransport(t *testing.T) { defer server.Close() req, err = http.NewRequest(http.MethodGet, server.URL, nil) - if err != nil { - t.Fatalf("http.NewRequest: %v", err) - } + require.NoError(t, err, "http.NewRequest") resp, err := transport.RoundTrip(req) - if err != nil { - t.Fatalf("RoundTrip through the nil base: %v", err) - } + require.NoError(t, err, "RoundTrip through the nil base") defer resp.Body.Close() - if resp.StatusCode != http.StatusNoContent { - t.Errorf("status = %d, want %d", resp.StatusCode, http.StatusNoContent) - } + assert.Equal(t, http.StatusNoContent, resp.StatusCode) } // recordingRoundTripper counts the requests that made it past the guard. @@ -305,11 +265,12 @@ func (t *cleartextRouteTransport) RoundTrip(req *http.Request) (*http.Response, return base.RoundTrip(routed) } -// TestCompatibleConfig_InstallsRedirectGuardOnBothPaths asserts neither -// construction path can reach the network without the redirect policy. Both -// New and NewWithHTTPClient build their transport through compatibleConfig, so -// covering it here covers both. -func TestCompatibleConfig_InstallsRedirectGuardOnBothPaths(t *testing.T) { +// TestCompatibleConfig_InstallsRedirectGuardForDefaultAndCallerClients asserts +// neither kind of client can reach the network without the redirect policy. +// New builds its transport through compatibleConfig whether opts.HTTPClient is +// nil (the gateway default) or caller-supplied, so covering compatibleConfig +// here covers both. +func TestCompatibleConfig_InstallsRedirectGuardForDefaultAndCallerClients(t *testing.T) { tests := []struct { name string client *http.Client @@ -321,12 +282,8 @@ func TestCompatibleConfig_InstallsRedirectGuardOnBothPaths(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { cfg := compatibleConfig(defaultBaseURL, tc.client) - if cfg.HTTPClient == nil { - t.Fatal("HTTPClient = nil, want a client carrying the redirect guard") - } - if cfg.HTTPClient.CheckRedirect == nil { - t.Error("CheckRedirect = nil, want Eden's redirect guard installed") - } + require.NotNil(t, cfg.HTTPClient, "HTTPClient = nil, want a client carrying the redirect guard") + assert.NotNil(t, cfg.HTTPClient.CheckRedirect, "CheckRedirect = nil, want Eden's redirect guard installed") }) } } @@ -338,15 +295,9 @@ func TestGuardedHTTPClient_DoesNotMutateCallerClient(t *testing.T) { caller := &http.Client{} guarded := guardedHTTPClient(caller) - if caller.CheckRedirect != nil { - t.Error("caller's CheckRedirect was set; the guard must be installed on a copy") - } - if guarded == caller { - t.Error("guardedHTTPClient returned the caller's client; want a copy") - } - if guarded.CheckRedirect == nil { - t.Error("guarded client has no CheckRedirect") - } + assert.Nil(t, caller.CheckRedirect, "caller's CheckRedirect was set; the guard must be installed on a copy") + assert.NotSame(t, caller, guarded, "guardedHTTPClient returned the caller's client; want a copy") + assert.NotNil(t, guarded.CheckRedirect, "guarded client has no CheckRedirect") } // TestGuardedHTTPClient_WrapsCallerTransport asserts the guard is layered over @@ -356,16 +307,11 @@ func TestGuardedHTTPClient_WrapsCallerTransport(t *testing.T) { transport := &http.Transport{} caller := &http.Client{Transport: transport} - guarded, ok := guardedHTTPClient(caller).Transport.(*secureTransport) - if !ok { - t.Fatalf("Transport = %T, want *secureTransport wrapping the caller's", guardedHTTPClient(caller).Transport) - } - if guarded.base != transport { - t.Errorf("wrapped base = %v, want the caller's transport preserved", guarded.base) - } - if caller.Transport != transport { - t.Error("the caller's client was modified; the guard must go on a copy") - } + guarded := guardedHTTPClient(caller) + wrapped, ok := guarded.Transport.(*secureTransport) + require.True(t, ok, "Transport = %T, want *secureTransport wrapping the caller's", guarded.Transport) + assert.Same(t, transport, wrapped.base, "wrapped base should be the caller's transport preserved") + assert.Same(t, transport, caller.Transport, "the caller's client was modified; the guard must go on a copy") } // TestCheckRedirect_AllowsSameHostHTTPS asserts the one redirect shape a REST @@ -375,14 +321,10 @@ func TestCheckRedirect_AllowsSameHostHTTPS(t *testing.T) { origin := mustRequest(t, "https://api.edenai.run/v3/chat/completions") target := mustRequest(t, "https://api.edenai.run/v3/chat/completions-moved") - if err := checkRedirect(target, []*http.Request{origin}); err != nil { - t.Errorf("checkRedirect same-host = %v, want nil", err) - } + require.NoError(t, checkRedirect(target, []*http.Request{origin}), "checkRedirect same-host") // Host comparison is case-insensitive, as hostnames are. upper := mustRequest(t, "https://API.EdenAI.run/v3/chat/completions-moved") - if err := checkRedirect(upper, []*http.Request{origin}); err != nil { - t.Errorf("checkRedirect differing-case host = %v, want nil", err) - } + require.NoError(t, checkRedirect(upper, []*http.Request{origin}), "checkRedirect differing-case host") } // TestCheckRedirect_RefusesCrossHostHTTPS covers the credential-exposure path @@ -400,13 +342,7 @@ func TestCheckRedirect_RefusesCrossHostHTTPS(t *testing.T) { } { req := mustRequest(t, target) err := checkRedirect(req, []*http.Request{origin}) - if err == nil { - t.Errorf("checkRedirect(%q) = nil, want a refusal", target) - continue - } - if !strings.Contains(err.Error(), "cross-host") { - t.Errorf("checkRedirect(%q) = %v, want a cross-host refusal", target, err) - } + assert.ErrorContains(t, err, "cross-host", "checkRedirect(%q), want a cross-host refusal", target) } } @@ -418,9 +354,7 @@ func TestCheckRedirect_ComparesAgainstOriginalHost(t *testing.T) { hop := mustRequest(t, "https://api.edenai.run/v3/models-moved") target := mustRequest(t, "https://elsewhere.edenai.run/v3/models") - if err := checkRedirect(target, []*http.Request{origin, hop}); err == nil { - t.Error("checkRedirect = nil for a second hop leaving the original host, want a refusal") - } + require.Error(t, checkRedirect(target, []*http.Request{origin, hop}), "checkRedirect = nil for a second hop leaving the original host, want a refusal") } // TestRedirectPolicy_PreservesCallerPolicy asserts the guard composes with the @@ -439,12 +373,9 @@ func TestRedirectPolicy_PreservesCallerPolicy(t *testing.T) { }) // Eden allows this same-host hop, so the caller's policy decides. - if err := policy(target, []*http.Request{origin}); !errors.Is(err, callerErr) { - t.Errorf("policy = %v, want the caller's error returned verbatim", err) - } - if called != 1 { - t.Errorf("caller policy invoked %d times, want 1", called) - } + err := policy(target, []*http.Request{origin}) + require.ErrorIs(t, err, callerErr, "want the caller's error returned verbatim") + assert.Equal(t, 1, called, "caller policy invocation count") } // TestRedirectPolicy_PropagatesErrUseLastResponse asserts the sentinel a caller @@ -458,9 +389,8 @@ func TestRedirectPolicy_PropagatesErrUseLastResponse(t *testing.T) { return http.ErrUseLastResponse }) - if err := policy(target, []*http.Request{origin}); !errors.Is(err, http.ErrUseLastResponse) { - t.Errorf("policy = %v, want http.ErrUseLastResponse propagated", err) - } + err := policy(target, []*http.Request{origin}) + assert.ErrorIs(t, err, http.ErrUseLastResponse, "want http.ErrUseLastResponse propagated") } // TestRedirectPolicy_RecheckAfterCallerMutation is the check-then-mutate case. @@ -511,12 +441,8 @@ func TestRedirectPolicy_RecheckAfterCallerMutation(t *testing.T) { }) err := policy(target, []*http.Request{origin}) - if err == nil { - t.Fatalf("policy = nil, want a refusal after the callback rewrote the target to %s", target.URL) - } - if !strings.Contains(err.Error(), tc.wantErr) { - t.Errorf("policy = %v, want an error mentioning %q", err, tc.wantErr) - } + require.Error(t, err, "policy = nil, want a refusal after the callback rewrote the target to %s", target.URL) + assert.ErrorContains(t, err, tc.wantErr, "want an error mentioning %q", tc.wantErr) }) } } @@ -533,9 +459,7 @@ func TestRedirectPolicy_AllowsHarmlessCallerMutation(t *testing.T) { return nil }) - if err := policy(target, []*http.Request{origin}); err != nil { - t.Errorf("policy = %v, want nil: a same-host path rewrite is fine", err) - } + require.NoError(t, policy(target, []*http.Request{origin}), "a same-host path rewrite is fine") } // TestRedirectPolicy_CallerErrorsPropagateUnchanged asserts every non-nil @@ -562,9 +486,8 @@ func TestRedirectPolicy_CallerErrorsPropagateUnchanged(t *testing.T) { policy := redirectPolicy(func(*http.Request, []*http.Request) error { return tc.err }) - if err := policy(target, []*http.Request{origin}); !errors.Is(err, tc.err) { - t.Errorf("policy = %v, want %v returned unchanged", err, tc.err) - } + err := policy(target, []*http.Request{origin}) + assert.ErrorIs(t, err, tc.err, "want %v returned unchanged", tc.err) }) } } @@ -581,9 +504,8 @@ func TestRedirectPolicy_CallerErrorSurvivesAMutation(t *testing.T) { return http.ErrUseLastResponse }) - if err := policy(target, []*http.Request{origin}); !errors.Is(err, http.ErrUseLastResponse) { - t.Errorf("policy = %v, want http.ErrUseLastResponse propagated unchanged", err) - } + err := policy(target, []*http.Request{origin}) + assert.ErrorIs(t, err, http.ErrUseLastResponse, "want http.ErrUseLastResponse propagated unchanged") } // TestRedirectPolicy_EdenRefusalWinsOverPermissiveCaller asserts a caller @@ -603,13 +525,9 @@ func TestRedirectPolicy_EdenRefusalWinsOverPermissiveCaller(t *testing.T) { "http://api.edenai.run/v3/models", // cleartext downgrade "https://evil.api.edenai.run/v3/models", // cross-host } { - if err := policy(mustRequest(t, target), []*http.Request{origin}); err == nil { - t.Errorf("policy(%q) = nil, want Eden's refusal to stand", target) - } - } - if called != 0 { - t.Errorf("caller policy invoked %d times, want 0: Eden refuses before delegating", called) + require.Error(t, policy(mustRequest(t, target), []*http.Request{origin}), "policy(%q) = nil, want Eden's refusal to stand", target) } + assert.Zero(t, called, "caller policy was invoked, want 0 calls: Eden refuses before delegating") } // TestRedirectPolicy_NilCallerAllowsEdenApprovedHop asserts the common case — @@ -618,9 +536,7 @@ func TestRedirectPolicy_NilCallerAllowsEdenApprovedHop(t *testing.T) { origin := mustRequest(t, "https://api.edenai.run/v3/models") target := mustRequest(t, "https://api.edenai.run/v3/models-moved") - if err := redirectPolicy(nil)(target, []*http.Request{origin}); err != nil { - t.Errorf("redirectPolicy(nil) = %v, want nil for a same-host TLS hop", err) - } + require.NoError(t, redirectPolicy(nil)(target, []*http.Request{origin}), "redirectPolicy(nil), want nil for a same-host TLS hop") } // TestGuardedHTTPClient_ComposesCallerRedirectPolicy asserts the composition is @@ -636,21 +552,15 @@ func TestGuardedHTTPClient_ComposesCallerRedirectPolicy(t *testing.T) { origin := mustRequest(t, "https://api.edenai.run/v3/models") target := mustRequest(t, "https://api.edenai.run/v3/models-moved") - if err := guarded.CheckRedirect(target, []*http.Request{origin}); err != nil { - t.Fatalf("CheckRedirect = %v, want nil", err) - } - if !called { - t.Error("the caller's redirect policy was not invoked; the guard must compose, not replace") - } + require.NoError(t, guarded.CheckRedirect(target, []*http.Request{origin}), "CheckRedirect") + assert.True(t, called, "the caller's redirect policy was not invoked; the guard must compose, not replace") } // mustRequest builds a GET request for a URL a test controls. func mustRequest(t *testing.T, rawURL string) *http.Request { t.Helper() req, err := http.NewRequest(http.MethodGet, rawURL, nil) - if err != nil { - t.Fatalf("http.NewRequest(%q): %v", rawURL, err) - } + require.NoError(t, err, "http.NewRequest(%q)", rawURL) return req } @@ -660,35 +570,23 @@ func mustRequest(t *testing.T, rawURL string) *http.Request { // redirect would otherwise forward the bearer token in the clear. func TestCheckRedirect_RefusesSchemeDowngrade(t *testing.T) { req, err := http.NewRequest(http.MethodGet, "http://api.edenai.run/v3/chat/completions", nil) - if err != nil { - t.Fatalf("http.NewRequest: %v", err) - } + require.NoError(t, err, "http.NewRequest") err = checkRedirect(req, nil) - if err == nil { - t.Fatal("checkRedirect = nil, want a refusal for a cleartext redirect target") - } - if !strings.Contains(err.Error(), "api.edenai.run") { - t.Errorf("error %q should name the refused host", err) - } + require.Error(t, err, "checkRedirect = nil, want a refusal for a cleartext redirect target") + assert.ErrorContains(t, err, "api.edenai.run", "the error should name the refused host") } // TestCheckRedirect_ErrorOmitsCredentials asserts the refusal message cannot // leak a secret carried in the redirect URL's userinfo or query string. func TestCheckRedirect_ErrorOmitsCredentials(t *testing.T) { req, err := http.NewRequest(http.MethodGet, "http://user:s3cret@api.edenai.run/v3?api_key=leaked", nil) - if err != nil { - t.Fatalf("http.NewRequest: %v", err) - } + require.NoError(t, err, "http.NewRequest") err = checkRedirect(req, nil) - if err == nil { - t.Fatal("checkRedirect = nil, want a refusal") - } + require.Error(t, err, "checkRedirect = nil, want a refusal") for _, secret := range []string{"s3cret", "leaked"} { - if strings.Contains(err.Error(), secret) { - t.Errorf("error %q leaked %q from the redirect URL", err, secret) - } + assert.NotContains(t, err.Error(), secret, "error leaked %q from the redirect URL", secret) } } @@ -696,14 +594,10 @@ func TestCheckRedirect_ErrorOmitsCredentials(t *testing.T) { // silently remove net/http's own protection against endless redirect chains. func TestCheckRedirect_EnforcesRedirectBudget(t *testing.T) { req, err := http.NewRequest(http.MethodGet, defaultBaseURL, nil) - if err != nil { - t.Fatalf("http.NewRequest: %v", err) - } + require.NoError(t, err, "http.NewRequest") via := make([]*http.Request, maxRedirects) - if err := checkRedirect(req, via); err == nil { - t.Fatalf("checkRedirect with %d prior hops = nil, want the redirect budget enforced", maxRedirects) - } + require.Error(t, checkRedirect(req, via), "checkRedirect with %d prior hops = nil, want the redirect budget enforced", maxRedirects) } // TestRedirect_HTTPSToHTTPDoesNotForwardCredential exercises the guard through @@ -731,27 +625,19 @@ func TestRedirect_HTTPSToHTTPDoesNotForwardCredential(t *testing.T) { defer secure.Close() cleartextTarget, err := url.Parse(cleartext.URL) - if err != nil { - t.Fatalf("url.Parse: %v", err) - } + require.NoError(t, err, "url.Parse") client := secure.Client() client.Transport = &cleartextRouteTransport{ base: client.Transport, cleartext: cleartextTarget.Host, } - provider := NewWithHTTPClient("eden-key", secure.URL+"/v3", client, llmclient.Hooks{}) + provider := newTestProvider("eden-key", secure.URL+"/v3", client, llmclient.Hooks{}) _, err = provider.Embeddings(context.Background(), embeddingRequest()) - if err == nil { - t.Fatal("Embeddings succeeded through an HTTPS -> HTTP redirect, want the redirect refused") - } + require.Error(t, err, "Embeddings succeeded through an HTTPS -> HTTP redirect, want the redirect refused") - if cleartextHits != 0 { - t.Errorf("cleartext endpoint received %d request(s), want 0", cleartextHits) - } - if cleartextAuth != "" { - t.Errorf("cleartext endpoint saw Authorization = %q, want empty", cleartextAuth) - } + assert.Zero(t, cleartextHits, "cleartext endpoint received requests, want 0") + assert.Empty(t, cleartextAuth, "cleartext endpoint saw an Authorization header, want empty") } // TestRedirect_CleartextLoopbackStillFollowed pins the exemption's scope: a @@ -774,16 +660,11 @@ func TestRedirect_CleartextLoopbackStillFollowed(t *testing.T) { server := httptest.NewServer(mux) defer server.Close() - provider := NewWithHTTPClient("eden-key", server.URL+"/v3", server.Client(), llmclient.Hooks{}) - if _, err := provider.Embeddings(context.Background(), embeddingRequest()); err != nil { - t.Fatalf("Embeddings through a loopback redirect: %v", err) - } - if !served { - t.Fatal("redirect target was never reached") - } - if finalAuth != "Bearer eden-key" { - t.Errorf("redirect target saw Authorization = %q, want %q", finalAuth, "Bearer eden-key") - } + provider := newTestProvider("eden-key", server.URL+"/v3", server.Client(), llmclient.Hooks{}) + _, err := provider.Embeddings(context.Background(), embeddingRequest()) + require.NoError(t, err, "Embeddings through a loopback redirect") + require.True(t, served, "redirect target was never reached") + assert.Equal(t, "Bearer eden-key", finalAuth, "redirect target Authorization header") } // TestRedirect_HTTPSToHTTPSStillFollowed asserts the guard is narrow: a @@ -806,20 +687,12 @@ func TestRedirect_HTTPSToHTTPSStillFollowed(t *testing.T) { secure := httptest.NewTLSServer(mux) defer secure.Close() - provider := NewWithHTTPClient("eden-key", secure.URL+"/v3", secure.Client(), llmclient.Hooks{}) + provider := newTestProvider("eden-key", secure.URL+"/v3", secure.Client(), llmclient.Hooks{}) resp, err := provider.Embeddings(context.Background(), embeddingRequest()) - if err != nil { - t.Fatalf("Embeddings through an HTTPS -> HTTPS redirect: %v", err) - } - if !served { - t.Fatal("redirect target was never reached") - } - if resp == nil { - t.Fatal("response = nil") - } - if finalAuth != "Bearer eden-key" { - t.Errorf("redirect target saw Authorization = %q, want %q", finalAuth, "Bearer eden-key") - } + require.NoError(t, err, "Embeddings through an HTTPS -> HTTPS redirect") + require.True(t, served, "redirect target was never reached") + require.NotNil(t, resp, "response = nil") + assert.Equal(t, "Bearer eden-key", finalAuth, "redirect target Authorization header") } // TestRedirect_HTTPSSubdomainDoesNotForwardCredential is the end-to-end form of @@ -847,26 +720,18 @@ func TestRedirect_HTTPSSubdomainDoesNotForwardCredential(t *testing.T) { defer server.Close() target, err := url.Parse(server.URL) - if err != nil { - t.Fatalf("url.Parse: %v", err) - } + require.NoError(t, err, "url.Parse") client := server.Client() // Route both hostnames to the one test server; the TLS config from // server.Client() already trusts its certificate. client.Transport = &hostPinnedTransport{base: client.Transport, addr: target.Host} - provider := NewWithHTTPClient("eden-key", "https://eden.test/v3", client, llmclient.Hooks{}) + provider := newTestProvider("eden-key", "https://eden.test/v3", client, llmclient.Hooks{}) _, err = provider.Embeddings(context.Background(), embeddingRequest()) - if err == nil { - t.Fatal("Embeddings followed an HTTPS subdomain redirect, want it refused") - } + require.Error(t, err, "Embeddings followed an HTTPS subdomain redirect, want it refused") - if secondHopHits != 0 { - t.Errorf("subdomain endpoint received %d request(s), want 0", secondHopHits) - } - if secondHopAuth != "" { - t.Errorf("subdomain endpoint saw Authorization = %q, want empty: the Eden key must not follow an upstream-chosen host", secondHopAuth) - } + assert.Zero(t, secondHopHits, "subdomain endpoint received requests, want 0") + assert.Empty(t, secondHopAuth, "subdomain endpoint saw an Authorization header, want empty: the Eden key must not follow an upstream-chosen host") } // hostPinnedTransport dials one fixed address whatever hostname the request @@ -893,9 +758,7 @@ func (t *hostPinnedTransport) RoundTrip(req *http.Request) (*http.Response, erro // http:// would silently withhold the API key on ordinary traffic. // TestNew_DefaultsBaseURL already covers New resolving to this value. func TestDefaultBaseURLIsHTTPS(t *testing.T) { - if !strings.HasPrefix(defaultBaseURL, "https://") { - t.Fatalf("defaultBaseURL = %q, want an https:// endpoint", defaultBaseURL) - } + require.True(t, strings.HasPrefix(defaultBaseURL, "https://"), "defaultBaseURL = %q, want an https:// endpoint", defaultBaseURL) } // TestSetBaseURL_OverridesResolvedEndpoint covers the public base-URL override, @@ -903,89 +766,80 @@ func TestDefaultBaseURLIsHTTPS(t *testing.T) { // endpoint the provider was constructed with. func TestSetBaseURL_OverridesResolvedEndpoint(t *testing.T) { provider, ok := New(providers.ProviderConfig{APIKey: "eden-key"}, providers.ProviderOptions{}).(*Provider) - if !ok { - t.Fatal("New did not return *Provider") - } + require.True(t, ok, "New did not return *Provider") const override = "https://eden.eu.example.com/v3" provider.SetBaseURL(override) - if got := provider.GetBaseURL(); got != override { - t.Errorf("GetBaseURL() = %q, want %q after SetBaseURL", got, override) - } + assert.Equal(t, override, provider.GetBaseURL(), "GetBaseURL() after SetBaseURL") } -// TestGuardedHTTPClient_PreservesDefaultClientSemantics pins what -// NewWithHTTPClient's nil path produces. Every other chat-compatible provider -// documents "if httpClient is nil, http.DefaultClient is used", and gets that -// from NewCompatibleProviderWithHTTPClient; Eden hands that helper an -// already-guarded client, so it substitutes http.DefaultClient itself. Guarding -// that client must leave its transport and its absent timeout alone, or the -// constructor would quietly diverge from every peer. +// TestGuardedHTTPClient_PreservesDefaultClientSemantics pins what happens when +// a caller hands New http.DefaultClient through ProviderOptions.HTTPClient. +// Guarding that client must leave its transport and its absent timeout alone: +// the caller chose those semantics, and the guard's only job is to add the +// redirect and cleartext policies on top of them. // // The guard also has to go on a copy: writing CheckRedirect onto // http.DefaultClient would change redirect behavior for every other user of // that global in the process. That is the assertion that matters most here. // -// The substitution itself is not observable from outside the provider (the +// The guarded client itself is not observable from outside the provider (the // transport lives on an unexported field of openai.CompatibleProvider), so it // is pinned by this test together with -// TestGuardedHTTPClient_NilBuildsGatewayDefault, which shows the two callers -// deliberately get different clients. +// TestGuardedHTTPClient_NilBuildsGatewayDefault, which shows the nil path +// deliberately gets a different client. func TestGuardedHTTPClient_PreservesDefaultClientSemantics(t *testing.T) { guarded := guardedHTTPClient(http.DefaultClient) - if guarded == http.DefaultClient { - t.Fatal("guardedHTTPClient returned http.DefaultClient itself; want a copy") - } - if http.DefaultClient.CheckRedirect != nil { - t.Error("http.DefaultClient.CheckRedirect was set; the process-wide client must not be modified") - } - if guarded.CheckRedirect == nil { - t.Error("CheckRedirect = nil, want Eden's redirect guard on the copy") - } + require.NotSame(t, http.DefaultClient, guarded, "guardedHTTPClient returned http.DefaultClient itself; want a copy") + assert.Nil(t, http.DefaultClient.CheckRedirect, "http.DefaultClient.CheckRedirect was set; the process-wide client must not be modified") + assert.NotNil(t, guarded.CheckRedirect, "CheckRedirect = nil, want Eden's redirect guard on the copy") wrapped, ok := guarded.Transport.(*secureTransport) - if !ok { - t.Fatalf("Transport = %T, want *secureTransport", guarded.Transport) - } + require.True(t, ok, "Transport = %T, want *secureTransport", guarded.Transport) // http.DefaultClient leaves Transport nil, meaning http.DefaultTransport; // the wrapper preserves that by delegating to it when its base is nil. - if wrapped.base != http.DefaultClient.Transport { - t.Errorf("wrapped base = %v, want http.DefaultClient's transport (nil)", wrapped.base) - } - if guarded.Timeout != http.DefaultClient.Timeout { - t.Errorf("Timeout = %v, want http.DefaultClient's %v, not the gateway client's", guarded.Timeout, http.DefaultClient.Timeout) - } + assert.Equal(t, http.DefaultClient.Transport, wrapped.base, "wrapped base should be http.DefaultClient's transport (nil)") + assert.Equal(t, http.DefaultClient.Timeout, guarded.Timeout, "Timeout should be http.DefaultClient's, not the gateway client's") } -// TestGuardedHTTPClient_NilBuildsGatewayDefault covers the other caller: New -// passes nil because the factory path has no client of its own, and must get -// the tuned transport and timeouts llmclient would otherwise have installed — -// not http.DefaultClient's absent timeout. +// TestGuardedHTTPClient_NilBuildsGatewayDefault covers the nil path: New +// passes opts.HTTPClient, which is nil unless the factory applied an outbound +// proxy or a test injected a client, and that path must produce the tuned +// gateway client — the transport and timeouts llmclient would otherwise have +// installed — not http.DefaultClient's absent timeout. func TestGuardedHTTPClient_NilBuildsGatewayDefault(t *testing.T) { guarded := guardedHTTPClient(nil) - if guarded.CheckRedirect == nil { - t.Error("CheckRedirect = nil, want Eden's redirect guard") - } - if guarded.Timeout <= 0 { - t.Errorf("Timeout = %v, want the gateway default client's positive timeout", guarded.Timeout) - } + assert.NotNil(t, guarded.CheckRedirect, "CheckRedirect = nil, want Eden's redirect guard") + assert.Positive(t, guarded.Timeout, "Timeout should be the gateway default client's positive timeout") } -// TestNewWithHTTPClient_NilClientStillGuardsRedirects proves the nil path is -// wired end to end: a provider built with no client still refuses a -// credential-leaking redirect rather than following it. -func TestNewWithHTTPClient_NilClientStillGuardsRedirects(t *testing.T) { - provider := NewWithHTTPClient("eden-key", "http://eden.example.com/v3", nil, llmclient.Hooks{}) - if provider == nil { - t.Fatal("NewWithHTTPClient(..., nil, ...) returned nil") - } +// TestNew_NilClientStillGuardsRedirects proves the nil-client path is wired end +// to end: a provider built with no client of its own runs on the guarded +// gateway default client, so a credential-leaking redirect is refused rather +// than followed. The guard refuses the hop before dialing, so no DNS lookup of +// the redirect target ever happens. +// +// llmclient treats the refusal as a transport error and retries it, so the +// loopback server sees one request per attempt; every attempt stops at the +// same guard, and none of them is a followed redirect. +func TestNew_NilClientStillGuardsRedirects(t *testing.T) { + var hits int + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hits++ + http.Redirect(w, r, "http://eden.example.com/v3/embeddings", http.StatusFound) + })) + defer server.Close() - req, err := http.NewRequest(http.MethodGet, "http://api.edenai.run/v3/chat/completions", nil) - if err != nil { - t.Fatalf("http.NewRequest: %v", err) - } - if err := checkRedirect(req, nil); err == nil { - t.Error("checkRedirect = nil for a cleartext target; the guard must apply on the nil-client path too") - } + provider := newTestProvider("eden-key", server.URL+"/v3", nil, llmclient.Hooks{}) + _, err := provider.Embeddings(context.Background(), embeddingRequest()) + require.Error(t, err, "Embeddings followed a cleartext redirect on the nil-client path, want it refused") + // The client-facing message hides upstream details by design; the guard's + // refusal is the cause net/http wrapped in a *url.Error. + var urlErr *url.Error + require.ErrorAs(t, err, &urlErr, "want the redirect refusal retained as the cause") + require.ErrorContains(t, urlErr.Err, "follow redirect", "want the redirect guard, not the request guard, to have refused") + require.ErrorContains(t, urlErr.Err, "cleartext", "want the redirect guard's cleartext refusal") + attempts := 1 + providertest.Resilience().Retry.MaxRetries + assert.Equal(t, attempts, hits, "the loopback server should see exactly the original request once per attempt") } diff --git a/internal/providers/registry_provider_pricing_test.go b/internal/providers/registry_provider_pricing_test.go index 7957abb35..6a91e7151 100644 --- a/internal/providers/registry_provider_pricing_test.go +++ b/internal/providers/registry_provider_pricing_test.go @@ -4,6 +4,9 @@ import ( "context" "testing" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/enterpilot/gomodel/config" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/modeldata" @@ -50,9 +53,7 @@ func TestInitialize_ProviderReportedPricingSurvivesCatalogMiss(t *testing.T) { } registry.RegisterProviderWithNameAndType(provider, "edenai", "edenai") - if err := registry.Initialize(context.Background()); err != nil { - t.Fatalf("Initialize: %v", err) - } + require.NoError(t, registry.Initialize(context.Background()), "Initialize") // A populated catalog that knows nothing about this provider: the exact // condition Eden models are always in. @@ -64,25 +65,18 @@ func TestInitialize_ProviderReportedPricingSurvivesCatalogMiss(t *testing.T) { }, nil, "etag", "https://example.invalid/models.json") info := registry.GetModel("edenai/openai/gpt-4") - if info == nil { - t.Fatal("model not registered under its provider-qualified ID") - } - if info.Discovered == nil || info.Discovered.Pricing == nil { - t.Fatalf("Discovered = %+v, want the provider's own report retained", info.Discovered) - } + require.NotNil(t, info, "model not registered under its provider-qualified ID") + require.NotNil(t, info.Discovered, "want the provider's own report retained") + require.NotNil(t, info.Discovered.Pricing, "Discovered = %+v, want the provider's own report retained", info.Discovered) meta := info.Model.Metadata - if meta == nil || meta.Pricing == nil { - t.Fatalf("Metadata = %+v, want provider pricing to survive enrichment", meta) - } + require.NotNil(t, meta, "want provider pricing to survive enrichment") + require.NotNil(t, meta.Pricing, "Metadata = %+v, want provider pricing to survive enrichment", meta) assertPricePtr(t, "InputPerMtok", meta.Pricing.InputPerMtok, 0.06) assertPricePtr(t, "OutputPerMtok", meta.Pricing.OutputPerMtok, 0.18) assertPricePtr(t, "CachedInputPerMtok", meta.Pricing.CachedInputPerMtok, 0.012) - if meta.ContextWindow == nil || *meta.ContextWindow != 131072 { - t.Errorf("ContextWindow = %v, want 131072", meta.ContextWindow) - } - if !meta.Capabilities["reasoning"] { - t.Errorf("Capabilities = %v, want reasoning retained", meta.Capabilities) - } + require.NotNil(t, meta.ContextWindow, "ContextWindow = nil, want 131072") + assert.Equal(t, 131072, *meta.ContextWindow) + assert.True(t, meta.Capabilities["reasoning"], "Capabilities = %v, want reasoning retained", meta.Capabilities) } // TestResolvePricing_UsesProviderReportedPricing closes the loop to the cost @@ -99,15 +93,11 @@ func TestResolvePricing_UsesProviderReportedPricing(t *testing.T) { } registry.RegisterProviderWithNameAndType(provider, "edenai", "edenai") - if err := registry.Initialize(context.Background()); err != nil { - t.Fatalf("Initialize: %v", err) - } + require.NoError(t, registry.Initialize(context.Background()), "Initialize") for _, selector := range []string{"edenai/openai/gpt-4", "openai/gpt-4"} { pricing := registry.ResolvePricing(selector, "edenai") - if pricing == nil { - t.Fatalf("ResolvePricing(%q) = nil, want the provider-reported rates", selector) - } + require.NotNil(t, pricing, "ResolvePricing(%q) = nil, want the provider-reported rates", selector) assertPricePtr(t, selector+" InputPerMtok", pricing.InputPerMtok, 0.06) assertPricePtr(t, selector+" OutputPerMtok", pricing.OutputPerMtok, 0.18) } @@ -122,24 +112,13 @@ func TestModelFilter_AdmitsProviderPricedModels(t *testing.T) { unpriced := core.Model{ID: "openai/gpt-5", Object: "model"} filter, active := newModelFilter(config.ModelFilter{MaxPricePerMtok: new(1.0)}) - if !active { - t.Fatal("newModelFilter reported an inactive filter for a price cap") - } - if !filter.keep(priced) { - t.Error("priced model rejected by a 1.0/MTok cap, want admitted (0.18 max rate)") - } - if filter.keep(unpriced) { - t.Error("unpriced model admitted by a price cap, want dropped") - } + require.True(t, active, "newModelFilter reported an inactive filter for a price cap") + assert.True(t, filter.keep(priced), "priced model rejected by a 1.0/MTok cap, want admitted (0.18 max rate)") + assert.False(t, filter.keep(unpriced), "unpriced model admitted by a price cap, want dropped") } func assertPricePtr(t *testing.T, name string, got *float64, want float64) { t.Helper() - if got == nil { - t.Errorf("%s = nil, want %v", name, want) - return - } - if *got != want { - t.Errorf("%s = %v, want %v", name, *got, want) - } + require.NotNil(t, got, "%s = nil, want %v", name, want) + assert.Equal(t, want, *got, name) } diff --git a/internal/server/passthrough_support_test.go b/internal/server/passthrough_support_test.go index 9cc4a0c22..c58458f3e 100644 --- a/internal/server/passthrough_support_test.go +++ b/internal/server/passthrough_support_test.go @@ -11,7 +11,6 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - "github.com/enterpilot/gomodel/config" "github.com/enterpilot/gomodel/internal/core" "github.com/enterpilot/gomodel/internal/echotest" "github.com/enterpilot/gomodel/internal/usage" @@ -50,41 +49,7 @@ func TestDefaultEnabledPassthroughProvidersIncludesMatrixProviders(t *testing.T) // and the default handler must not reject those requests before contacting the // upstream. func TestDefaultEnabledPassthroughProvidersIncludesEdenAI(t *testing.T) { - found := slices.Contains(defaultEnabledPassthroughProviders, "edenai") - if !found { - t.Fatalf("defaultEnabledPassthroughProviders = %v, want edenai included", defaultEnabledPassthroughProviders) - } -} - -// TestDefaultEnabledPassthroughProvidersMatchesConfigDefault keeps the two -// passthrough allowlist defaults from drifting apart. -// -// This package's slice is only the fallback for a Handler built without -// config; the list a running gateway actually enforces comes from -// config.Config.Server.EnabledPassthroughProviders, which http.go applies over -// the fallback. Adding a provider to one and not the other compiles, passes -// every handler test (they construct Handlers directly and so read the -// fallback), and still rejects the provider at runtime with "passthrough for -// X is not enabled" — which is exactly how the edenai entry was first missed. -func TestDefaultEnabledPassthroughProvidersMatchesConfigDefault(t *testing.T) { - // Load() with no config file present yields the built-in defaults; any - // ENABLED_PASSTHROUGH_PROVIDERS in the environment would mask them. - t.Setenv("ENABLED_PASSTHROUGH_PROVIDERS", "") - loaded, err := config.Load() - if err != nil { - t.Fatalf("config.Load() error = %v", err) - } - - fromConfig := append([]string(nil), loaded.Config.Server.EnabledPassthroughProviders...) - fromServer := append([]string(nil), defaultEnabledPassthroughProviders...) - slices.Sort(fromConfig) - slices.Sort(fromServer) - - if !slices.Equal(fromConfig, fromServer) { - t.Fatalf("passthrough allowlist defaults disagree:\n config/config.go: %v\n internal/server: %v\n"+ - "both must list the same provider types, or the runtime default silently differs from the tested one", - fromConfig, fromServer) - } + require.Contains(t, defaultEnabledPassthroughProviders, "edenai", "defaultEnabledPassthroughProviders = %v, want edenai included", defaultEnabledPassthroughProviders) } // A successful non-streaming JSON passthrough response must produce a usage diff --git a/internal/usage/cost_test.go b/internal/usage/cost_test.go index 33f346de2..3d4d698a1 100644 --- a/internal/usage/cost_test.go +++ b/internal/usage/cost_test.go @@ -790,23 +790,18 @@ func TestCalculateUsageCost_EdenAICostOverridesStaticPricing(t *testing.T) { result := CalculateUsageCost(1170, 99, map[string]any{"cost": 0.0002349}, "edenai", pricing) assertCostNear(t, "TotalCost", result.TotalCost, 0.0002349) - if result.Source != CostSourceEdenAICost { - t.Fatalf("Source = %q, want %q", result.Source, CostSourceEdenAICost) - } + require.Equal(t, CostSourceEdenAICost, result.Source) // Eden reports no input/output split, so only the total is claimed rather // than inventing a division of it. - if result.InputCost != nil || result.OutputCost != nil { - t.Fatalf("InputCost/OutputCost = %v/%v, want nil (Eden reports no split)", result.InputCost, result.OutputCost) - } + require.Nil(t, result.InputCost, "Eden reports no split") + require.Nil(t, result.OutputCost, "Eden reports no split") } func TestCalculateUsageCost_EdenAIAcceptsZeroCost(t *testing.T) { result := CalculateUsageCost(10, 4, map[string]any{"cost": 0.0}, "edenai", nil) assertCostNear(t, "TotalCost", result.TotalCost, 0) - if result.Source != CostSourceEdenAICost { - t.Fatalf("Source = %q, want %q", result.Source, CostSourceEdenAICost) - } + require.Equal(t, CostSourceEdenAICost, result.Source) } // Without a usable cost the request must fall back to the ordinary token math @@ -836,9 +831,7 @@ func TestCalculateUsageCost_EdenAIFallsBackToModelPricingWhenCostUnusable(t *tes assertCostNear(t, "InputCost", result.InputCost, 1.0) assertCostNear(t, "OutputCost", result.OutputCost, 1.0) assertCostNear(t, "TotalCost", result.TotalCost, 2.0) - if result.Source != CostSourceModelPricing { - t.Fatalf("Source = %q, want %q", result.Source, CostSourceModelPricing) - } + require.Equal(t, CostSourceModelPricing, result.Source) }) } } @@ -854,9 +847,7 @@ func TestCalculateUsageCost_EdenAICostIgnoredForOtherProviders(t *testing.T) { result := CalculateUsageCost(1_000_000, 500_000, map[string]any{"cost": 0.0002349}, "openai", pricing) assertCostNear(t, "TotalCost", result.TotalCost, 2.0) - if result.Source != CostSourceModelPricing { - t.Fatalf("Source = %q, want %q", result.Source, CostSourceModelPricing) - } + require.Equal(t, CostSourceModelPricing, result.Source) } // With no pricing and no usable cost the entry stays uncosted rather than @@ -864,10 +855,6 @@ func TestCalculateUsageCost_EdenAICostIgnoredForOtherProviders(t *testing.T) { func TestCalculateUsageCost_EdenAIWithoutCostOrPricingRecordsNothing(t *testing.T) { result := CalculateUsageCost(10, 4, map[string]any{}, "edenai", nil) - if result.TotalCost != nil { - t.Fatalf("TotalCost = %v, want nil", *result.TotalCost) - } - if result.Source != "" { - t.Fatalf("Source = %q, want empty", result.Source) - } + require.Nil(t, result.TotalCost) + require.Empty(t, result.Source) } diff --git a/internal/usage/extractor_test.go b/internal/usage/extractor_test.go index 389da1c10..2b09482d6 100644 --- a/internal/usage/extractor_test.go +++ b/internal/usage/extractor_test.go @@ -840,21 +840,13 @@ func TestExtractFromChatResponse_EdenAIExactCostReachesTotalCost(t *testing.T) { entry := ExtractFromChatResponse(resp, "req-eden", "edenai", "/v1/chat/completions", pricing) - if entry == nil { - t.Fatal("ExtractFromChatResponse() = nil") - } - if entry.RawData["cost"] != 0.0002349 { - t.Fatalf("RawData[cost] = %#v, want the lifted 0.0002349", entry.RawData["cost"]) - } - if entry.TotalCost == nil || math.Abs(*entry.TotalCost-0.0002349) > 1e-12 { - t.Fatalf("TotalCost = %v, want 0.0002349", entry.TotalCost) - } - if entry.CostSource != CostSourceEdenAICost { - t.Fatalf("CostSource = %q, want %q", entry.CostSource, CostSourceEdenAICost) - } - if entry.InputTokens != 1170 || entry.OutputTokens != 99 { - t.Errorf("token counts = %d/%d, want 1170/99", entry.InputTokens, entry.OutputTokens) - } + require.NotNil(t, entry) + require.Equal(t, 0.0002349, entry.RawData["cost"], "RawData[cost] should be the lifted value") + require.NotNil(t, entry.TotalCost) + require.InDelta(t, 0.0002349, *entry.TotalCost, 1e-12) + require.Equal(t, CostSourceEdenAICost, entry.CostSource) + assert.Equal(t, 1170, entry.InputTokens) + assert.Equal(t, 99, entry.OutputTokens) } // TestExtractFromChatResponse_EdenAIWithoutCostFallsBackToPricing asserts the @@ -870,16 +862,11 @@ func TestExtractFromChatResponse_EdenAIWithoutCostFallsBackToPricing(t *testing. entry := ExtractFromChatResponse(resp, "req-eden", "edenai", "/v1/chat/completions", pricing) - if entry == nil { - t.Fatal("ExtractFromChatResponse() = nil") - } - if entry.CostSource != CostSourceModelPricing { - t.Fatalf("CostSource = %q, want %q", entry.CostSource, CostSourceModelPricing) - } + require.NotNil(t, entry) + require.Equal(t, CostSourceModelPricing, entry.CostSource) // 1M * 0.06/1M + 0.5M * 0.18/1M = 0.06 + 0.09 - if entry.TotalCost == nil || math.Abs(*entry.TotalCost-0.15) > 1e-9 { - t.Fatalf("TotalCost = %v, want 0.15 from discovered per-model pricing", entry.TotalCost) - } + require.NotNil(t, entry.TotalCost) + require.InDelta(t, 0.15, *entry.TotalCost, 1e-9, "TotalCost should come from discovered per-model pricing") } // TestExtractFromEmbeddingResponse_ForwardsRawUsage asserts the provider's @@ -897,18 +884,11 @@ func TestExtractFromEmbeddingResponse_ForwardsRawUsage(t *testing.T) { } entry := ExtractFromEmbeddingResponse(resp, "req", "edenai", "/v1/embeddings") - if entry == nil { - t.Fatal("ExtractFromEmbeddingResponse returned nil") - } - if got := entry.RawData["cost"]; got != 0.0000012 { - t.Errorf("RawData[cost] = %v, want 0.0000012 forwarded from RawUsage", got) - } - if entry.TotalCost == nil || *entry.TotalCost != 0.0000012 { - t.Errorf("TotalCost = %v, want the provider-reported 0.0000012", entry.TotalCost) - } - if entry.CostSource != CostSourceEdenAICost { - t.Errorf("CostSource = %q, want %q", entry.CostSource, CostSourceEdenAICost) - } + require.NotNil(t, entry) + assert.Equal(t, 0.0000012, entry.RawData["cost"], "RawData[cost] should be forwarded from RawUsage") + require.NotNil(t, entry.TotalCost) + assert.Equal(t, 0.0000012, *entry.TotalCost, "TotalCost should be the provider-reported figure") + assert.Equal(t, CostSourceEdenAICost, entry.CostSource) } // TestExtractFromEmbeddingResponse_NilRawUsageStaysNil asserts a provider that @@ -923,15 +903,10 @@ func TestExtractFromEmbeddingResponse_NilRawUsageStaysNil(t *testing.T) { entry := ExtractFromEmbeddingResponse(resp, "req", "openai", "/v1/embeddings", &core.ModelPricing{Currency: "USD", InputPerMtok: &rate}) - if entry.RawData != nil { - t.Errorf("RawData = %v, want nil when the provider reported no extra usage", entry.RawData) - } - if entry.TotalCost == nil || *entry.TotalCost != 0.00002 { - t.Errorf("TotalCost = %v, want 0.00002 from the token rate", entry.TotalCost) - } - if entry.CostSource != CostSourceModelPricing { - t.Errorf("CostSource = %q, want %q", entry.CostSource, CostSourceModelPricing) - } + assert.Nil(t, entry.RawData, "RawData should stay nil when the provider reported no extra usage") + require.NotNil(t, entry.TotalCost) + assert.Equal(t, 0.00002, *entry.TotalCost, "TotalCost should come from the token rate") + assert.Equal(t, CostSourceModelPricing, entry.CostSource) } // TestExtractFromEmbeddingResponse_ExactCostSuppressesMissingUsageCaveat @@ -951,12 +926,9 @@ func TestExtractFromEmbeddingResponse_ExactCostSuppressesMissingUsageCaveat(t *t entry := ExtractFromEmbeddingResponse(resp, "req", "edenai", "/v1/embeddings", &core.ModelPricing{Currency: "USD", InputPerMtok: &rate}) - if entry.CostsCalculationCaveat != "" { - t.Errorf("CostsCalculationCaveat = %q, want empty: the cost was reported by the provider", entry.CostsCalculationCaveat) - } - if entry.TotalCost == nil || *entry.TotalCost != 0.0000012 { - t.Errorf("TotalCost = %v, want the provider-reported 0.0000012", entry.TotalCost) - } + assert.Empty(t, entry.CostsCalculationCaveat, "the cost was reported by the provider, so no caveat is expected") + require.NotNil(t, entry.TotalCost) + assert.Equal(t, 0.0000012, *entry.TotalCost, "TotalCost should be the provider-reported figure") } // TestExtractFromEmbeddingResponse_ZeroTokenCaveatStillApplies asserts the @@ -968,9 +940,7 @@ func TestExtractFromEmbeddingResponse_ZeroTokenCaveatStillApplies(t *testing.T) entry := ExtractFromEmbeddingResponse(resp, "req", "gemini", "/v1/embeddings", &core.ModelPricing{Currency: "USD", InputPerMtok: &rate}) - if entry.CostsCalculationCaveat == "" { - t.Error("CostsCalculationCaveat = empty, want the zero-token row flagged") - } + assert.NotEmpty(t, entry.CostsCalculationCaveat, "the zero-token row should be flagged") } // TestisProviderReportedCostSource pins which cost sources count as figures the @@ -990,8 +960,6 @@ func TestIsProviderReportedCostSource(t *testing.T) { } for _, tc := range tests { - if got := isProviderReportedCostSource(tc.source); got != tc.want { - t.Errorf("isProviderReportedCostSource(%q) = %v, want %v", tc.source, got, tc.want) - } + assert.Equal(t, tc.want, isProviderReportedCostSource(tc.source), "isProviderReportedCostSource(%q)", tc.source) } } diff --git a/internal/usage/stream_observer_test.go b/internal/usage/stream_observer_test.go index 6737e6828..02d5a93b5 100644 --- a/internal/usage/stream_observer_test.go +++ b/internal/usage/stream_observer_test.go @@ -2,6 +2,7 @@ package usage import ( "io" + "math" "strings" "sync" "testing" @@ -639,19 +640,13 @@ func TestStreamUsageObserverEdenAIRootLevelCost(t *testing.T) { observer.OnStreamClose() entries := logger.getEntries() - if len(entries) != 1 { - t.Fatalf("expected 1 entry, got %d", len(entries)) - } + require.Len(t, entries, 1) entry := entries[0] - if entry.RawData == nil || entry.RawData["cost"] != 0.0002349 { - t.Fatalf("RawData[cost] = %#v, want 0.0002349 harvested from the chunk root", entry.RawData["cost"]) - } - if entry.TotalCost == nil || *entry.TotalCost != 0.0002349 { - t.Fatalf("TotalCost = %v, want 0.0002349", entry.TotalCost) - } - if entry.CostSource != CostSourceEdenAICost { - t.Fatalf("CostSource = %q, want %q", entry.CostSource, CostSourceEdenAICost) - } + require.NotNil(t, entry.RawData) + require.Equal(t, 0.0002349, entry.RawData["cost"], "RawData[cost] should be harvested from the chunk root") + require.NotNil(t, entry.TotalCost) + require.Equal(t, 0.0002349, *entry.TotalCost) + require.Equal(t, CostSourceEdenAICost, entry.CostSource) } // A usage-level cost is the conventional location, so it stays authoritative @@ -673,12 +668,8 @@ func TestStreamUsageObserverUsageCostWinsOverRootLevelCost(t *testing.T) { observer.OnStreamClose() entries := logger.getEntries() - if len(entries) != 1 { - t.Fatalf("expected 1 entry, got %d", len(entries)) - } - if got := entries[0].RawData["cost"]; got != 0.5 { - t.Fatalf("RawData[cost] = %#v, want the usage-level 0.5", got) - } + require.Len(t, entries, 1) + require.Equal(t, 0.5, entries[0].RawData["cost"], "RawData[cost] should keep the usage-level value") } // An unusable root-level cost must not reach rawData, so the entry falls back @@ -707,25 +698,14 @@ func TestStreamUsageObserverRejectsUnusableRootLevelCost(t *testing.T) { }, }) observer.OnStreamClose() + entries := logger.getEntries() - - if len(entries) != 1 { - t.Fatalf("expected 1 entry, got %d", len(entries)) - } + require.Len(t, entries, 1) entry := entries[0] - if _, ok := entry.RawData["cost"]; ok { - t.Fatalf("RawData[cost] = %#v, want the unusable value dropped", entry.RawData["cost"]) - } - if entry.CostSource != CostSourceModelPricing { - t.Fatalf("CostSource = %q, want %q", entry.CostSource, CostSourceModelPricing) - } - if entry.TotalCost == nil || math.Abs(*entry.TotalCost-2.0) > 1e-9 { - t.Fatalf("TotalCost = %v, want the token-priced 2.0", entry.TotalCost) - } - require.Len(t, entries, 1) - require.NotNil(t, entries[0].InputCost) - assert.InDelta(t, tt.wantInput, *entries[0].InputCost, 1e-9) - assert.Equal(t, "jev-1.13.0", entries[0].Model) + require.NotContains(t, entry.RawData, "cost", "unusable root-level cost should be dropped") + require.Equal(t, CostSourceModelPricing, entry.CostSource) + require.NotNil(t, entry.TotalCost) + require.InDelta(t, 2.0, *entry.TotalCost, 1e-9, "TotalCost should be token-priced") }) } } @@ -753,17 +733,13 @@ func TestStreamUsageObserverRootLevelCostDoesNotRepriceOtherProviders(t *testing observer.OnStreamClose() entries := logger.getEntries() - if len(entries) != 1 { - t.Fatalf("expected 1 entry, got %d", len(entries)) - } + require.Len(t, entries, 1) entry := entries[0] - if entry.CostSource != CostSourceModelPricing { - t.Fatalf("CostSource = %q, want %q", entry.CostSource, CostSourceModelPricing) - } - if entry.TotalCost == nil || math.Abs(*entry.TotalCost-2.0) > 1e-9 { - t.Fatalf("TotalCost = %v, want the token-priced 2.0", entry.TotalCost) - } + require.Equal(t, CostSourceModelPricing, entry.CostSource) + require.NotNil(t, entry.TotalCost) + require.InDelta(t, 2.0, *entry.TotalCost, 1e-9, "TotalCost should be token-priced") } + // A routed alias (jev-latest) is answered by a versioned model (jev-1.13.0); // the routed model's price wins, and the answered model's price applies when // only it is declared. @@ -787,7 +763,10 @@ func TestStreamUsageObserverPricesAnsweredModelWhenRoutedHasNone(t *testing.T) { observer.OnJSONEvent(map[string]any{ "model": "jev-1.13.0", "usage": map[string]any{"input_tokens": float64(1_000_000), "output_tokens": float64(0)}, - + }) + observer.OnStreamClose() + + entries := logger.getEntries() require.Len(t, entries, 1) require.NotNil(t, entries[0].InputCost) assert.InDelta(t, tt.wantInput, *entries[0].InputCost, 1e-9) diff --git a/run/lifecycle_test.go b/run/lifecycle_test.go index b3c354dc1..326a33972 100644 --- a/run/lifecycle_test.go +++ b/run/lifecycle_test.go @@ -232,18 +232,11 @@ func TestMain_EdenAIProviderRegistration(t *testing.T) { factory := defaultProviderFactory(&config.Config{}) registered := factory.RegisteredTypes() - found := slices.Contains(registered, "edenai") - if !found { - t.Fatalf("edenai not in RegisteredTypes() = %v", registered) - } + require.Contains(t, registered, "edenai", "edenai not in RegisteredTypes() = %v", registered) provider, err := factory.Create(providers.ProviderConfig{Type: "edenai", APIKey: "test"}) - if err != nil { - t.Fatalf("factory.Create(edenai) error = %v, want nil", err) - } - if provider == nil { - t.Fatal("factory.Create(edenai) returned nil provider") - } + require.NoError(t, err, "factory.Create(edenai)") + require.NotNil(t, provider, "factory.Create(edenai) returned nil provider") } func TestMain_HetznerProviderRegistration(t *testing.T) {