diff --git a/.env.template b/.env.template index 71f7f446b..375f6373b 100644 --- a/.env.template +++ b/.env.template @@ -106,11 +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) +# 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 +# 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 @@ -726,6 +726,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 1e9cafc0a..4af385622 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/cmd/gomodel/docs/docs.go b/cmd/gomodel/docs/docs.go index 90f72c5dc..22669ad2f 100644 --- a/cmd/gomodel/docs/docs.go +++ b/cmd/gomodel/docs/docs.go @@ -10251,6 +10251,10 @@ const docTemplate = `{ "prompt_tokens": { "type": "integer" }, + "raw_usage": { + "type": "object", + "additionalProperties": {} + }, "total_tokens": { "type": "integer" } diff --git a/config/config.example.yaml b/config/config.example.yaml index 1a177aaa6..526edfbdc 100644 --- a/config/config.example.yaml +++ b/config/config.example.yaml @@ -572,6 +572,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 146105ad7..52e0b8fb2 100644 --- a/config/config.go +++ b/config/config.go @@ -128,6 +128,7 @@ func buildDefaultConfig() *Config { "llamacpp", "llmd", "deepseek", + "edenai", "jev", }, }, diff --git a/config/config_test.go b/config/config_test.go index 45982cf67..7ee99005d 100644 --- a/config/config_test.go +++ b/config/config_test.go @@ -132,7 +132,7 @@ func TestBuildDefaultConfig(t *testing.T) { assert.Equal(t, DefaultStreamStallTimeoutSeconds, cfg.Server.StreamStallTimeout) assert.True(t, cfg.Server.EnablePassthroughRoutes) assert.True(t, cfg.Server.AllowPassthroughV1Alias) - assert.Equal(t, []string{"openai", "anthropic", "openrouter", "kilo", "zai", "sglang", "vllm", "llamacpp", "llmd", "deepseek", "jev"}, cfg.Server.EnabledPassthroughProviders) + assert.Equal(t, []string{"openai", "anthropic", "openrouter", "kilo", "zai", "sglang", "vllm", "llamacpp", "llmd", "deepseek", "edenai", "jev"}, cfg.Server.EnabledPassthroughProviders) assert.Equal(t, ConfiguredProviderModelsModeFallback, cfg.Models.ConfiguredProviderModelsMode) assert.Nil(t, cfg.Cache.Model.Local) assert.Equal(t, 3600, cfg.Cache.Model.RefreshInterval) diff --git a/docs/advanced/configuration.mdx b/docs/advanced/configuration.mdx index 8d7e6e7d6..9547ae94e 100644 --- a/docs/advanced/configuration.mdx +++ b/docs/advanced/configuration.mdx @@ -434,6 +434,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 | | `ZAI_API_KEY` | Z.ai | | `XAI_API_KEY` | xAI (Grok) | | `GROQ_API_KEY` | Groq | @@ -446,7 +447,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`, @@ -594,6 +595,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 511eb42e5..b515fbfe1 100644 --- a/docs/docs.json +++ b/docs/docs.json @@ -226,6 +226,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 bc409fdaa..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`, 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. @@ -160,7 +160,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,jev +ENABLED_PASSTHROUGH_PROVIDERS=openai,anthropic,openrouter,kilo,zai,sglang,vllm,llamacpp,llmd,deepseek,edenai,jev ``` Set `ENABLED_PASSTHROUGH_PROVIDERS` to the provider types you want to expose. diff --git a/docs/openapi.json b/docs/openapi.json index a6cb55a95..2995d02dc 100644 --- a/docs/openapi.json +++ b/docs/openapi.json @@ -14141,6 +14141,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 new file mode 100644 index 000000000..d490bb9c4 --- /dev/null +++ b/docs/providers/edenai.mdx @@ -0,0 +1,219 @@ +--- +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 +``` + + + 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`: + +```yaml +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 +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, 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: 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` 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. + +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 — 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 +`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" } +``` + +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 +`/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 d8ea5d3bb..d39621942 100644 --- a/docs/providers/overview.mdx +++ b/docs/providers/overview.mdx @@ -56,6 +56,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) | @@ -134,6 +135,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/core/types.go b/internal/core/types.go index a13f42671..6f4545ef2 100644 --- a/internal/core/types.go +++ b/internal/core/types.go @@ -530,7 +530,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/config_test.go b/internal/providers/config_test.go index a983bea2d..bd6ff7e3c 100644 --- a/internal/providers/config_test.go +++ b/internal/providers/config_test.go @@ -90,6 +90,9 @@ var testDiscoveryConfigs = map[string]DiscoveryConfig{ "hetzner": { DefaultBaseURL: "https://inference.hetzner.com/api/v1", }, + "edenai": { + DefaultBaseURL: "https://api.edenai.run/v3", + }, } // --- buildProviderConfig --- @@ -1467,6 +1470,37 @@ func TestBuildProviderConfig_Hetzner_ResolvesBaseURL(t *testing.T) { assert.Equal(t, testDiscoveryConfigs["hetzner"].DefaultBaseURL, p.BaseURL) } +// 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"] + 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 +// 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"] + 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) { t.Setenv("OPENROUTER_API_KEY", "sk-openrouter") t.Setenv("OPENROUTER_MODEL_FILTER_INCLUDE", "*:free, *:nitro") diff --git a/internal/providers/edenai/capabilities_test.go b/internal/providers/edenai/capabilities_test.go new file mode 100644 index 000000000..f7b3fb708 --- /dev/null +++ b/internal/providers/edenai/capabilities_test.go @@ -0,0 +1,185 @@ +package edenai + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +// 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"} { + 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. + 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"} { + 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"} { + assert.NotContains(t, capabilities, name, "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", + } { + assert.True(t, model.Metadata.Capabilities[name], "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 + 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 +// 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"} { + 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 +// 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} + }]}`) + + 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 +// 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"} { + assert.True(t, capabilities[name], "capability %q = false, want true", name) + } + for _, name := range []string{"file", "text"} { + assert.NotContains(t, capabilities, name, "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"] + } + }]}`) + + 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 new file mode 100644 index 000000000..7d4e1b309 --- /dev/null +++ b/internal/providers/edenai/edenai.go @@ -0,0 +1,192 @@ +// 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. +// +// 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), + opts.HTTPClient, + ))} +} + +// compatibleConfig returns the shared OpenAI-compatible transport settings for +// 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, + BaseURL: baseURL, + HTTPClient: guardedHTTPClient(httpClient), + 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. +// +// 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 !secureDestination(req.URL) { + return + } + 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, providerType) +} + +// 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. +// +// 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) { + 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 + } + core.EnsureModel(&resp.Model, req.Model) + normalizeEmbeddingResponse(&resp) + return &resp.EmbeddingResponse, 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..4f7a081ec --- /dev/null +++ b/internal/providers/edenai/edenai_test.go @@ -0,0 +1,360 @@ +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" + "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 +// 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{}) + + require.NotNil(t, provider, "provider should not be nil") + + concrete, ok := provider.(*Provider) + 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) + 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 +// 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) + require.True(t, ok, "New() did not return *edenai.Provider") + assert.Equal(t, custom, provider.GetBaseURL()) +} + +// 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 +// 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) { + 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 +// 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 := 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"}}, + }) + 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 +// 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"}}` + require.NoError(t, json.Unmarshal([]byte(raw), &req)) + + 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) + 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) + 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 +// /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 := 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"}}, + }) + require.NoError(t, err) + defer stream.Close() + body, err := io.ReadAll(stream) + 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 +// 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 := 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", + }) + 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. + 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 +// 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 := newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{ + Model: slashedModel, + Input: "hi", + }) + require.NoError(t, err) + require.Equal(t, []string{"/chat/completions"}, paths, "upstream paths = %v, want exactly [/chat/completions]", paths) + for _, path := range paths { + 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 +// 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 := newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + stream, err := provider.StreamResponses(context.Background(), &core.ResponsesRequest{ + Model: slashedModel, + Input: "hi", + }) + require.NoError(t, err) + defer stream.Close() + _, 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 +// 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 := 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"}`)), + }) + require.NoError(t, err) + defer resp.Body.Close() + require.Equal(t, "/chat/completions", gotPath) + require.Equal(t, "Bearer edenai-key", gotAuth) + require.Equal(t, http.StatusOK, 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 := newTestProvider("edenai-key", "", nil, llmclient.Hooks{}) + + 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 new file mode 100644 index 000000000..d6cb9faef --- /dev/null +++ b/internal/providers/edenai/embeddings_cost_test.go @@ -0,0 +1,217 @@ +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" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// 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 newTestProvider("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()) + require.NoError(t, err) + + cost, ok := resp.Usage.RawUsage["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. + 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. + 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 +// 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()) + 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) + 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 +// 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()) + 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) + 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 +// 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()) + 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}) + assert.Equal(t, usage.CostSourceModelPricing, entry.CostSource, "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()) + require.NoError(t, err) + + rate := 0.02 + entry := usage.ExtractFromEmbeddingResponse(resp, "req", providerType, "/v1/embeddings", + &core.ModelPricing{Currency: "USD", InputPerMtok: &rate}) + 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 +// 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()) + 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 +// provide is still in place now that the provider issues the request itself. +func TestEmbeddings_NilRequestRejected(t *testing.T) { + provider := embeddingsProvider(t, edenEmbeddingBody) + + _, err := provider.Embeddings(context.Background(), nil) + require.Error(t, err, "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()) + 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.go b/internal/providers/edenai/models.go new file mode 100644 index 000000000..e06c2e314 --- /dev/null +++ b/internal/providers/edenai/models.go @@ -0,0 +1,314 @@ +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. 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"` + CacheCreationInputTokenCost *float64 `json:"cache_creation_input_token_cost"` +} + +// 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 { + 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 field(block) +} + +// 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 +} + +// 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 + } + 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..8de01d0be --- /dev/null +++ b/internal/providers/edenai/models_test.go @@ -0,0 +1,539 @@ +package edenai + +import ( + "context" + "math" + "net/http" + "net/http/httptest" + "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. +// 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 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()) + require.NoError(t, err) + require.Len(t, resp.Data, 1, "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()) + 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] + 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 + 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). + 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} + assert.Len(t, meta.Capabilities, len(wantCapabilities), "Capabilities = %v, want %v", meta.Capabilities, wantCapabilities) + for name := range wantCapabilities { + 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) + assert.Equal(t, "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) + 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 +// 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} + }]}`) + + 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) +} + +// TestListModels_ListPricingDoesNotMaskApplicablePricing asserts the fallback +// 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) + 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 +// 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) + 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 +// 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) + 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 +// 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 { + require.Nil(t, model.Metadata.Pricing, "Pricing = %+v, want nil", model.Metadata.Pricing) + } + return + } + 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) + }) + } +} + +// 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} { + _, 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)) + 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 +// 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"}, + }, + { + // 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"]}`, + 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 + } + assert.Equal(t, tt.wantModes, modes) + for _, capability := range tt.wantCapabilities { + require.NotNil(t, model.Metadata, "Capabilities missing %q", capability) + assert.True(t, model.Metadata.Capabilities[capability], "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()) + 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 +// 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 := newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + resp, err := provider.ListModels(context.Background()) + 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() + 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 { + 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.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..40ced510b --- /dev/null +++ b/internal/providers/edenai/passthrough_semantics_test.go @@ -0,0 +1,67 @@ +package edenai + +import ( + "testing" + + "github.com/enterpilot/gomodel/internal/core" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestPassthroughSemanticEnricher(t *testing.T) { + require.Equal(t, "edenai", passthroughSemanticEnricher.ProviderType()) + + 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, + }) + 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) + }) + } +} + +// 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", + }) + 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.go b/internal/providers/edenai/response.go new file mode 100644 index 000000000..a378413ed --- /dev/null +++ b/internal/providers/edenai/response.go @@ -0,0 +1,162 @@ +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 +} + +// 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 new file mode 100644 index 000000000..976cdc25c --- /dev/null +++ b/internal/providers/edenai/response_test.go @@ -0,0 +1,216 @@ +package edenai + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + + "github.com/goccy/go-json" + + "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. +// 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 := 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"}}, + }) + require.NoError(t, 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"] + require.True(t, ok, "Usage.RawUsage = %v, want Eden's root-level cost lifted in", resp.Usage.RawUsage) + value, ok := cost.(float64) + 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 +// 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) + require.NoError(t, err) + var decoded map[string]any + 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 +// 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) + require.NotContains(t, resp.Usage.RawUsage, "cost", "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) + 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 +// 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) + + 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 +// 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) + require.NotEmpty(t, raw, "ExtraFields missing %q", upstreamProviderField) + var upstream string + 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 +// 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}}`) + + 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 +// 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 := 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", + }) + 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 +// 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 := newTestProvider("edenai-key", server.URL, server.Client(), llmclient.Hooks{}) + resp, err := provider.Responses(context.Background(), &core.ResponsesRequest{ + Model: slashedModel, + Input: "hi", + }) + require.NoError(t, err) + value, ok := resp.Usage.RawUsage["cost"].(float64) + 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 +// 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 := 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"}}, + }) + + 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 new file mode 100644 index 000000000..bb52e7ca0 --- /dev/null +++ b/internal/providers/edenai/transport.go @@ -0,0 +1,192 @@ +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 + +// 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 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 + } + 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, 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. +// +// 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). 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 +// 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.Transport = &secureTransport{base: guarded.Transport} + guarded.CheckRedirect = redirectPolicy(base.CheckRedirect) + return &guarded +} + +// redirectPolicy composes Eden's redirect rules with whatever policy the +// caller's client already carried. +// +// 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. +// +// 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 { + return err + } + if caller == nil { + return nil + } + 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) + } +} + +// 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 !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 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 new file mode 100644 index 000000000..fec701088 --- /dev/null +++ b/internal/providers/edenai/transport_test.go @@ -0,0 +1,845 @@ +package edenai + +import ( + "context" + "errors" + "io" + "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" + "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 +// 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) + require.NoError(t, err, "url.Parse(%q)", tc.raw) + assert.Equal(t, tc.want, secureDestination(parsed), "secureDestination(%q)", tc.raw) + }) + } +} + +// 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) { + 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) + require.NoError(t, err, "http.NewRequest") + setHeaders(req, "eden-key") + + assert.Equal(t, "Bearer eden-key", req.Header.Get("Authorization")) +} + +// 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) + require.NoError(t, err, "http.NewRequest") + setHeaders(req, "eden-key") + + assert.Empty(t, req.Header.Get("Authorization"), "the credential must not be sent in cleartext") +} + +// 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) + require.NoError(t, err, "http.NewRequest") + setHeaders(req, "eden-key") + + assert.Equal(t, "Bearer eden-key", req.Header.Get("Authorization")) +} + +// TestCleartextBaseURL_RequestRefusedBeforeSending is the payload half of the +// cleartext guarantee, and the reason withholding the credential is not enough +// on its own. +// +// 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_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(`{"id":"x","choices":[]}`)) + })) + defer server.Close() + + target, err := url.Parse(server.URL) + require.NoError(t, err, "url.Parse") + client := &http.Client{Transport: &cleartextRouteTransport{cleartext: target.Host}} + + 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"}}, + }) + require.Error(t, err, "ChatCompletion succeeded against a cleartext endpoint, want the request refused") + + 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: +// 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 := 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 +// 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) + require.NoError(t, err, "http.NewRequest") + resp, err := transport.RoundTrip(req) + if resp != nil && resp.Body != nil { + _ = resp.Body.Close() + } + + 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: 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) + 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. + 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) + require.NoError(t, err, "http.NewRequest") + resp, err := transport.RoundTrip(req) + require.NoError(t, err, "RoundTrip through the nil base") + defer resp.Body.Close() + assert.Equal(t, http.StatusNoContent, resp.StatusCode) +} + +// 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 +// 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_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 + }{ + {"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) + 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") + }) + } +} + +// 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) + + 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 +// 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} + + 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 +// 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") + + 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") + require.NoError(t, checkRedirect(upper, []*http.Request{origin}), "checkRedirect differing-case host") +} + +// 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://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 := mustRequest(t, target) + err := checkRedirect(req, []*http.Request{origin}) + assert.ErrorContains(t, err, "cross-host", "checkRedirect(%q), want a cross-host refusal", target) + } +} + +// 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") + + 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 +// 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. + 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 +// 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 + }) + + err := policy(target, []*http.Request{origin}) + assert.ErrorIs(t, err, http.ErrUseLastResponse, "want http.ErrUseLastResponse propagated") +} + +// 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}) + 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) + }) + } +} + +// 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 + }) + + require.NoError(t, policy(target, []*http.Request{origin}), "a same-host path rewrite is fine") +} + +// 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 + }) + err := policy(target, []*http.Request{origin}) + assert.ErrorIs(t, err, tc.err, "want %v returned unchanged", 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 + }) + + err := policy(target, []*http.Request{origin}) + assert.ErrorIs(t, err, http.ErrUseLastResponse, "want http.ErrUseLastResponse propagated unchanged") +} + +// 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 + } { + 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 — +// 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") + + 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 +// 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") + + 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) + require.NoError(t, err, "http.NewRequest(%q)", rawURL) + return req +} + +// 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) + require.NoError(t, err, "http.NewRequest") + + err = checkRedirect(req, nil) + 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) + require.NoError(t, err, "http.NewRequest") + + err = checkRedirect(req, nil) + require.Error(t, err, "checkRedirect = nil, want a refusal") + for _, secret := range []string{"s3cret", "leaked"} { + assert.NotContains(t, err.Error(), secret, "error leaked %q from the redirect URL", 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) + require.NoError(t, err, "http.NewRequest") + via := make([]*http.Request, 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 +// 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) + require.NoError(t, err, "url.Parse") + client := secure.Client() + client.Transport = &cleartextRouteTransport{ + base: client.Transport, + cleartext: cleartextTarget.Host, + } + + provider := newTestProvider("eden-key", secure.URL+"/v3", client, llmclient.Hooks{}) + _, err = provider.Embeddings(context.Background(), embeddingRequest()) + require.Error(t, err, "Embeddings succeeded through an HTTPS -> HTTP redirect, want the redirect refused") + + 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 +// 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 := 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 +// 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 := newTestProvider("eden-key", secure.URL+"/v3", secure.Client(), llmclient.Hooks{}) + resp, err := provider.Embeddings(context.Background(), embeddingRequest()) + 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 +// 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) + 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 := newTestProvider("eden-key", "https://eden.test/v3", client, llmclient.Hooks{}) + _, err = provider.Embeddings(context.Background(), embeddingRequest()) + require.Error(t, err, "Embeddings followed an HTTPS subdomain redirect, want it refused") + + 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 +// 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. +// TestNew_DefaultsBaseURL already covers New resolving to this value. +func TestDefaultBaseURLIsHTTPS(t *testing.T) { + require.True(t, strings.HasPrefix(defaultBaseURL, "https://"), "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) + require.True(t, ok, "New did not return *Provider") + + const override = "https://eden.eu.example.com/v3" + provider.SetBaseURL(override) + assert.Equal(t, override, provider.GetBaseURL(), "GetBaseURL() after SetBaseURL") +} + +// 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 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 nil path +// deliberately gets a different client. +func TestGuardedHTTPClient_PreservesDefaultClientSemantics(t *testing.T) { + guarded := guardedHTTPClient(http.DefaultClient) + + 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) + 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. + 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 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) + + 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") +} + +// 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() + + 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 new file mode 100644 index 000000000..6a91e7151 --- /dev/null +++ b/internal/providers/registry_provider_pricing_test.go @@ -0,0 +1,124 @@ +package providers + +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" +) + +// 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") + + 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. + 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") + 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 + 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) + 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 +// 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") + + require.NoError(t, registry.Initialize(context.Background()), "Initialize") + + for _, selector := range []string{"edenai/openai/gpt-4", "openai/gpt-4"} { + pricing := registry.ResolvePricing(selector, "edenai") + 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) + } +} + +// 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)}) + 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() + require.NotNil(t, got, "%s = nil, want %v", name, want) + assert.Equal(t, want, *got, name) +} diff --git a/internal/server/handlers_test.go b/internal/server/handlers_test.go index b8cd9e565..2bb998e1e 100644 --- a/internal/server/handlers_test.go +++ b/internal/server/handlers_test.go @@ -5939,7 +5939,7 @@ func TestProviderPassthrough_RejectsUnsupportedProvider(t *testing.T) { require.Equal(t, http.StatusBadRequest, rec.Code) require.Contains(t, rec.Body.String(), `provider passthrough for \"groq\" is not enabled`) - require.Contains(t, rec.Body.String(), "anthropic, deepseek, hetzner, jev, kilo, llamacpp, llmd, openai, openrouter, sglang, vllm, zai") + require.Contains(t, rec.Body.String(), "anthropic, deepseek, edenai, hetzner, jev, kilo, llamacpp, llmd, openai, openrouter, sglang, vllm, zai") } func TestProviderPassthrough_ChutesRequiresExplicitOptIn(t *testing.T) { diff --git a/internal/server/passthrough_support.go b/internal/server/passthrough_support.go index 51ba65e79..ea99d0a31 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", "jev"} +var defaultEnabledPassthroughProviders = []string{"openai", "anthropic", "openrouter", "kilo", "zai", "sglang", "vllm", "llamacpp", "llmd", "deepseek", "hetzner", "edenai", "jev"} 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 09fde6a7a..c58458f3e 100644 --- a/internal/server/passthrough_support_test.go +++ b/internal/server/passthrough_support_test.go @@ -44,6 +44,14 @@ func TestDefaultEnabledPassthroughProvidersIncludesMatrixProviders(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) { + require.Contains(t, defaultEnabledPassthroughProviders, "edenai", "defaultEnabledPassthroughProviders = %v, want edenai included", defaultEnabledPassthroughProviders) +} + // 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 0d5bda78e..71d3e69ea 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 @@ -567,6 +568,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) } @@ -616,6 +620,48 @@ 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") +} + +// 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/cost_test.go b/internal/usage/cost_test.go index 6f9fdf3b1..3d4d698a1 100644 --- a/internal/usage/cost_test.go +++ b/internal/usage/cost_test.go @@ -774,3 +774,87 @@ func assertCostNear(t *testing.T, name string, got *float64, want float64) { require.NotNil(t, got, "%s is nil, want %f", name, want) require.InDelta(t, want, *got, 1e-9, "%s", name) } + +// --- 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) + 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. + 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) + require.Equal(t, CostSourceEdenAICost, result.Source) +} + +// 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) + require.Equal(t, CostSourceModelPricing, result.Source) + }) + } +} + +// 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) + require.Equal(t, CostSourceModelPricing, result.Source) +} + +// 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) + + require.Nil(t, result.TotalCost) + require.Empty(t, result.Source) +} diff --git a/internal/usage/extractor.go b/internal/usage/extractor.go index 732a45b5d..1898a179c 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, entry.Timestamp, pricing...)) && entry.CostsCalculationCaveat == "" { entry.CostsCalculationCaveat = caveatEmbeddingMissingUsage diff --git a/internal/usage/extractor_test.go b/internal/usage/extractor_test.go index c781d0185..2b09482d6 100644 --- a/internal/usage/extractor_test.go +++ b/internal/usage/extractor_test.go @@ -817,3 +817,149 @@ func TestExtractFromEmbeddingResponse_NoUsageCaveat(t *testing.T) { retained := retainedMissingUsageCaveat(caveatEmbeddingMissingUsage, 0, nil, tiered) require.Equal(t, caveatEmbeddingMissingUsage, 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) + + 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 +// 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) + + require.NotNil(t, entry) + require.Equal(t, CostSourceModelPricing, entry.CostSource) + // 1M * 0.06/1M + 0.5M * 0.18/1M = 0.06 + 0.09 + 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 +// 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") + 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 +// 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}) + 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 +// 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}) + 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 +// 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}) + assert.NotEmpty(t, entry.CostsCalculationCaveat, "the zero-token row should be 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 { + assert.Equal(t, tc.want, isProviderReportedCostSource(tc.source), "isProviderReportedCostSource(%q)", tc.source) + } +} diff --git a/internal/usage/stream_observer.go b/internal/usage/stream_observer.go index ae8e99c65..eb74d700e 100644 --- a/internal/usage/stream_observer.go +++ b/internal/usage/stream_observer.go @@ -252,6 +252,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_") @@ -320,6 +321,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 79c27c896..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" @@ -618,6 +619,127 @@ func TestStreamUsageObserverAnthropicNativeEvents(t *testing.T) { assert.Equal(t, 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() + require.Len(t, entries, 1) + entry := entries[0] + 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 +// 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() + 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 +// 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() + require.Len(t, entries, 1) + entry := entries[0] + 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") + }) + } +} + +// 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() + require.Len(t, entries, 1) + entry := entries[0] + 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. diff --git a/run/lifecycle_test.go b/run/lifecycle_test.go index b3eaeab7e..326a33972 100644 --- a/run/lifecycle_test.go +++ b/run/lifecycle_test.go @@ -228,6 +228,17 @@ func TestMain_KimicodeProviderRegistration(t *testing.T) { require.NotNil(t, provider) } +func TestMain_EdenAIProviderRegistration(t *testing.T) { + factory := defaultProviderFactory(&config.Config{}) + + registered := factory.RegisteredTypes() + require.Contains(t, registered, "edenai", "edenai not in RegisteredTypes() = %v", registered) + + provider, err := factory.Create(providers.ProviderConfig{Type: "edenai", APIKey: "test"}) + require.NoError(t, err, "factory.Create(edenai)") + require.NotNil(t, provider, "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 a1bffd594..cf384d4ac 100644 --- a/run/providers.go +++ b/run/providers.go @@ -14,6 +14,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" @@ -61,6 +62,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 aca6429fe..ae3995c9c 100644 --- a/run/providers_test.go +++ b/run/providers_test.go @@ -168,7 +168,7 @@ var credentialPayloadFields = []string{ func TestDefaultProviderFactoryRegistersAllProviderTypes(t *testing.T) { expected := []string{ - "anthropic", "audiocpp", "azure", "bailian", "bedrock", "bedrock-mantle", "chatgpt", "chutes", "cohere", "deepseek", "elevenlabs", + "anthropic", "audiocpp", "azure", "bailian", "bedrock", "bedrock-mantle", "chatgpt", "chutes", "cohere", "deepseek", "edenai", "elevenlabs", "fireworks", "gemini", "groq", "hetzner", "jev", "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",