diff --git a/docs/superpowers/plans/2026-05-02-settings-driven-llm-providers-models.md b/docs/superpowers/plans/2026-05-02-settings-driven-llm-providers-models.md index 3d766d6dc..fa1841850 100644 --- a/docs/superpowers/plans/2026-05-02-settings-driven-llm-providers-models.md +++ b/docs/superpowers/plans/2026-05-02-settings-driven-llm-providers-models.md @@ -23,6 +23,13 @@ The foundation slice of this plan is implemented and shipped on the run branch: The remainder of the plan — replacing `fabro_model::Provider` with `ProviderId` across 80+ files, regenerating the OpenAPI clients, swapping the auth resolver to use `ProviderId`, replacing the 25 `Catalog::builtin()` production call sites with a settings-resolved `Arc` injected through server/workflow/CLI state, the `bootstrap_catalog` install hatch, the typed `Request.speed`/`GenerateParams.speed` swap, and the per-speed billing rows — is **deferred to follow-up sessions**. Each deferred step is marked individually below. +## Phase 1 gateway header update (2026-05-12 session) + +Phase 1 gateway header work is motivated by @haroldolivieri's Portkey/Bedrock report on PR #207: +https://github.com/fabro-sh/fabro/pull/207#issuecomment-4377929769 + +This follow-up adds provider-level `extra_headers` with typed `literal`, `env`, and `credential` values, whole-map replacement semantics across settings layers, and adapter-registry pass-through coverage. It remains schema/seam work only; runtime credential resolution and settings-defined provider registration stay deferred to the resolved catalog/client phases. + --- ## Summary @@ -115,6 +122,9 @@ speed = "fast" - [x] Preserve sparse field-merge semantics for `[llm.providers.]` and `[llm.models.]`. Arrays such as `credentials`, `aliases`, `controls.reasoning_effort`, and `controls.speed` replace as whole arrays. (Backed by `MergeMap` per-key field-merge; arrays are `Option>` with `or` combine semantics.) - [x] Keep the targeted legacy `[llm]` migration error for old keys such as `provider` or `model`; accept only the new `[llm.providers]` and `[llm.models]` subtrees. (`LEGACY_LLM_KEYS` matched in `parse_settings` before the strict deserialize.) - [x] Parse adapter keys as strings in `fabro-config`. Do not make `fabro-config` depend on `fabro-llm`. (Adapter is `Option`; resolution happens against `fabro_model::adapter` metadata.) + - [x] Add provider-level `extra_headers` with typed literal/env/credential values. + - [x] Make `extra_headers` replace as a whole map across settings layers. + - [x] Keep gateway headers as schema/seam work only; runtime credential resolution and provider registration remain deferred to the resolved catalog/client phases. - [~] **Catalog model** — partially landed. Remaining items are **deferred** because they require breaking changes across 80+ files and the OpenAPI regeneration step. - [x] Add `ProviderId` and `ModelId` string newtypes where they improve type clarity across crates. (`fabro_model::ids`.) @@ -234,4 +244,4 @@ speed = "fast" - Field-merge for provider/model tables is intentional. Whole-array replacement for controls can mask future built-in values; more granular array merge operations are deferred. - V1 does not support custom auth schemes, data-driven profile templates, provider-level CLI backend routing, data-driven adapter implementations, or new request control kinds. - Adding a new value to an existing Rust-owned control enum, such as a new speed value beyond `standard` and `fast`, remains a Rust change. -- Existing imprecise knowledge cutoff labels migrate to exact normalized dates, e.g. `May 2025` becomes `2025-05-01`; presentation can render lower precision. \ No newline at end of file +- Existing imprecise knowledge cutoff labels migrate to exact normalized dates, e.g. `May 2025` becomes `2025-05-01`; presentation can render lower precision. diff --git a/lib/crates/fabro-config/src/layers/combine.rs b/lib/crates/fabro-config/src/layers/combine.rs index f3dda42c7..94eb345e9 100644 --- a/lib/crates/fabro-config/src/layers/combine.rs +++ b/lib/crates/fabro-config/src/layers/combine.rs @@ -14,7 +14,7 @@ use fabro_types::settings::{Duration, InterpString, Size}; use super::LogFilter; use super::cli::{CliAuthLayer, CliLoggingLayer, CliTargetLayer}; use super::features::FeaturesLayer; -use super::llm::{CostRates, CredentialRef}; +use super::llm::{CostRates, CredentialRef, HeaderValueRef}; use super::run::{ DaytonaSnapshotLayer, HookAgentMarker, HookEntry, HookTlsMode, InterviewProviderLayer, ModelRefOrSplice, NotificationProviderLayer, RunArtifactsLayer, RunCheckpointLayer, @@ -117,6 +117,12 @@ impl Combine for Option> { } } +impl Combine for Option> { + fn combine(self, other: Self) -> Self { + self.or(other) + } +} + macro_rules! impl_combine_self { ($($ty:ty),+ $(,)?) => { $( diff --git a/lib/crates/fabro-config/src/layers/llm.rs b/lib/crates/fabro-config/src/layers/llm.rs index 109b98143..5a347732f 100644 --- a/lib/crates/fabro-config/src/layers/llm.rs +++ b/lib/crates/fabro-config/src/layers/llm.rs @@ -26,10 +26,10 @@ //! Resolution against the static adapter registry happens in `fabro-model` //! when the resolved [`Catalog`](fabro_model::Catalog) is built. -use std::collections::BTreeMap; +use std::collections::{BTreeMap, HashMap}; use chrono::NaiveDate; -use serde::{Deserialize, Deserializer, Serialize}; +use serde::{Deserialize, Deserializer, Serialize, Serializer}; use super::maps::MergeMap; @@ -53,24 +53,29 @@ pub struct LlmLayer { #[serde(deny_unknown_fields)] pub struct ProviderSettings { #[serde(default, skip_serializing_if = "Option::is_none")] - pub display_name: Option, + pub display_name: Option, /// Adapter registry key (e.g. `"openai_compatible"`). #[serde(default, skip_serializing_if = "Option::is_none")] - pub adapter: Option, + pub adapter: Option, #[serde(default, skip_serializing_if = "Option::is_none")] - pub base_url: Option, + pub base_url: Option, /// Ordered list of credential references — first successful wins. Each /// entry must be a typed `CredentialRef` (`credential:` or /// `env:`); literal secret strings fail deserialization. #[serde(default, skip_serializing_if = "Option::is_none")] - pub credentials: Option>, + pub credentials: Option>, + /// Extra HTTP headers attached to every outgoing provider request after + /// credential resolution. Header values are typed so secret-bearing values + /// stay as references until a later resolution phase. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub extra_headers: Option>, /// Higher wins; missing → `0`; ties broken by canonical provider ID. #[serde(default, skip_serializing_if = "Option::is_none")] - pub priority: Option, + pub priority: Option, #[serde(default, skip_serializing_if = "Option::is_none")] - pub enabled: Option, + pub enabled: Option, #[serde(default, skip_serializing_if = "Option::is_none")] - pub aliases: Option>, + pub aliases: Option>, } /// One entry in `[llm.models.]`. @@ -314,6 +319,143 @@ impl TryFrom for CredentialRef { } } +// --------------------------------------------------------------------------- +// HeaderValueRef - typed extra header value +// --------------------------------------------------------------------------- + +/// A typed provider extra-header value. +/// +/// Literal values are intended for non-secret routing metadata. Secret-bearing +/// values must use `env` or `credential` references so settings never need to +/// carry raw API keys as successful values. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum HeaderValueRef { + Literal(String), + Env(String), + Credential(String), +} + +impl Serialize for HeaderValueRef { + fn serialize(&self, serializer: S) -> Result + where + S: Serializer, + { + use serde::ser::SerializeMap; + + let mut map = serializer.serialize_map(Some(1))?; + match self { + Self::Literal(value) => map.serialize_entry("literal", value)?, + Self::Env(value) => map.serialize_entry("env", value)?, + Self::Credential(value) => map.serialize_entry("credential", value)?, + } + map.end() + } +} + +#[derive(Debug, Deserialize)] +#[serde(untagged)] +enum HeaderValueRefInput { + Table(HeaderValueRefSerde), + BareString(String), +} + +#[derive(Debug, Clone, PartialEq, Eq, Deserialize)] +#[serde(rename_all = "snake_case", deny_unknown_fields)] +struct HeaderValueRefSerde { + #[serde(default)] + literal: Option, + #[serde(default)] + env: Option, + #[serde(default)] + credential: Option, +} + +impl<'de> Deserialize<'de> for HeaderValueRef { + fn deserialize(deserializer: D) -> Result + where + D: Deserializer<'de>, + { + use serde::de::Error as _; + + match HeaderValueRefInput::deserialize(deserializer)? { + HeaderValueRefInput::Table(value) => value.try_into().map_err(D::Error::custom), + HeaderValueRefInput::BareString(value) => { + drop(value); + Err(D::Error::custom(HeaderValueRefParseError::WrongFieldCount)) + } + } + } +} + +impl TryFrom for HeaderValueRef { + type Error = HeaderValueRefParseError; + + fn try_from(value: HeaderValueRefSerde) -> Result { + let populated = [ + value.literal.as_ref(), + value.env.as_ref(), + value.credential.as_ref(), + ] + .into_iter() + .flatten() + .count(); + + if populated != 1 { + return Err(HeaderValueRefParseError::WrongFieldCount); + } + + if let Some(value) = value.literal { + if value.is_empty() { + return Err(HeaderValueRefParseError::EmptyValue); + } + return Ok(Self::Literal(value)); + } + if let Some(value) = value.env { + if value.is_empty() { + return Err(HeaderValueRefParseError::EmptyValue); + } + return Ok(Self::Env(value)); + } + if let Some(value) = value.credential { + if value.is_empty() { + return Err(HeaderValueRefParseError::EmptyValue); + } + return Ok(Self::Credential(value)); + } + + unreachable!("populated field count was already checked"); + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum HeaderValueRefParseError { + WrongFieldCount, + EmptyValue, +} + +impl std::fmt::Display for HeaderValueRefParseError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::WrongFieldCount => f.write_str( + "header value must be a table with exactly one of `literal`, `env`, or `credential`; bare strings are rejected", + ), + Self::EmptyValue => f.write_str("header value reference must not be empty"), + } + } +} + +impl std::error::Error for HeaderValueRefParseError {} + +impl std::fmt::Display for HeaderValueRef { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Literal(_) => f.write_str("literal:"), + Self::Env(name) => write!(f, "env:{name}"), + Self::Credential(id) => write!(f, "credential:{id}"), + } + } +} + #[cfg(test)] mod tests { use std::str::FromStr; @@ -406,6 +548,142 @@ mod tests { assert!(err.is_err(), "literal secret strings must fail to parse"); } + // ---- HeaderValueRef -------------------------------------------------- + + #[test] + fn header_value_ref_parses_literal_form() { + let parsed: HeaderValueRef = toml::from_str(r#"value = { literal = "@bedrock-prod" }"#) + .map(|v: toml::Value| { + v.as_table() + .unwrap() + .get("value") + .unwrap() + .clone() + .try_into() + .unwrap() + }) + .unwrap(); + + assert_eq!(parsed, HeaderValueRef::Literal("@bedrock-prod".to_string())); + assert_eq!(parsed.to_string(), "literal:"); + } + + #[test] + fn header_value_ref_parses_env_form() { + let parsed: HeaderValueRef = toml::from_str(r#"value = { env = "PORTKEY_API_KEY" }"#) + .map(|v: toml::Value| { + v.as_table() + .unwrap() + .get("value") + .unwrap() + .clone() + .try_into() + .unwrap() + }) + .unwrap(); + + assert_eq!(parsed, HeaderValueRef::Env("PORTKEY_API_KEY".to_string())); + assert_eq!(parsed.to_string(), "env:PORTKEY_API_KEY"); + } + + #[test] + fn header_value_ref_parses_credential_form() { + let parsed: HeaderValueRef = toml::from_str(r#"value = { credential = "portkey_config" }"#) + .map(|v: toml::Value| { + v.as_table() + .unwrap() + .get("value") + .unwrap() + .clone() + .try_into() + .unwrap() + }) + .unwrap(); + + assert_eq!( + parsed, + HeaderValueRef::Credential("portkey_config".to_string()) + ); + assert_eq!(parsed.to_string(), "credential:portkey_config"); + } + + #[test] + fn header_value_ref_rejects_bare_string() { + #[derive(Debug, Deserialize)] + struct Wrap { + #[expect( + dead_code, + reason = "field exists only to drive the deserializer; we assert on the parse error" + )] + value: HeaderValueRef, + } + + let err = toml::from_str::(r#"value = "sk-portkey-literal""#).unwrap_err(); + let message = err.message(); + + assert!(message.contains("header value")); + assert!( + !message.contains("sk-portkey-literal"), + "deserializer message must not echo a possible literal secret", + ); + } + + #[test] + fn header_value_ref_rejects_ambiguous_table() { + #[derive(Debug, Deserialize)] + struct Wrap { + #[expect( + dead_code, + reason = "field exists only to drive the deserializer; we assert on the parse error" + )] + value: HeaderValueRef, + } + + let err = toml::from_str::( + r#"value = { env = "PORTKEY_API_KEY", literal = "@bedrock-prod" }"#, + ) + .unwrap_err(); + + assert!(err.to_string().contains("exactly one")); + } + + #[test] + fn header_value_ref_rejects_unknown_keys() { + #[derive(Deserialize)] + #[expect( + dead_code, + reason = "field exists only to drive the deserializer; we assert on the parse error" + )] + struct Wrap { + value: HeaderValueRef, + } + + let err: Result = toml::from_str(r#"value = { secret = "PORTKEY_API_KEY" }"#); + + assert!(err.is_err(), "unknown header value keys must fail"); + } + + #[test] + fn header_value_ref_rejects_empty_values() { + #[derive(Debug, Deserialize)] + struct Wrap { + #[expect( + dead_code, + reason = "field exists only to drive the deserializer; we assert on the parse error" + )] + value: HeaderValueRef, + } + + for source in [ + r#"value = { literal = "" }"#, + r#"value = { env = "" }"#, + r#"value = { credential = "" }"#, + ] { + let err = toml::from_str::(source).unwrap_err(); + assert!(err.to_string().contains("must not be empty")); + } + } + // ---- LlmLayer parsing ------------------------------------------------- #[test] @@ -434,6 +712,56 @@ aliases = ["moonshot"] ]); } + #[test] + fn parses_provider_extra_headers() { + let toml = r#" +[providers.portkey] +display_name = "Portkey Bedrock" +adapter = "anthropic" +base_url = "https://api.portkey.ai/v1" + +[providers.portkey.extra_headers] +x-portkey-api-key = { env = "PORTKEY_API_KEY" } +x-portkey-provider = { literal = "@bedrock-prod" } +x-portkey-config = { credential = "portkey_config" } +"#; + + let layer: LlmLayer = toml::from_str(toml).unwrap(); + let portkey = layer.providers.get("portkey").unwrap(); + + assert!(portkey.credentials.is_none()); + let headers = portkey.extra_headers.as_ref().unwrap(); + assert_eq!( + headers.get("x-portkey-api-key"), + Some(&HeaderValueRef::Env("PORTKEY_API_KEY".to_string())), + ); + assert_eq!( + headers.get("x-portkey-provider"), + Some(&HeaderValueRef::Literal("@bedrock-prod".to_string())), + ); + assert_eq!( + headers.get("x-portkey-config"), + Some(&HeaderValueRef::Credential("portkey_config".to_string())), + ); + } + + #[test] + fn provider_extra_headers_reject_bare_string_values() { + let toml = r#" +[providers.portkey.extra_headers] +x-portkey-api-key = "sk-portkey-literal" +"#; + + let err = toml::from_str::(toml).unwrap_err(); + let message = err.message(); + + assert!(message.contains("header value")); + assert!( + !message.contains("sk-portkey-literal"), + "deserializer message must not echo a possible literal secret", + ); + } + #[test] fn parses_full_model_entry() { let toml = r#" @@ -603,6 +931,78 @@ mystery = 1 )]); } + #[test] + fn provider_extra_headers_map_replaces_wholesale() { + let high = ProviderSettings { + extra_headers: Some(HashMap::from([( + "x-portkey-provider".to_string(), + HeaderValueRef::Literal("@bedrock-prod".to_string()), + )])), + ..ProviderSettings::default() + }; + let low = ProviderSettings { + extra_headers: Some(HashMap::from([ + ( + "x-portkey-api-key".to_string(), + HeaderValueRef::Env("PORTKEY_API_KEY".to_string()), + ), + ( + "x-portkey-provider".to_string(), + HeaderValueRef::Literal("@bedrock-default".to_string()), + ), + ])), + ..ProviderSettings::default() + }; + + let merged = high.combine(low); + + let headers = merged.extra_headers.unwrap(); + assert_eq!(headers.len(), 1); + assert_eq!( + headers.get("x-portkey-provider"), + Some(&HeaderValueRef::Literal("@bedrock-prod".to_string())), + ); + assert!(!headers.contains_key("x-portkey-api-key")); + } + + #[test] + fn provider_extra_headers_inherit_when_unset() { + let high = ProviderSettings::default(); + let low = ProviderSettings { + extra_headers: Some(HashMap::from([( + "x-portkey-api-key".to_string(), + HeaderValueRef::Env("PORTKEY_API_KEY".to_string()), + )])), + ..ProviderSettings::default() + }; + + let merged = high.combine(low); + + assert_eq!( + merged.extra_headers.unwrap().get("x-portkey-api-key"), + Some(&HeaderValueRef::Env("PORTKEY_API_KEY".to_string())), + ); + } + + #[test] + fn provider_extra_headers_empty_map_clears_lower_layer() { + let high = ProviderSettings { + extra_headers: Some(HashMap::new()), + ..ProviderSettings::default() + }; + let low = ProviderSettings { + extra_headers: Some(HashMap::from([( + "x-portkey-api-key".to_string(), + HeaderValueRef::Env("PORTKEY_API_KEY".to_string()), + )])), + ..ProviderSettings::default() + }; + + let merged = high.combine(low); + + assert!(merged.extra_headers.unwrap().is_empty()); + } + #[test] fn merge_map_field_merges_per_provider_id() { let mut high_map: std::collections::HashMap = diff --git a/lib/crates/fabro-config/src/layers/mod.rs b/lib/crates/fabro-config/src/layers/mod.rs index ab739b75a..97783b531 100644 --- a/lib/crates/fabro-config/src/layers/mod.rs +++ b/lib/crates/fabro-config/src/layers/mod.rs @@ -18,9 +18,9 @@ pub use cli::{ pub(crate) use combine::Combine; pub use features::FeaturesLayer; pub use llm::{ - CostRates, CredentialRef, CredentialRefParseError, LlmLayer, ModelControls, ModelCostTable, - ModelFeatures as LlmModelFeatures, ModelLimits as LlmModelLimits, ModelSettings, - ProviderSettings, + CostRates, CredentialRef, CredentialRefParseError, HeaderValueRef, LlmLayer, ModelControls, + ModelCostTable, ModelFeatures as LlmModelFeatures, ModelLimits as LlmModelLimits, + ModelSettings, ProviderSettings, }; pub use log_filter::LogFilter; pub use maps::{MergeMap, ReplaceMap, StickyMap}; diff --git a/lib/crates/fabro-config/src/lib.rs b/lib/crates/fabro-config/src/lib.rs index c3e67da69..4bd767eff 100644 --- a/lib/crates/fabro-config/src/lib.rs +++ b/lib/crates/fabro-config/src/lib.rs @@ -41,20 +41,20 @@ pub use layers::{ CliAuthLayer, CliExecAgentLayer, CliExecLayer, CliExecModelLayer, CliLayer, CliLoggingLayer, CliOutputLayer, CliTargetLayer, CliUpdatesLayer, CostRates, CredentialRef, CredentialRefParseError, DaytonaDockerfileLayer, DaytonaSandboxLayer, DaytonaSnapshotLayer, - DockerSandboxLayer, FeaturesLayer, GitAuthorLayer, GithubIntegrationLayer, HookAgentMarker, - HookEntry, HookTlsMode, IntegrationWebhooksLayer, InterviewProviderLayer, InterviewsLayer, - LlmLayer, LlmModelFeatures, LlmModelLimits, LogFilter, McpEntryLayer, MergeMap, ModelControls, - ModelCostTable, ModelRefOrSplice, ModelSettings, NotificationProviderLayer, - NotificationRouteLayer, ObjectStoreLocalLayer, ObjectStoreS3Layer, PrepareStep, ProjectLayer, - ProviderSettings, ReplaceMap, RunAgentLayer, RunArtifactsLayer, RunCheckpointLayer, - RunCloneLayer, RunExecutionLayer, RunGitLayer, RunGoalLayer, RunIntegrationsGithubLayer, - RunIntegrationsLayer, RunLayer, RunMetaBranchLayer, RunModelControlsLayer, RunModelLayer, - RunPrepareLayer, RunPullRequestLayer, RunRunBranchLayer, RunSandboxLayer, RunScmLayer, - ScmGitHubLayer, ServerApiLayer, ServerArtifactsLayer, ServerAuthGithubLayer, ServerAuthLayer, - ServerIntegrationsLayer, ServerIpAllowlistLayer, ServerIpAllowlistOverrideLayer, ServerLayer, - ServerListenLayer, ServerLoggingLayer, ServerSchedulerLayer, ServerSlateDbLayer, - ServerStorageLayer, ServerWebLayer, SlackIntegrationLayer, StickyMap, StringOrSplice, - WorkflowLayer, + DockerSandboxLayer, FeaturesLayer, GitAuthorLayer, GithubIntegrationLayer, HeaderValueRef, + HookAgentMarker, HookEntry, HookTlsMode, IntegrationWebhooksLayer, InterviewProviderLayer, + InterviewsLayer, LlmLayer, LlmModelFeatures, LlmModelLimits, LogFilter, McpEntryLayer, + MergeMap, ModelControls, ModelCostTable, ModelRefOrSplice, ModelSettings, + NotificationProviderLayer, NotificationRouteLayer, ObjectStoreLocalLayer, ObjectStoreS3Layer, + PrepareStep, ProjectLayer, ProviderSettings, ReplaceMap, RunAgentLayer, RunArtifactsLayer, + RunCheckpointLayer, RunCloneLayer, RunExecutionLayer, RunGitLayer, RunGoalLayer, + RunIntegrationsGithubLayer, RunIntegrationsLayer, RunLayer, RunMetaBranchLayer, + RunModelControlsLayer, RunModelLayer, RunPrepareLayer, RunPullRequestLayer, RunRunBranchLayer, + RunSandboxLayer, RunScmLayer, ScmGitHubLayer, ServerApiLayer, ServerArtifactsLayer, + ServerAuthGithubLayer, ServerAuthLayer, ServerIntegrationsLayer, ServerIpAllowlistLayer, + ServerIpAllowlistOverrideLayer, ServerLayer, ServerListenLayer, ServerLoggingLayer, + ServerSchedulerLayer, ServerSlateDbLayer, ServerStorageLayer, ServerWebLayer, + SlackIntegrationLayer, StickyMap, StringOrSplice, WorkflowLayer, }; pub(crate) use layers::{Combine, SettingsLayer}; pub use logging::{resolve_log_destination, resolve_log_destination_with_env}; diff --git a/lib/crates/fabro-llm/src/adapter_registry.rs b/lib/crates/fabro-llm/src/adapter_registry.rs index 2707022c0..fe2629864 100644 --- a/lib/crates/fabro-llm/src/adapter_registry.rs +++ b/lib/crates/fabro-llm/src/adapter_registry.rs @@ -69,7 +69,7 @@ impl AdapterConfig { /// rather than re-shaping every existing factory. pub type AdapterFactory = fn(AdapterConfig) -> Arc; -fn build_anthropic(config: AdapterConfig) -> Arc { +fn build_anthropic_adapter(config: AdapterConfig) -> providers::AnthropicAdapter { let mut adapter = providers::AnthropicAdapter::new(auth_value(&config.auth_header)); if let Some(base_url) = config.base_url { adapter = adapter.with_base_url(base_url); @@ -77,10 +77,14 @@ fn build_anthropic(config: AdapterConfig) -> Arc { if !config.extra_headers.is_empty() { adapter = adapter.with_default_headers(config.extra_headers); } - Arc::new(adapter) + adapter } -fn build_openai(config: AdapterConfig) -> Arc { +fn build_anthropic(config: AdapterConfig) -> Arc { + Arc::new(build_anthropic_adapter(config)) +} + +fn build_openai_adapter(config: AdapterConfig) -> providers::OpenAiAdapter { let mut adapter = providers::OpenAiAdapter::new(auth_value(&config.auth_header)); if let Some(base_url) = config.base_url { adapter = adapter.with_base_url(base_url); @@ -97,10 +101,14 @@ fn build_openai(config: AdapterConfig) -> Arc { if let Some(project_id) = config.project_id { adapter = adapter.with_project_id(project_id); } - Arc::new(adapter) + adapter } -fn build_gemini(config: AdapterConfig) -> Arc { +fn build_openai(config: AdapterConfig) -> Arc { + Arc::new(build_openai_adapter(config)) +} + +fn build_gemini_adapter(config: AdapterConfig) -> providers::GeminiAdapter { let mut adapter = providers::GeminiAdapter::new(auth_value(&config.auth_header)); if let Some(base_url) = config.base_url { adapter = adapter.with_base_url(base_url); @@ -108,10 +116,14 @@ fn build_gemini(config: AdapterConfig) -> Arc { if !config.extra_headers.is_empty() { adapter = adapter.with_default_headers(config.extra_headers); } - Arc::new(adapter) + adapter } -fn build_openai_compatible(config: AdapterConfig) -> Arc { +fn build_gemini(config: AdapterConfig) -> Arc { + Arc::new(build_gemini_adapter(config)) +} + +fn build_openai_compatible_adapter(config: AdapterConfig) -> providers::OpenAiCompatibleAdapter { // `openai_compatible` providers vary widely in base URL; the catalog must // pre-resolve `[llm.providers.].base_url` before constructing // `AdapterConfig`. There is no sensible default — silently routing to one @@ -126,7 +138,11 @@ fn build_openai_compatible(config: AdapterConfig) -> Arc { if !config.extra_headers.is_empty() { adapter = adapter.with_default_headers(config.extra_headers); } - Arc::new(adapter) + adapter +} + +fn build_openai_compatible(config: AdapterConfig) -> Arc { + Arc::new(build_openai_compatible_adapter(config)) } /// Single source of truth pairing every adapter key with its factory. Both @@ -223,6 +239,67 @@ mod tests { assert_eq!(adapter.name(), "kimi"); } + #[test] + fn openai_compatible_factory_preserves_extra_headers() { + let config = AdapterConfig { + provider_id: "portkey".to_string(), + auth_header: ApiKeyHeader::Bearer("unused-primary-key".to_string()), + base_url: Some("https://api.portkey.ai/v1".to_string()), + extra_headers: HashMap::from([ + ( + "x-portkey-api-key".to_string(), + "resolved-portkey-key".to_string(), + ), + ( + "x-portkey-provider".to_string(), + "@bedrock-prod".to_string(), + ), + ]), + codex_mode: false, + org_id: None, + project_id: None, + }; + + let adapter = build_openai_compatible_adapter(config); + + assert_eq!(adapter.name(), "portkey"); + assert_eq!( + adapter.http.default_headers.get("x-portkey-api-key"), + Some(&"resolved-portkey-key".to_string()), + ); + assert_eq!( + adapter.http.default_headers.get("x-portkey-provider"), + Some(&"@bedrock-prod".to_string()), + ); + } + + #[test] + fn anthropic_factory_preserves_extra_headers() { + let config = AdapterConfig { + provider_id: "anthropic-through-portkey".to_string(), + auth_header: ApiKeyHeader::Custom { + name: "x-api-key".to_string(), + value: "unused-primary-key".to_string(), + }, + base_url: Some("https://api.portkey.ai/v1".to_string()), + extra_headers: HashMap::from([( + "x-portkey-api-key".to_string(), + "resolved-portkey-key".to_string(), + )]), + codex_mode: false, + org_id: None, + project_id: None, + }; + + let adapter = build_anthropic_adapter(config); + + assert_eq!(adapter.name(), "anthropic"); + assert_eq!( + adapter.http.default_headers.get("x-portkey-api-key"), + Some(&"resolved-portkey-key".to_string()), + ); + } + #[test] #[should_panic(expected = "openai_compatible adapter requires a base_url")] fn openai_compatible_factory_panics_without_base_url() {