diff --git a/lib/foundation/fabro-auth/Cargo.toml b/lib/foundation/fabro-auth/Cargo.toml index 8fcfa13b0..ed3f5f2b0 100644 --- a/lib/foundation/fabro-auth/Cargo.toml +++ b/lib/foundation/fabro-auth/Cargo.toml @@ -18,12 +18,12 @@ async-trait.workspace = true base64.workspace = true chrono = { workspace = true, features = ["serde"] } fabro-http.workspace = true -fabro-model = { path = "../fabro-model" } fabro-oauth = { path = "../fabro-oauth" } fabro-redact.workspace = true fabro-static.workspace = true fabro-types = { path = "../fabro-types" } fabro-vault = { path = "../fabro-vault" } +lithos-llm = { workspace = true, features = ["runtime"] } serde.workspace = true serde_json.workspace = true thiserror.workspace = true @@ -32,6 +32,7 @@ tracing.workspace = true [dev-dependencies] httpmock = "0.8" +lithos-llm = { workspace = true, features = ["runtime", "builtin-catalog"] } tempfile = "3" tokio = { workspace = true, features = ["macros", "test-util"] } toml.workspace = true diff --git a/lib/foundation/fabro-auth/src/api_key_source.rs b/lib/foundation/fabro-auth/src/api_key_source.rs new file mode 100644 index 000000000..8cebcd4e3 --- /dev/null +++ b/lib/foundation/fabro-auth/src/api_key_source.rs @@ -0,0 +1,72 @@ +//! A credential source holding one operator-supplied API key. +//! +//! Used to validate a pasted key before it is stored: the key is shaped into +//! the provider's declared auth scheme and offered for that provider only. + +use std::collections::HashMap; +use std::sync::Arc; + +use async_trait::async_trait; +use fabro_vault::Vault; +use lithos_llm::catalog::{Catalog, CatalogProvider, ProviderId}; +use lithos_llm::credentials::Credentials; +use tokio::sync::RwLock as AsyncRwLock; + +use crate::credential_source::CredentialSource; +use crate::resolve::{ResolveError, credentials_for_api_key}; + +pub struct ApiKeyCredentialSource { + provider: ProviderId, + key: String, + vault: Arc>, +} + +impl ApiKeyCredentialSource { + /// A source for `provider` with no vault behind it, so extra headers that + /// interpolate vault secrets fail to resolve. + #[must_use] + pub fn new(provider: ProviderId, key: String) -> Self { + Self::with_vault( + provider, + key, + Arc::new(AsyncRwLock::new(Vault::from_entries(HashMap::new()))), + ) + } + + /// A source for `provider` whose extra headers resolve against `vault`. + #[must_use] + pub fn with_vault(provider: ProviderId, key: String, vault: Arc>) -> Self { + Self { + provider, + key, + vault, + } + } +} + +impl std::fmt::Debug for ApiKeyCredentialSource { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ApiKeyCredentialSource") + .field("provider", &self.provider) + .finish_non_exhaustive() + } +} + +#[async_trait] +impl CredentialSource for ApiKeyCredentialSource { + async fn credentials(&self, provider: &CatalogProvider) -> Result { + if provider.id() != &self.provider { + return Err(ResolveError::NotConfigured(provider.id().clone())); + } + let vault = self.vault.read().await; + credentials_for_api_key(provider, self.key.clone(), &vault) + } + + async fn configured_providers(&self, catalog: &Catalog) -> Vec { + catalog + .provider(self.provider.as_str()) + .ok() + .map(|provider| vec![provider.id().clone()]) + .unwrap_or_default() + } +} diff --git a/lib/foundation/fabro-auth/src/context.rs b/lib/foundation/fabro-auth/src/context.rs index 70c3b94eb..f8401b99e 100644 --- a/lib/foundation/fabro-auth/src/context.rs +++ b/lib/foundation/fabro-auth/src/context.rs @@ -1,4 +1,4 @@ -use fabro_model::ProviderId; +use lithos_llm::catalog::ProviderId; #[derive(Debug, Clone, PartialEq, Eq)] pub enum AuthContextRequest { diff --git a/lib/foundation/fabro-auth/src/credential.rs b/lib/foundation/fabro-auth/src/credential.rs index 95b195d61..0b185521f 100644 --- a/lib/foundation/fabro-auth/src/credential.rs +++ b/lib/foundation/fabro-auth/src/credential.rs @@ -1,5 +1,4 @@ use chrono::{DateTime, Duration, Utc}; -use fabro_redact::redact_string; pub use fabro_types::{OAuthConfig, OAuthCredential, OAuthTokens}; pub(crate) fn expires_at_from_now(expires_in: Option) -> DateTime { @@ -7,44 +6,6 @@ pub(crate) fn expires_at_from_now(expires_in: Option) -> DateTime { Utc::now() + Duration::seconds(seconds) } -#[derive(Clone, PartialEq, Eq)] -pub enum ApiKeyHeader { - Bearer(String), - Custom { - name: String, - value: String, - }, - /// No static header: the request is authenticated by AWS SigV4 signing, - /// with credentials resolved from the AWS default chain at request time. - AwsSigv4, -} - -fn redact_for_debug(value: &str) -> String { - let redacted = redact_string(value); - if redacted == value && !value.is_empty() { - "REDACTED".to_string() - } else { - redacted - } -} - -impl std::fmt::Debug for ApiKeyHeader { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - match self { - Self::Bearer(value) => f - .debug_tuple("Bearer") - .field(&redact_for_debug(value)) - .finish(), - Self::Custom { name, value } => f - .debug_struct("Custom") - .field("name", name) - .field("value", &redact_for_debug(value)) - .finish(), - Self::AwsSigv4 => f.write_str("AwsSigv4"), - } - } -} - #[cfg(test)] mod tests { use super::*; @@ -81,12 +42,4 @@ mod tests { assert!(fixture(Utc::now() + Duration::minutes(4)).needs_refresh()); assert!(!fixture(Utc::now() + Duration::minutes(6)).needs_refresh()); } - - #[test] - fn api_key_header_debug_redacts_secret_values() { - let header = ApiKeyHeader::Bearer("sk-test".to_string()); - let debug = format!("{header:?}"); - assert!(!debug.contains("sk-test")); - assert!(debug.contains("REDACTED")); - } } diff --git a/lib/foundation/fabro-auth/src/credential_ref.rs b/lib/foundation/fabro-auth/src/credential_ref.rs new file mode 100644 index 000000000..821637f99 --- /dev/null +++ b/lib/foundation/fabro-auth/src/credential_ref.rs @@ -0,0 +1,124 @@ +//! Credential references declared in `metadata.fabro.credentials`. + +use std::str::FromStr; + +use serde::{Deserialize, Deserializer, Serialize, de}; + +/// Where one provider secret comes from. +/// +/// A provider lists these in order; the first that resolves wins. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum CredentialRef { + /// A token or OAuth entry in the Fabro vault. + Vault(String), + /// A process environment variable. + Env(String), + /// The AWS default credential chain. Resolves without a secret; the + /// Bedrock adapter signs each request. + AwsSigv4, +} + +impl std::fmt::Display for CredentialRef { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Vault(name) => write!(f, "vault:{name}"), + Self::Env(name) => write!(f, "env:{name}"), + Self::AwsSigv4 => f.write_str("aws_sigv4"), + } + } +} + +impl FromStr for CredentialRef { + type Err = CredentialRefParseError; + + fn from_str(value: &str) -> Result { + if value == "aws_sigv4" { + return Ok(Self::AwsSigv4); + } + if let Some(name) = value.strip_prefix("vault:") { + return if name.is_empty() { + Err(CredentialRefParseError::EmptyVault) + } else { + Ok(Self::Vault(name.to_string())) + }; + } + if let Some(name) = value.strip_prefix("env:") { + return if name.is_empty() { + Err(CredentialRefParseError::EmptyEnv) + } else { + Ok(Self::Env(name.to_string())) + }; + } + Err(CredentialRefParseError::Invalid) + } +} + +impl Serialize for CredentialRef { + fn serialize(&self, serializer: S) -> Result { + serializer.serialize_str(&self.to_string()) + } +} + +impl<'de> Deserialize<'de> for CredentialRef { + fn deserialize>(deserializer: D) -> Result { + let value = String::deserialize(deserializer)?; + value.parse().map_err(de::Error::custom) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)] +pub enum CredentialRefParseError { + #[error("credential reference must be `vault:`, `env:`, or `aws_sigv4`")] + Invalid, + #[error("credential reference is missing a name after `vault:`")] + EmptyVault, + #[error("credential reference is missing a name after `env:`")] + EmptyEnv, +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parses_each_form() { + assert_eq!( + "vault:OPENAI_CODEX".parse::().unwrap(), + CredentialRef::Vault("OPENAI_CODEX".into()) + ); + assert_eq!( + "env:KIMI_API_KEY".parse::().unwrap(), + CredentialRef::Env("KIMI_API_KEY".into()) + ); + assert_eq!( + "aws_sigv4".parse::().unwrap(), + CredentialRef::AwsSigv4 + ); + } + + #[test] + fn rejects_literal_secrets_without_echoing_them() { + let err = "sk-ant-1234".parse::().unwrap_err(); + assert_eq!(err, CredentialRefParseError::Invalid); + assert!(!err.to_string().contains("sk-ant")); + assert_eq!( + "vault:".parse::().unwrap_err(), + CredentialRefParseError::EmptyVault + ); + assert_eq!( + "env:".parse::().unwrap_err(), + CredentialRefParseError::EmptyEnv + ); + } + + #[test] + fn round_trips_through_serde_strings() { + let value: Vec = + serde_json::from_str(r#"["env:A", "vault:b", "aws_sigv4"]"#).unwrap(); + assert_eq!( + serde_json::to_string(&value).unwrap(), + r#"["env:A","vault:b","aws_sigv4"]"# + ); + assert!(serde_json::from_str::>(r#"["sk-literal"]"#).is_err()); + } +} diff --git a/lib/foundation/fabro-auth/src/credential_source.rs b/lib/foundation/fabro-auth/src/credential_source.rs index f611e629b..6529bcef6 100644 --- a/lib/foundation/fabro-auth/src/credential_source.rs +++ b/lib/foundation/fabro-auth/src/credential_source.rs @@ -1,17 +1,102 @@ +//! Per-attempt credential lookup for LLM providers. +//! +//! [`CredentialSource`] is Fabro's storage-aware credential seam: the vault, +//! the SQL secret store, and the process environment each implement it. +//! [`lithos_credentials`] adapts a source into the lithos +//! [`CredentialProvider`] the client calls before every provider attempt, so a +//! refreshed OAuth token is picked up by the next retry. + +use std::sync::Arc; + use async_trait::async_trait; -use fabro_model::{Catalog, ProviderId}; +use fabro_types::catalog_policy; +use lithos_llm::catalog::{Catalog, CatalogProvider, ProviderId}; +use lithos_llm::credentials::{CredentialError, CredentialProvider, Credentials}; -use crate::{ApiCredential, ResolveError}; +use crate::{ResolveError, auth_issue_message}; -#[derive(Debug)] +/// Which providers a source can serve right now, and why the rest cannot. +#[derive(Debug, Default)] pub struct ResolvedCredentials { - pub credentials: Vec, + /// Enabled providers whose credentials resolved. + pub ready: Vec, + /// Enabled providers with credential material that failed to resolve, + /// such as an expired OAuth token that could not be refreshed. Providers + /// with no material at all are not issues; they are simply absent. pub auth_issues: Vec<(ProviderId, ResolveError)>, } +impl ResolvedCredentials { + /// A human-readable line per auth issue. + #[must_use] + pub fn issue_messages(&self) -> Vec { + self.auth_issues + .iter() + .map(|(provider, issue)| auth_issue_message(provider, issue)) + .collect() + } +} + #[async_trait] pub trait CredentialSource: Send + Sync { - async fn resolve(&self, catalog: &Catalog) -> anyhow::Result; + /// Resolves `provider`'s credentials for one request attempt. + async fn credentials(&self, provider: &CatalogProvider) -> Result; + /// Providers with credential material present. Does not refresh or + /// validate anything, so it is cheap enough for listings. async fn configured_providers(&self, catalog: &Catalog) -> Vec; + + /// Resolves every enabled provider once, separating the ready set from + /// the providers that have material but cannot use it. + async fn resolve_all(&self, catalog: &Catalog) -> ResolvedCredentials { + let mut resolved = ResolvedCredentials::default(); + for provider in catalog.providers() { + if !catalog_policy::provider_policy(provider).is_enabled() { + continue; + } + match self.credentials(provider).await { + Ok(_) => resolved.ready.push(provider.id().clone()), + Err(ResolveError::NotConfigured(_)) => {} + Err(err) => resolved.auth_issues.push((provider.id().clone(), err)), + } + } + resolved + } +} + +/// Adapts a [`CredentialSource`] into the lithos credential provider. +#[must_use] +pub fn lithos_credentials(source: Arc) -> Arc { + Arc::new(SourceCredentialProvider { source }) +} + +struct SourceCredentialProvider { + source: Arc, +} + +#[async_trait] +impl CredentialProvider for SourceCredentialProvider { + async fn credentials( + &self, + provider: &CatalogProvider, + ) -> Result { + self.source.credentials(provider).await.map_err(|err| { + tracing::warn!( + provider = %provider.id(), + error = %err, + "LLM credentials could not be resolved for this attempt" + ); + match err { + ResolveError::NotConfigured(provider) => { + CredentialError::NotConfigured { provider } + } + ResolveError::SchemeMismatch { provider, .. } => { + CredentialError::SchemeMismatch { provider } + } + other => CredentialError::NotConfigured { + provider: other.provider().clone(), + }, + } + }) + } } diff --git a/lib/foundation/fabro-auth/src/env_source.rs b/lib/foundation/fabro-auth/src/env_source.rs index afe73f3a5..b4b16762b 100644 --- a/lib/foundation/fabro-auth/src/env_source.rs +++ b/lib/foundation/fabro-auth/src/env_source.rs @@ -2,23 +2,21 @@ use std::collections::HashMap; use std::sync::Arc; use async_trait::async_trait; -use fabro_model::{Catalog, ProviderId}; -use fabro_static::EnvVars; use fabro_vault::Vault; +use lithos_llm::catalog::{Catalog, CatalogProvider, ProviderId}; +use lithos_llm::credentials::Credentials; use tokio::sync::RwLock as AsyncRwLock; -use crate::resolve::apply_openai_codex_api_context; -use crate::{CredentialSource, EnvLookup, ResolvedCredentials, VaultCredentialSource}; +use crate::{CredentialSource, EnvLookup, ResolveError, VaultCredentialSource}; /// A credential source for provider credentials declared as `env:`. /// -/// This public SDK facade does not resolve `{{ env.NAME }}` settings -/// interpolation. Provider extra headers can use literals, but secret -/// interpolation requires a vault-backed source. +/// This public SDK facade does not resolve `{{ secrets.NAME }}` header +/// interpolation, so providers whose headers come from the vault stay +/// unconfigured here. #[derive(Clone)] pub struct EnvCredentialSource { - inner: VaultCredentialSource, - env_lookup: EnvLookup, + inner: VaultCredentialSource, } impl EnvCredentialSource { @@ -34,13 +32,8 @@ impl EnvCredentialSource { #[must_use] pub fn with_env_lookup(env_lookup: EnvLookup) -> Self { let vault = Arc::new(AsyncRwLock::new(Vault::from_entries(HashMap::new()))); - let inner_lookup = Arc::clone(&env_lookup); - let inner = VaultCredentialSource::with_env_lookup(vault, move |name| inner_lookup(name)); - Self { inner, env_lookup } - } - - fn lookup(&self, name: &str) -> Option { - (self.env_lookup)(name) + let inner = VaultCredentialSource::with_env_lookup(vault, move |name| env_lookup(name)); + Self { inner } } } @@ -59,18 +52,8 @@ impl Default for EnvCredentialSource { #[async_trait] impl CredentialSource for EnvCredentialSource { - async fn resolve(&self, catalog: &Catalog) -> anyhow::Result { - let mut resolved = self.inner.resolve(catalog).await?; - if let (Some(account_id), Some(credential)) = ( - self.lookup(EnvVars::CHATGPT_ACCOUNT_ID), - resolved - .credentials - .iter_mut() - .find(|credential| credential.provider == ProviderId::openai()), - ) { - apply_openai_codex_api_context(credential, Some(&account_id), self.env_lookup.as_ref()); - } - Ok(resolved) + async fn credentials(&self, provider: &CatalogProvider) -> Result { + self.inner.credentials(provider).await } async fn configured_providers(&self, catalog: &Catalog) -> Vec { @@ -83,12 +66,11 @@ mod tests { use std::collections::HashMap; use std::sync::Arc; - use fabro_model::catalog::LlmCatalogSettings; - use fabro_model::{Catalog, ProviderId}; - use fabro_types::settings::interp::Namespace; + use lithos_llm::catalog::ProviderId; use super::EnvCredentialSource; use crate::CredentialSource; + use crate::test_support::test_catalog; fn test_source(entries: &[(&str, &str)]) -> EnvCredentialSource { let entries: HashMap = entries @@ -101,111 +83,24 @@ mod tests { #[tokio::test] async fn configured_providers_reads_injected_provider_env() { let source = test_source(&[("ANTHROPIC_API_KEY", "anthropic-key")]); - let catalog = Catalog::from_builtin().unwrap(); - - assert_eq!(source.configured_providers(&catalog).await, vec![ - ProviderId::anthropic() + assert_eq!(source.configured_providers(&test_catalog()).await, vec![ + ProviderId::new("anthropic") ]); } - #[tokio::test] - async fn resolve_builds_openai_codex_env_credential() { - let source = test_source(&[ - ("OPENAI_API_KEY", "openai-key"), - ("CHATGPT_ACCOUNT_ID", "acct_123"), - ("OPENAI_PROJECT_ID", "project_123"), - ]); - let catalog = Catalog::from_builtin().unwrap(); - - let resolved = source.resolve(&catalog).await.unwrap(); - let credential = resolved.credentials.first().unwrap(); - - assert_eq!(credential.provider, ProviderId::openai()); - assert!(credential.codex_mode); - assert_eq!( - credential.base_url.as_deref(), - Some("https://chatgpt.com/backend-api/codex") - ); - assert_eq!( - credential.extra_headers.get("ChatGPT-Account-Id"), - Some(&"acct_123".to_string()) - ); - assert_eq!(credential.project_id.as_deref(), Some("project_123")); - } - - #[tokio::test] - async fn env_settings_interpolation_remains_unsupported() { - let settings: LlmCatalogSettings = toml::from_str( - r#" -[providers.acme] -display_name = "Acme" -adapter = "openai_compatible" -agent_profile = "openai" -base_url = "https://api.acme.test/v1" - -[providers.acme.auth] -credentials = ["env:ACME_API_KEY"] - -[providers.acme.extra_headers] -x-account = "{{ env.ACME_ACCOUNT }}" -"#, - ) - .unwrap(); - let catalog = Catalog::from_builtin_with_overrides(&settings).unwrap(); - let source = test_source(&[("ACME_API_KEY", "acme-key"), ("ACME_ACCOUNT", "account-id")]); - - let resolved = source.resolve(&catalog).await.unwrap(); - - assert!( - resolved - .credentials - .iter() - .all(|credential| credential.provider != ProviderId::new("acme")) - ); - assert!(resolved.auth_issues.iter().any(|(provider, issue)| { - provider == &ProviderId::new("acme") - && matches!( - issue, - crate::ResolveError::Interpolation { source, .. } - if source.namespace == Namespace::Env - ) - })); - } - #[tokio::test] async fn modal_env_vars_do_not_replace_vault_secrets() { - let settings: LlmCatalogSettings = toml::from_str( - r#" -[providers.modal] -enabled = true -base_url = "https://example--kimi-k3.modal.run/v1" -"#, - ) - .unwrap(); - let catalog = Catalog::from_builtin_with_overrides(&settings).unwrap(); let source = test_source(&[ ("MODAL_TOKEN_ID", "wk-test"), ("MODAL_TOKEN_SECRET", "ws-test"), ]); + let catalog = test_catalog(); let modal = ProviderId::new("modal"); - assert!(!source.configured_providers(&catalog).await.contains(&modal)); - - let resolved = source.resolve(&catalog).await.unwrap(); - - assert!( - resolved - .credentials - .iter() - .all(|credential| credential.provider != modal) - ); - assert!(resolved.auth_issues.iter().any(|(provider, issue)| { - provider == &modal - && matches!( - issue, - crate::ResolveError::Interpolation { source, .. } - if source.namespace == Namespace::Secrets - ) - })); + let err = source + .credentials(catalog.provider("modal").unwrap()) + .await + .unwrap_err(); + assert!(matches!(err, crate::ResolveError::Interpolation { .. })); } } diff --git a/lib/foundation/fabro-auth/src/extra_headers_source.rs b/lib/foundation/fabro-auth/src/extra_headers_source.rs index 29ad806cd..0226e1fee 100644 --- a/lib/foundation/fabro-auth/src/extra_headers_source.rs +++ b/lib/foundation/fabro-auth/src/extra_headers_source.rs @@ -2,15 +2,18 @@ use std::collections::HashMap; use std::sync::Arc; use async_trait::async_trait; -use fabro_model::{Catalog, ProviderId}; +use lithos_llm::catalog::{Catalog, CatalogProvider, ProviderId}; +use lithos_llm::credentials::{CredentialHeader, Credentials, SecretValue}; -use crate::credential_source::{CredentialSource, ResolvedCredentials}; +use crate::ResolveError; +use crate::credential_source::CredentialSource; /// Decorates another [`CredentialSource`] by appending fixed extra headers to -/// every credential it resolves. +/// every HTTP credential it resolves. /// /// Headers already present on a credential (for example from explicit -/// provider configuration) are left untouched. +/// provider configuration) are left untouched. AWS-signed credentials carry +/// no header list and pass through unchanged. pub struct ExtraHeadersCredentialSource { inner: Arc, headers: HashMap, @@ -25,21 +28,24 @@ impl ExtraHeadersCredentialSource { #[async_trait] impl CredentialSource for ExtraHeadersCredentialSource { - async fn resolve(&self, catalog: &Catalog) -> anyhow::Result { - let mut resolved = self.inner.resolve(catalog).await?; - for credential in &mut resolved.credentials { + async fn credentials(&self, provider: &CatalogProvider) -> Result { + let mut credentials = self.inner.credentials(provider).await?; + if let Credentials::Http(http) = &mut credentials { for (name, value) in &self.headers { - if credential + if http .extra_headers - .keys() - .any(|existing| existing.eq_ignore_ascii_case(name)) + .iter() + .any(|existing| existing.name.eq_ignore_ascii_case(name)) { continue; } - credential.extra_headers.insert(name.clone(), value.clone()); + http.extra_headers.push(CredentialHeader::new( + name.clone(), + SecretValue::new(value.clone()), + )); } } - Ok(resolved) + Ok(credentials) } async fn configured_providers(&self, catalog: &Catalog) -> Vec { @@ -49,31 +55,35 @@ impl CredentialSource for ExtraHeadersCredentialSource { #[cfg(test)] mod tests { + use lithos_llm::credentials::HttpAuthentication; + use super::*; - use crate::{ApiCredential, ResolveError}; + use crate::test_support::test_catalog; struct StubSource { - credentials: Vec, - auth_issue_provider: Option, configured_providers: Vec, + existing_header: Option<(String, String)>, } #[async_trait] impl CredentialSource for StubSource { - async fn resolve(&self, _catalog: &Catalog) -> anyhow::Result { - Ok(ResolvedCredentials { - credentials: self.credentials.clone(), - auth_issues: self - .auth_issue_provider - .iter() - .map(|provider| { - ( - provider.clone(), - ResolveError::RefreshTokenMissing(provider.clone()), - ) - }) - .collect(), - }) + async fn credentials( + &self, + provider: &CatalogProvider, + ) -> Result { + if provider.id().as_str() == "bedrock" { + return Ok(Credentials::AwsDefaultChain { region: None }); + } + let mut credentials = Credentials::bearer(SecretValue::new("key")); + if let (Credentials::Http(http), Some((name, value))) = + (&mut credentials, &self.existing_header) + { + http.extra_headers.push(CredentialHeader::new( + name.clone(), + SecretValue::new(value.clone()), + )); + } + Ok(credentials) } async fn configured_providers(&self, _catalog: &Catalog) -> Vec { @@ -81,95 +91,60 @@ mod tests { } } - fn credential(provider: ProviderId, extra_headers: HashMap) -> ApiCredential { - ApiCredential { - provider, - auth_header: None, - extra_headers, - base_url: None, - codex_mode: false, - org_id: None, - project_id: None, - } - } - - #[tokio::test] - async fn appends_headers_to_every_resolved_credential() { - let source = ExtraHeadersCredentialSource::new( - Arc::new(StubSource { - credentials: vec![ - credential(ProviderId::anthropic(), HashMap::new()), - credential(ProviderId::openai(), HashMap::new()), - ], - auth_issue_provider: None, - configured_providers: Vec::new(), - }), - HashMap::from([("x-session-id".to_string(), "run-123".to_string())]), - ); - - let resolved = source.resolve(Catalog::builtin()).await.unwrap(); - - assert_eq!(resolved.credentials.len(), 2); - for credential in &resolved.credentials { - assert_eq!( - credential - .extra_headers - .get("x-session-id") - .map(String::as_str), - Some("run-123") - ); - } - } - - #[tokio::test] - async fn preserves_case_insensitive_headers_already_set_on_a_credential() { - let source = ExtraHeadersCredentialSource::new( - Arc::new(StubSource { - credentials: vec![credential( - ProviderId::new("openrouter"), - HashMap::from([("X-Session-Id".to_string(), "configured".to_string())]), - )], - auth_issue_provider: None, - configured_providers: Vec::new(), - }), - HashMap::from([("x-session-id".to_string(), "run-123".to_string())]), - ); - - let resolved = source.resolve(Catalog::builtin()).await.unwrap(); - - assert_eq!( - resolved.credentials[0] + fn header<'a>(credentials: &'a Credentials, name: &str) -> Option<&'a str> { + match credentials { + Credentials::Http(http) => http .extra_headers - .get("X-Session-Id") - .map(String::as_str), - Some("configured") - ); - assert_eq!(resolved.credentials[0].extra_headers.len(), 1); + .iter() + .find(|header| header.name.eq_ignore_ascii_case(name)) + .map(|header| header.value.expose_secret()), + _ => None, + } } #[tokio::test] - async fn passes_through_auth_issues_and_configured_providers() { - let auth_issue_provider = ProviderId::anthropic(); - let configured_provider = ProviderId::gemini(); + async fn appends_headers_to_http_credentials_only() { + let catalog = test_catalog(); let source = ExtraHeadersCredentialSource::new( Arc::new(StubSource { - credentials: vec![credential(ProviderId::openai(), HashMap::new())], - auth_issue_provider: Some(auth_issue_provider.clone()), - configured_providers: vec![configured_provider.clone()], + configured_providers: Vec::new(), + existing_header: None, }), HashMap::from([("x-session-id".to_string(), "run-123".to_string())]), ); + let openai = source + .credentials(catalog.provider("openai").unwrap()) + .await + .unwrap(); + assert!(matches!( + &openai, + Credentials::Http(http) if matches!(http.auth, HttpAuthentication::Bearer(_)) + )); + assert_eq!(header(&openai, "x-session-id"), Some("run-123")); + let bedrock = source + .credentials(catalog.provider("bedrock").unwrap()) + .await + .unwrap(); + assert!(matches!(bedrock, Credentials::AwsDefaultChain { .. })); + } - let resolved = source.resolve(Catalog::builtin()).await.unwrap(); - let [(reported_provider, ResolveError::RefreshTokenMissing(error_provider))] = - resolved.auth_issues.as_slice() - else { - panic!("expected the inner source's refresh-token issue"); - }; - assert_eq!(reported_provider, &auth_issue_provider); - assert_eq!(error_provider, &auth_issue_provider); - - let providers = source.configured_providers(Catalog::builtin()).await; - assert_eq!(providers, vec![configured_provider]); + #[tokio::test] + async fn preserves_case_insensitive_headers_already_set() { + let catalog = test_catalog(); + let source = ExtraHeadersCredentialSource::new( + Arc::new(StubSource { + configured_providers: vec![ProviderId::new("openai")], + existing_header: Some(("X-Session-Id".to_string(), "configured".to_string())), + }), + HashMap::from([("x-session-id".to_string(), "run-123".to_string())]), + ); + let credentials = source + .credentials(catalog.provider("openai").unwrap()) + .await + .unwrap(); + assert_eq!(header(&credentials, "x-session-id"), Some("configured")); + assert_eq!(source.configured_providers(&catalog).await, vec![ + ProviderId::new("openai") + ]); } } diff --git a/lib/foundation/fabro-auth/src/lib.rs b/lib/foundation/fabro-auth/src/lib.rs index e54266f86..a4e656ea2 100644 --- a/lib/foundation/fabro-auth/src/lib.rs +++ b/lib/foundation/fabro-auth/src/lib.rs @@ -1,5 +1,7 @@ +mod api_key_source; mod context; mod credential; +mod credential_ref; mod credential_source; mod env_source; mod extra_headers_source; @@ -14,15 +16,17 @@ mod vault_source; pub mod strategies; +pub use api_key_source::ApiKeyCredentialSource; pub use context::{AuthContextRequest, AuthContextResponse}; -pub use credential::{ApiKeyHeader, OAuthConfig, OAuthCredential, OAuthTokens}; -pub use credential_source::{CredentialSource, ResolvedCredentials}; +pub use credential::{OAuthConfig, OAuthCredential, OAuthTokens}; +pub use credential_ref::{CredentialRef, CredentialRefParseError}; +pub use credential_source::{CredentialSource, ResolvedCredentials, lithos_credentials}; pub use env_source::EnvCredentialSource; pub use extra_headers_source::ExtraHeadersCredentialSource; pub use refresh::refresh_oauth_credential; pub use resolve::{ - ApiCredential, CredentialResolver, CredentialUsage, EnvLookup, ResolveError, - ResolvedCredential, auth_issue_message, build_api_key_header, + CredentialResolver, EnvLookup, ResolveError, accepts_api_key, auth_issue_message, + credential_refs, credentials_for_api_key, env_var_names, expected_vault_secret_name, }; pub use sql_vault_source::SqlVaultCredentialSource; pub use strategy::{ diff --git a/lib/foundation/fabro-auth/src/resolve.rs b/lib/foundation/fabro-auth/src/resolve.rs index dd12d1a1e..3e8f2e5f9 100644 --- a/lib/foundation/fabro-auth/src/resolve.rs +++ b/lib/foundation/fabro-auth/src/resolve.rs @@ -1,15 +1,28 @@ -use std::collections::HashMap; +//! Secret resolution for one catalog provider. +//! +//! A provider's `metadata.fabro.credentials` names where its secret lives. +//! [`CredentialResolver`] walks that list against the vault and the process +//! environment, refreshes an expired OAuth credential, and shapes the result +//! into the lithos [`Credentials`] the provider's declared auth scheme +//! expects. + +use std::collections::BTreeMap; use std::sync::Arc; -use fabro_model::catalog::CatalogProvider; -use fabro_model::{ApiKeyHeaderPolicy, Catalog, CredentialRef, ProviderId}; use fabro_static::EnvVars; +use fabro_types::catalog_policy::{self, ProviderPolicy}; +use fabro_types::provider_ids; use fabro_types::settings::{InterpString, ResolveCtx, ResolveError as InterpResolveError}; use fabro_vault::{SecretType, Vault}; +use lithos_llm::catalog::{AuthScheme, Catalog, CatalogProvider, ProviderId}; +use lithos_llm::credentials::{ + CredentialHeader, Credentials, HttpAuthentication, HttpCredentials, SecretValue, +}; use tokio::sync::RwLock as AsyncRwLock; use tokio::task::spawn_blocking; -use crate::credential::{ApiKeyHeader, OAuthCredential}; +use crate::credential::OAuthCredential; +use crate::credential_ref::{CredentialRef, CredentialRefParseError}; use crate::refresh::refresh_oauth_credential; use crate::vault_ext::{ VaultLookupError, vault_get_oauth, vault_get_token, vault_set_oauth, vault_token_lookup, @@ -17,136 +30,46 @@ use crate::vault_ext::{ pub type EnvLookup = Arc Option + Send + Sync>; -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum CredentialUsage { - ApiRequest, -} +const CHATGPT_ACCOUNT_ID_HEADER: &str = "ChatGPT-Account-Id"; +const OPENAI_ORGANIZATION_HEADER: &str = "OpenAI-Organization"; +const OPENAI_PROJECT_HEADER: &str = "OpenAI-Project"; -#[derive(Debug, Clone, PartialEq, Eq)] +/// A secret found for a provider, before it is shaped into [`Credentials`]. +#[derive(Clone, PartialEq, Eq)] pub(crate) enum ResolvedSecret { ApiKey(String), OAuth { credential: Box, vault_name: String, }, - /// Opaque AWS SigV4 source: no static secret; the adapter signs requests - /// using the AWS default credential chain. - AwsSigv4, + /// No static secret: the adapter signs with the AWS default chain. + AwsDefaultChain, } -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct ApiCredential { - pub provider: ProviderId, - pub auth_header: Option, - pub extra_headers: HashMap, - pub base_url: Option, - pub codex_mode: bool, - pub org_id: Option, - pub project_id: Option, -} - -impl ApiCredential { - /// Build an `ApiCredential` from an API key using the supplied catalog for - /// auth header policy and provider base URL. - pub fn from_api_key( - provider: impl Into, - key: String, - catalog: &Catalog, - ) -> Result { - let provider_id = provider.into(); - let provider = catalog - .provider(&provider_id) - .ok_or_else(|| ResolveError::NotConfigured(provider_id.clone()))?; - let auth_header = auth_header_for_catalog_provider(provider, key)?; - Ok(Self { - provider: provider_id, - auth_header: Some(auth_header), - extra_headers: HashMap::new(), - base_url: provider.base_url.clone(), - codex_mode: false, - org_id: None, - project_id: None, - }) - } - - /// Build an `ApiCredential` for a provider that authenticates with request - /// headers instead of an API key, such as Modal's proxy-token pair. - #[must_use] - pub fn with_extra_headers( - provider: impl Into, - extra_headers: HashMap, - ) -> Self { - Self { - provider: provider.into(), - auth_header: None, - extra_headers, - base_url: None, - codex_mode: false, - org_id: None, - project_id: None, +impl std::fmt::Debug for ResolvedSecret { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::ApiKey(_) => f.write_str("ApiKey()"), + Self::OAuth { vault_name, .. } => f + .debug_struct("OAuth") + .field("vault_name", vault_name) + .finish_non_exhaustive(), + Self::AwsDefaultChain => f.write_str("AwsDefaultChain"), } } } -const OPENAI_CODEX_BASE_URL: &str = "https://chatgpt.com/backend-api/codex"; -const CHATGPT_ACCOUNT_ID_HEADER: &str = "ChatGPT-Account-Id"; -const ORIGINATOR_HEADER: &str = "originator"; -const FABRO_ORIGINATOR: &str = "fabro"; - -pub(crate) fn apply_openai_api_env_context( - credential: &mut ApiCredential, - env_lookup: &(dyn Fn(&str) -> Option + Send + Sync), -) { - credential.org_id = env_lookup(EnvVars::OPENAI_ORG_ID); - credential.project_id = env_lookup(EnvVars::OPENAI_PROJECT_ID); -} - -pub(crate) fn apply_openai_codex_api_context( - credential: &mut ApiCredential, - account_id: Option<&str>, - env_lookup: &(dyn Fn(&str) -> Option + Send + Sync), -) { - apply_openai_api_env_context(credential, env_lookup); - if let Some(account_id) = account_id { - credential.extra_headers.insert( - CHATGPT_ACCOUNT_ID_HEADER.to_string(), - account_id.to_string(), - ); - } - credential - .extra_headers - .insert(ORIGINATOR_HEADER.to_string(), FABRO_ORIGINATOR.to_string()); - credential.base_url = Some(OPENAI_CODEX_BASE_URL.to_string()); - credential.codex_mode = true; -} - -#[must_use] -pub fn build_api_key_header(policy: ApiKeyHeaderPolicy, key: String) -> ApiKeyHeader { - match policy { - ApiKeyHeaderPolicy::Bearer => ApiKeyHeader::Bearer(key), - ApiKeyHeaderPolicy::Custom { name } => ApiKeyHeader::Custom { name, value: key }, - } -} - -fn auth_header_for_catalog_provider( - provider: &CatalogProvider, - key: String, -) -> Result { - let Some(auth) = &provider.auth else { - return Err(ResolveError::NotConfigured(provider.id.clone())); - }; - Ok(build_api_key_header(auth.header.clone(), key)) -} - -#[derive(Debug, Clone, PartialEq, Eq)] -pub enum ResolvedCredential { - Api(ApiCredential), -} - #[derive(Debug, thiserror::Error)] pub enum ResolveError { #[error("{0} is not configured")] NotConfigured(ProviderId), + #[error("{provider} declares an invalid credential reference `{reference}`: {source}")] + InvalidCredentialRef { + provider: ProviderId, + reference: String, + #[source] + source: CredentialRefParseError, + }, #[error("{provider} header interpolation failed: {source}")] Interpolation { provider: ProviderId, @@ -174,35 +97,240 @@ pub enum ResolveError { }, #[error("{0} requires re-authentication: missing refresh token")] RefreshTokenMissing(ProviderId), + #[error("{provider} resolved a secret its `{scheme}` auth scheme cannot use")] + SchemeMismatch { + provider: ProviderId, + scheme: String, + }, +} + +impl ResolveError { + #[must_use] + pub fn provider(&self) -> &ProviderId { + match self { + Self::NotConfigured(provider) + | Self::RefreshTokenMissing(provider) + | Self::InvalidCredentialRef { provider, .. } + | Self::Interpolation { provider, .. } + | Self::VaultSchemaMismatch { provider, .. } + | Self::VaultDecodeFailed { provider, .. } + | Self::RefreshFailed { provider, .. } + | Self::SchemeMismatch { provider, .. } => provider, + } + } } #[must_use] pub fn auth_issue_message(provider: &ProviderId, err: &ResolveError) -> String { - let provider_name = provider.display_name(); match err { - ResolveError::NotConfigured(_) => { - format!("{provider_name} is not configured") - } + ResolveError::NotConfigured(_) => format!("{provider} is not configured"), + ResolveError::InvalidCredentialRef { + reference, source, .. + } => format!("{provider} declares an invalid credential reference `{reference}`: {source}"), ResolveError::Interpolation { source, .. } => { - format!("{provider_name} header interpolation failed: {source}") - } - ResolveError::VaultSchemaMismatch { name, actual, .. } => { - format!( - "{provider_name} vault credential '{name}' has schema {actual:?}, expected Token or Oauth" - ) + format!("{provider} header interpolation failed: {source}") } + ResolveError::VaultSchemaMismatch { name, actual, .. } => format!( + "{provider} vault credential '{name}' has schema {actual:?}, expected Token or Oauth" + ), ResolveError::VaultDecodeFailed { name, source, .. } => { - format!("{provider_name} vault credential '{name}' is not valid OAuth JSON: {source}") + format!("{provider} vault credential '{name}' is not valid OAuth JSON: {source}") } ResolveError::RefreshFailed { source, .. } => { - format!("{provider_name} requires re-authentication: {source}") + format!("{provider} requires re-authentication: {source}") } ResolveError::RefreshTokenMissing(_) => { - format!("{provider_name} requires re-authentication: refresh token missing") + format!("{provider} requires re-authentication: refresh token missing") + } + ResolveError::SchemeMismatch { scheme, .. } => { + format!("{provider} resolved a secret its `{scheme}` auth scheme cannot use") } } } +/// The credential references a provider declares, in resolution order. +pub fn credential_refs(provider: &CatalogProvider) -> Result, ResolveError> { + credential_refs_from_policy(provider.id(), &catalog_policy::provider_policy(provider)) +} + +fn credential_refs_from_policy( + provider: &ProviderId, + policy: &ProviderPolicy, +) -> Result, ResolveError> { + policy + .credentials + .iter() + .map(|reference| { + reference + .parse() + .map_err(|source| ResolveError::InvalidCredentialRef { + provider: provider.clone(), + reference: reference.clone(), + source, + }) + }) + .collect() +} + +/// The vault entry an operator should create to configure `provider`, when +/// the provider reads a vault secret. +#[must_use] +pub fn expected_vault_secret_name(provider: &CatalogProvider) -> Option { + credential_refs(provider) + .ok()? + .into_iter() + .find_map(|reference| match reference { + CredentialRef::Vault(name) => Some(name), + CredentialRef::Env(_) | CredentialRef::AwsSigv4 => None, + }) +} + +/// The environment variables an operator can set to configure `provider`. +#[must_use] +pub fn env_var_names(provider: &CatalogProvider) -> Vec { + credential_refs(provider) + .unwrap_or_default() + .into_iter() + .filter_map(|reference| match reference { + CredentialRef::Env(name) => Some(name), + CredentialRef::Vault(_) | CredentialRef::AwsSigv4 => None, + }) + .collect() +} + +/// Whether the provider takes a single API key an operator can paste in. +#[must_use] +pub fn accepts_api_key(provider: &CatalogProvider) -> bool { + matches!( + provider.auth(), + AuthScheme::Bearer { .. } | AuthScheme::Header { .. } | AuthScheme::BedrockBearer + ) && credential_refs(provider).is_ok_and(|refs| { + refs.iter() + .any(|reference| matches!(reference, CredentialRef::Vault(_) | CredentialRef::Env(_))) + }) +} + +fn auth_scheme_name(scheme: &AuthScheme) -> &'static str { + match scheme { + AuthScheme::None => "none", + AuthScheme::Bearer { .. } => "bearer", + AuthScheme::Header { .. } => "header", + AuthScheme::Headers => "headers", + AuthScheme::Aws { .. } => "aws", + AuthScheme::BedrockBearer => "bedrock_bearer", + _ => "unknown", + } +} + +/// Shapes a caller-supplied API key into the provider's credentials. +/// +/// Used to validate a key before it is stored. Extra headers that need vault +/// secrets are resolved against `vault`. +pub fn credentials_for_api_key( + provider: &CatalogProvider, + key: String, + vault: &Vault, +) -> Result { + let extra_headers = resolved_extra_headers(vault, provider)?; + shape_secret(provider, ResolvedSecret::ApiKey(key), extra_headers, None) +} + +/// Resolves a provider's `extra_headers` interpolation against the vault. +/// +/// Resolved header values may contain secrets; keep this path free of value +/// logging. +fn resolved_extra_headers( + vault: &Vault, + provider: &CatalogProvider, +) -> Result, ResolveError> { + let policy = catalog_policy::provider_policy(provider); + let mut ctx = + ResolveCtx::new().with_secrets(|secret_name| vault_token_lookup(vault, secret_name)); + resolve_extra_headers(provider.id(), &policy.extra_headers, &mut ctx) +} + +pub(crate) fn resolve_extra_headers( + provider: &ProviderId, + headers: &BTreeMap, + ctx: &mut ResolveCtx<'_>, +) -> Result, ResolveError> { + headers + .iter() + .map(|(name, source)| { + let value = InterpString::parse(source) + .resolve_with(ctx) + .map_err(|source| ResolveError::Interpolation { + provider: provider.clone(), + source, + })?; + Ok(CredentialHeader::new(name.clone(), SecretValue::new(value))) + }) + .collect() +} + +fn shape_secret( + provider: &CatalogProvider, + secret: ResolvedSecret, + mut extra_headers: Vec, + env_lookup: Option<&EnvLookup>, +) -> Result { + let scheme = provider.auth(); + let mismatch = || ResolveError::SchemeMismatch { + provider: provider.id().clone(), + scheme: auth_scheme_name(scheme).to_string(), + }; + if let Some(env_lookup) = env_lookup.filter(|_| provider.id().as_str() == provider_ids::OPENAI) + { + for (variable, header) in [ + (EnvVars::OPENAI_ORG_ID, OPENAI_ORGANIZATION_HEADER), + (EnvVars::OPENAI_PROJECT_ID, OPENAI_PROJECT_HEADER), + ] { + if let Some(value) = env_lookup(variable) { + extra_headers.push(CredentialHeader::new(header, SecretValue::new(value))); + } + } + } + match (scheme, secret) { + (AuthScheme::Bearer { .. }, ResolvedSecret::ApiKey(key)) => Ok(http_credentials( + HttpAuthentication::Bearer(SecretValue::new(key)), + extra_headers, + )), + (AuthScheme::Bearer { .. }, ResolvedSecret::OAuth { credential, .. }) => { + if let Some(account_id) = &credential.account_id { + extra_headers.push(CredentialHeader::new( + CHATGPT_ACCOUNT_ID_HEADER, + SecretValue::new(account_id.clone()), + )); + } + Ok(http_credentials( + HttpAuthentication::Bearer(SecretValue::new( + credential.tokens.access_token.clone(), + )), + extra_headers, + )) + } + (AuthScheme::Header { name }, ResolvedSecret::ApiKey(key)) => Ok(http_credentials( + HttpAuthentication::Header(CredentialHeader::new(name.clone(), SecretValue::new(key))), + extra_headers, + )), + (AuthScheme::BedrockBearer | AuthScheme::Aws { .. }, ResolvedSecret::ApiKey(key)) => { + Ok(Credentials::BedrockBearer(SecretValue::new(key))) + } + (AuthScheme::Aws { region }, ResolvedSecret::AwsDefaultChain) => { + Ok(Credentials::AwsDefaultChain { + region: region.clone(), + }) + } + _ => Err(mismatch()), + } +} + +fn http_credentials(auth: HttpAuthentication, extra_headers: Vec) -> Credentials { + let mut credentials = HttpCredentials::new(auth); + credentials.extra_headers = extra_headers; + Credentials::Http(credentials) +} + #[derive(Clone)] pub struct CredentialResolver { vault: Arc>, @@ -224,126 +352,132 @@ impl CredentialResolver { Self { vault, env_lookup } } - pub async fn resolve( - &self, - provider: impl Into, - _usage: CredentialUsage, - catalog: &Catalog, - ) -> Result { - let provider_id = provider.into(); - let Some(catalog_provider) = catalog.provider(&provider_id) else { - return Err(ResolveError::NotConfigured(provider_id)); - }; - if catalog_provider.auth.is_none() { - let vault = self.vault.read().await; - return Self::api_credential_from_provider_auth(&vault, catalog_provider, catalog) - .map(ResolvedCredential::Api); + /// Resolves `provider`'s credentials for one request attempt. + /// + /// An expired OAuth credential is refreshed and the refreshed tokens are + /// written back to the vault before the credentials are returned. + pub async fn resolve(&self, provider: &CatalogProvider) -> Result { + let provider_id = provider.id().clone(); + match provider.auth() { + AuthScheme::None => { + let vault = self.vault.read().await; + let headers = resolved_extra_headers(&vault, provider)?; + return Ok(Credentials::headers(headers)); + } + AuthScheme::Headers => { + let vault = self.vault.read().await; + let headers = resolved_extra_headers(&vault, provider)?; + if headers.is_empty() { + return Err(ResolveError::NotConfigured(provider_id)); + } + return Ok(Credentials::headers(headers)); + } + _ => {} } - let initial_secret = { + + let (initial_secret, extra_headers) = { let vault = self.vault.read().await; - self.find_credential(&vault, catalog_provider)? + ( + self.find_secret(&vault, provider)?, + resolved_extra_headers(&vault, provider)?, + ) }; - let secret = if let ResolvedSecret::OAuth { - credential, - vault_name, - } = &initial_secret - { - if !credential.needs_refresh() { - initial_secret - } else if credential.tokens.refresh_token.is_none() { - return Err(ResolveError::RefreshTokenMissing(provider_id.clone())); - } else { - let refreshed = refresh_oauth_credential(credential) + let secret = match initial_secret { + ResolvedSecret::OAuth { + credential, + vault_name, + } if credential.needs_refresh() => { + if credential.tokens.refresh_token.is_none() { + return Err(ResolveError::RefreshTokenMissing(provider_id)); + } + let refreshed = refresh_oauth_credential(&credential) .await .map_err(|source| ResolveError::RefreshFailed { provider: provider_id.clone(), source, })?; - let refreshed_for_store = refreshed.clone(); - let vault_name_for_store = vault_name.clone(); - let vault = Arc::clone(&self.vault); - spawn_blocking(move || { - let mut vault = vault.blocking_write(); - vault_set_oauth(&mut vault, &vault_name_for_store, &refreshed_for_store) - .map(|_| ()) - .map_err(anyhow::Error::from) - }) - .await - .map_err(|join_err| ResolveError::RefreshFailed { - provider: provider_id.clone(), - source: anyhow::Error::from(join_err), - })? - .map_err(|source| ResolveError::RefreshFailed { - provider: provider_id.clone(), - source, - })?; + self.persist_oauth(&provider_id, &vault_name, &refreshed) + .await?; ResolvedSecret::OAuth { credential: Box::new(refreshed), - vault_name: vault_name.clone(), + vault_name, } } - } else { - initial_secret + secret => secret, }; - let vault = self.vault.read().await; - self.to_api_credential(&vault, &provider_id, &secret, catalog) - .map(ResolvedCredential::Api) + shape_secret(provider, secret, extra_headers, Some(&self.env_lookup)) } + async fn persist_oauth( + &self, + provider: &ProviderId, + vault_name: &str, + refreshed: &OAuthCredential, + ) -> Result<(), ResolveError> { + let refreshed = refreshed.clone(); + let vault_name = vault_name.to_string(); + let vault = Arc::clone(&self.vault); + spawn_blocking(move || { + let mut vault = vault.blocking_write(); + vault_set_oauth(&mut vault, &vault_name, &refreshed) + .map(|_| ()) + .map_err(anyhow::Error::from) + }) + .await + .map_err(|join_err| ResolveError::RefreshFailed { + provider: provider.clone(), + source: anyhow::Error::from(join_err), + })? + .map_err(|source| ResolveError::RefreshFailed { + provider: provider.clone(), + source, + }) + } + + /// Providers with credential material present, without refreshing + /// anything. Disabled providers are skipped. #[must_use] pub fn configured_providers(&self, vault: &Vault, catalog: &Catalog) -> Vec { catalog .providers() - .iter() - .filter(|provider| self.has_credential_material(vault, provider, catalog)) - .map(|provider| provider.id.clone()) + .filter(|provider| catalog_policy::provider_policy(provider).is_enabled()) + .filter(|provider| self.has_credential_material(vault, provider)) + .map(|provider| provider.id().clone()) .collect() } - fn find_credential( + fn has_credential_material(&self, vault: &Vault, provider: &CatalogProvider) -> bool { + match provider.auth() { + AuthScheme::None => resolved_extra_headers(vault, provider).is_ok(), + AuthScheme::Headers => { + resolved_extra_headers(vault, provider).is_ok_and(|headers| !headers.is_empty()) + } + _ => self.find_secret(vault, provider).is_ok(), + } + } + + fn find_secret( &self, vault: &Vault, provider: &CatalogProvider, ) -> Result { - let Some(auth) = &provider.auth else { - return Err(ResolveError::NotConfigured(provider.id.clone())); - }; - - for credential_ref in &auth.credentials { - if let Some(credential) = - self.credential_from_ref(vault, &provider.id, credential_ref)? - { - return Ok(credential); + for reference in credential_refs(provider)? { + if let Some(secret) = self.secret_from_ref(vault, provider.id(), &reference)? { + return Ok(secret); } } - - Err(ResolveError::NotConfigured(provider.id.clone())) + Err(ResolveError::NotConfigured(provider.id().clone())) } - fn has_credential_material( - &self, - vault: &Vault, - provider: &CatalogProvider, - catalog: &Catalog, - ) -> bool { - let Some(auth) = &provider.auth else { - return Self::resolved_extra_headers_for_catalog(vault, &provider.id, catalog).is_ok(); - }; - auth.credentials.iter().any(|credential_ref| { - self.credential_from_ref(vault, &provider.id, credential_ref) - .is_ok_and(|credential| credential.is_some()) - }) - } - - fn credential_from_ref( + fn secret_from_ref( &self, vault: &Vault, provider: &ProviderId, - credential_ref: &CredentialRef, + reference: &CredentialRef, ) -> Result, ResolveError> { - match credential_ref { + match reference { CredentialRef::Vault(name) => match vault_get_token(vault, name) { Ok(Some(token)) => Ok(Some(ResolvedSecret::ApiKey(token))), Ok(None) => Ok(None), @@ -361,148 +495,9 @@ impl CredentialResolver { Err(err) => Err(vault_lookup_error(provider, name, err)), }, CredentialRef::Env(name) => Ok((self.env_lookup)(name).map(ResolvedSecret::ApiKey)), - // AWS SigV4 is an opaque source: it always "resolves" (the adapter - // signs at request time from the AWS chain), no vault/env lookup. - CredentialRef::AwsSigv4 => Ok(Some(ResolvedSecret::AwsSigv4)), + CredentialRef::AwsSigv4 => Ok(Some(ResolvedSecret::AwsDefaultChain)), } } - - fn provider_base_url_for_catalog(provider: &ProviderId, catalog: &Catalog) -> Option { - catalog - .provider(provider) - .and_then(|provider| provider.base_url.clone()) - } - - fn resolved_extra_headers_for_catalog( - vault: &Vault, - provider: &ProviderId, - catalog: &Catalog, - ) -> Result, ResolveError> { - let Some(catalog_provider) = catalog.provider(provider) else { - return Ok(HashMap::new()); - }; - let mut ctx = - ResolveCtx::new().with_secrets(|secret_name| vault_token_lookup(vault, secret_name)); - resolve_extra_headers(provider, &catalog_provider.extra_headers, &mut ctx) - } - - fn to_api_credential( - &self, - vault: &Vault, - provider_id: &ProviderId, - secret: &ResolvedSecret, - catalog: &Catalog, - ) -> Result { - let base_url = Self::provider_base_url_for_catalog(provider_id, catalog); - match secret { - // Opaque AWS SigV4 source: carry the marker so the adapter signs - // with the AWS chain; no static secret resolved here. - ResolvedSecret::AwsSigv4 => Ok(ApiCredential { - provider: provider_id.clone(), - auth_header: Some(ApiKeyHeader::AwsSigv4), - extra_headers: Self::resolved_extra_headers_for_catalog( - vault, - provider_id, - catalog, - )?, - base_url, - codex_mode: false, - org_id: None, - project_id: None, - }), - ResolvedSecret::ApiKey(key) => { - let provider = catalog - .provider(provider_id) - .ok_or_else(|| ResolveError::NotConfigured(provider_id.clone()))?; - let auth_header = auth_header_for_catalog_provider(provider, key.clone())?; - let mut cred = ApiCredential { - provider: provider_id.clone(), - auth_header: Some(auth_header), - extra_headers: Self::resolved_extra_headers_for_catalog( - vault, - provider_id, - catalog, - )?, - base_url, - codex_mode: false, - org_id: None, - project_id: None, - }; - if provider_id == &ProviderId::openai() { - apply_openai_api_env_context(&mut cred, &*self.env_lookup); - } - Ok(cred) - } - ResolvedSecret::OAuth { credential, .. } => { - let mut api_credential = ApiCredential { - provider: provider_id.clone(), - auth_header: Some(ApiKeyHeader::Bearer(credential.tokens.access_token.clone())), - extra_headers: Self::resolved_extra_headers_for_catalog( - vault, - provider_id, - catalog, - )?, - base_url, - codex_mode: false, - org_id: None, - project_id: None, - }; - if provider_id == &ProviderId::openai() { - apply_openai_codex_api_context( - &mut api_credential, - credential.account_id.as_deref(), - &*self.env_lookup, - ); - } - Ok(api_credential) - } - } - } - - fn api_credential_from_provider_auth( - vault: &Vault, - provider: &CatalogProvider, - catalog: &Catalog, - ) -> Result { - if provider.auth.is_some() { - return Err(ResolveError::NotConfigured(provider.id.clone())); - } - let extra_headers = Self::resolved_extra_headers_for_catalog(vault, &provider.id, catalog)?; - Ok(ApiCredential { - provider: provider.id.clone(), - auth_header: None, - extra_headers, - base_url: Self::provider_base_url_for_catalog(&provider.id, catalog), - codex_mode: false, - org_id: None, - project_id: None, - }) - } -} - -/// Resolve a provider's `extra_headers` interpolation sources with `ctx`. -/// -/// Resolved header values may contain secrets; keep this path free of value -/// logging. Content-based redaction covers credential-shaped values on output -/// surfaces, but no mechanism redacts these exact values, so a low-entropy -/// header value that does not look like a credential is not caught. -pub(crate) fn resolve_extra_headers( - provider: &ProviderId, - headers: &HashMap, - ctx: &mut ResolveCtx<'_>, -) -> Result, ResolveError> { - headers - .iter() - .map(|(name, source)| { - let value = InterpString::parse(source) - .resolve_with(ctx) - .map_err(|source| ResolveError::Interpolation { - provider: provider.clone(), - source, - })?; - Ok((name.clone(), value)) - }) - .collect() } fn vault_lookup_error(provider: &ProviderId, name: &str, err: VaultLookupError) -> ResolveError { @@ -522,15 +517,13 @@ fn vault_lookup_error(provider: &ProviderId, name: &str, err: VaultLookupError) #[cfg(test)] mod tests { - use std::error::Error as _; - use chrono::{Duration, Utc}; - use fabro_model::catalog::LlmCatalogSettings; use httpmock::Method::POST; use httpmock::MockServer; use super::*; use crate::credential::{OAuthConfig, OAuthCredential, OAuthTokens}; + use crate::test_support::test_catalog; use crate::vault_ext::{vault_get_oauth, vault_set_oauth, vault_set_token}; fn oauth_credential(token_url: String, expires_at: chrono::DateTime) -> OAuthCredential { @@ -556,138 +549,86 @@ mod tests { CredentialResolver::with_env_lookup(Arc::new(AsyncRwLock::new(vault)), env_lookup) } - fn catalog_with(overrides: &str) -> Catalog { - let settings: LlmCatalogSettings = toml::from_str(overrides).unwrap(); - Catalog::from_builtin_with_overrides(&settings).unwrap() + fn empty_vault() -> Vault { + Vault::from_entries(std::collections::HashMap::new()) } - fn default_catalog() -> Catalog { - catalog_with("") + fn bearer_secret(credentials: &Credentials) -> &str { + match credentials { + Credentials::Http(HttpCredentials { + auth: HttpAuthentication::Bearer(secret), + .. + }) => secret.expose_secret(), + other => panic!("expected bearer credentials, got {other:?}"), + } } - /// A no-auth portkey provider whose only variation is its `extra_headers` - /// TOML lines. - fn portkey_catalog(extra_headers: &str) -> Catalog { - catalog_with(&format!( - r#" -[providers.portkey] -display_name = "Portkey Bedrock" -adapter = "anthropic" -agent_profile = "anthropic" -base_url = "https://api.portkey.ai/v1" - -[providers.portkey.extra_headers] -{extra_headers} - -[models."portkey-claude"] -provider = "portkey" -display_name = "Portkey Claude" -family = "claude" -default = true - -[models."portkey-claude".limits] -context_window = 200000 - -[models."portkey-claude".features] -tools = true -vision = true -reasoning = true -reasoning_effort = "levels" -"# - )) - } - - fn modal_catalog() -> Catalog { - catalog_with( - r#" -[providers.modal] -enabled = true -base_url = "https://example--kimi-k3.modal.run/v1" -"#, - ) + fn header_value<'a>(credentials: &'a Credentials, name: &str) -> Option<&'a str> { + match credentials { + Credentials::Http(http) => http + .extra_headers + .iter() + .find(|header| header.name.eq_ignore_ascii_case(name)) + .map(|header| header.value.expose_secret()), + _ => None, + } } #[tokio::test] - async fn resolve_openai_api_request_prefers_env_when_listed_first() { - let dir = tempfile::tempdir().unwrap(); - let mut vault = Vault::load(dir.path().join("secrets.json")).unwrap(); + async fn env_listed_first_wins_over_vault() { + let mut vault = empty_vault(); vault_set_token(&mut vault, "OPENAI_API_KEY", "vault-key").unwrap(); let resolver = test_resolver( vault, Arc::new(|name| (name == "OPENAI_API_KEY").then(|| "env-key".to_string())), ); - let catalog = default_catalog(); - - let resolved = resolver - .resolve(ProviderId::openai(), CredentialUsage::ApiRequest, &catalog) + let catalog = test_catalog(); + let credentials = resolver + .resolve(catalog.provider("openai").unwrap()) .await .unwrap(); - - let ResolvedCredential::Api(api) = resolved; - assert_eq!( - api.auth_header, - Some(ApiKeyHeader::Bearer("env-key".to_string())) - ); + assert_eq!(bearer_secret(&credentials), "env-key"); } #[tokio::test] - async fn resolve_moonshot_api_request_prefers_moonshot_env_key() { - let dir = tempfile::tempdir().unwrap(); - let vault = Vault::load(dir.path().join("secrets.json")).unwrap(); + async fn moonshot_falls_back_to_kimi_env_key() { let resolver = test_resolver( - vault, - Arc::new(|name| match name { - EnvVars::MOONSHOT_API_KEY => Some("moonshot-key".to_string()), - EnvVars::KIMI_API_KEY => Some("kimi-key".to_string()), - _ => None, - }), - ); - - let resolved = resolver - .resolve( - ProviderId::new("moonshot"), - CredentialUsage::ApiRequest, - &default_catalog(), - ) - .await - .unwrap(); - - let ResolvedCredential::Api(api) = resolved; - assert_eq!( - api.auth_header, - Some(ApiKeyHeader::Bearer("moonshot-key".to_string())) - ); - } - - #[tokio::test] - async fn resolve_moonshot_api_request_falls_back_to_kimi_env_key() { - let dir = tempfile::tempdir().unwrap(); - let vault = Vault::load(dir.path().join("secrets.json")).unwrap(); - let resolver = test_resolver( - vault, + empty_vault(), Arc::new(|name| (name == EnvVars::KIMI_API_KEY).then(|| "kimi-key".to_string())), ); - - let resolved = resolver - .resolve( - ProviderId::new("moonshot"), - CredentialUsage::ApiRequest, - &default_catalog(), - ) + let catalog = test_catalog(); + let credentials = resolver + .resolve(catalog.provider("moonshot").unwrap()) .await .unwrap(); - - let ResolvedCredential::Api(api) = resolved; - assert_eq!( - api.auth_header, - Some(ApiKeyHeader::Bearer("kimi-key".to_string())) - ); + assert_eq!(bearer_secret(&credentials), "kimi-key"); } #[tokio::test] - async fn resolve_openai_api_request_falls_back_to_codex_oauth_credential() { - let dir = tempfile::tempdir().unwrap(); - let mut vault = Vault::load(dir.path().join("secrets.json")).unwrap(); + async fn anthropic_uses_its_header_scheme() { + let mut vault = empty_vault(); + vault_set_token(&mut vault, "ANTHROPIC_API_KEY", "anthropic-key").unwrap(); + let resolver = test_resolver(vault, Arc::new(|_| None)); + let catalog = test_catalog(); + let credentials = resolver + .resolve(catalog.provider("anthropic").unwrap()) + .await + .unwrap(); + match credentials { + Credentials::Http(HttpCredentials { + auth: HttpAuthentication::Header(header), + .. + }) => { + assert_eq!(header.name, "x-api-key"); + assert_eq!(header.value.expose_secret(), "anthropic-key"); + } + other => panic!("expected header credentials, got {other:?}"), + } + } + + #[tokio::test] + async fn codex_oauth_becomes_a_bearer_with_account_header() { + let mut vault = empty_vault(); vault_set_oauth( &mut vault, crate::OPENAI_CODEX_VAULT_SECRET_NAME, @@ -698,465 +639,133 @@ base_url = "https://example--kimi-k3.modal.run/v1" ) .unwrap(); let resolver = test_resolver(vault, Arc::new(|_| None)); - let catalog = default_catalog(); - - let resolved = resolver - .resolve(ProviderId::openai(), CredentialUsage::ApiRequest, &catalog) + let catalog = test_catalog(); + let credentials = resolver + .resolve(catalog.provider("openai-codex").unwrap()) .await .unwrap(); - - let ResolvedCredential::Api(api) = resolved; + assert_eq!(bearer_secret(&credentials), "expired-access"); assert_eq!( - api.auth_header, - Some(ApiKeyHeader::Bearer("expired-access".to_string())) - ); - assert!(api.codex_mode); - assert_eq!( - api.base_url.as_deref(), - Some("https://chatgpt.com/backend-api/codex") + header_value(&credentials, CHATGPT_ACCOUNT_ID_HEADER), + Some("acct_123") ); } #[tokio::test] - async fn sigv4_provider_resolves_to_aws_sigv4_credential() { - let dir = tempfile::tempdir().unwrap(); - let vault = Vault::load(dir.path().join("secrets.json")).unwrap(); - // No env credentials configured: SigV4 must still resolve. - let resolver = test_resolver(vault, Arc::new(|_| None)); - let catalog = catalog_with( - r#" -[providers.bedrock] -adapter = "bedrock" -enabled = true -base_url = "https://bedrock-runtime.eu-west-1.amazonaws.com" - -[providers.bedrock.auth] -credentials = ["aws_sigv4"] -"#, - ); - - let resolved = resolver - .resolve( - ProviderId::from("bedrock"), - CredentialUsage::ApiRequest, - &catalog, - ) - .await - .unwrap(); - - let ResolvedCredential::Api(api) = resolved; - assert_eq!(api.provider, ProviderId::from("bedrock")); - assert_eq!(api.auth_header, Some(ApiKeyHeader::AwsSigv4)); - } - - #[tokio::test] - async fn resolve_returns_not_configured_for_missing_provider() { - let dir = tempfile::tempdir().unwrap(); - let vault = Vault::load(dir.path().join("secrets.json")).unwrap(); - let resolver = test_resolver(vault, Arc::new(|_| None)); - let catalog = default_catalog(); - - let err = resolver - .resolve( - ProviderId::anthropic(), - CredentialUsage::ApiRequest, - &catalog, - ) - .await - .unwrap_err(); - - assert!(matches!( - err, - ResolveError::NotConfigured(provider) if provider == ProviderId::anthropic() - )); - } - - #[tokio::test] - async fn anthropic_api_credentials_use_x_api_key_header() { - let dir = tempfile::tempdir().unwrap(); - let mut vault = Vault::load(dir.path().join("secrets.json")).unwrap(); - vault_set_token(&mut vault, "ANTHROPIC_API_KEY", "anthropic-key").unwrap(); - let resolver = test_resolver(vault, Arc::new(|_| None)); - let catalog = default_catalog(); - - let resolved = resolver - .resolve( - ProviderId::anthropic(), - CredentialUsage::ApiRequest, - &catalog, - ) - .await - .unwrap(); - let ResolvedCredential::Api(api) = resolved; - - assert_eq!( - api.auth_header, - Some(ApiKeyHeader::Custom { - name: "x-api-key".to_string(), - value: "anthropic-key".to_string(), - }) - ); - } - - #[tokio::test] - async fn custom_openai_compatible_resolves_with_catalog_base_url_from_vault() { - let catalog = catalog_with( - r#" -[providers.acme] -display_name = "Acme" -adapter = "openai_compatible" -agent_profile = "openai" -base_url = "https://default.example.com/v1" - -[providers.acme.auth] -credentials = ["vault:acme"] - -[models."compat-model"] -provider = "acme" -display_name = "Compat Model" -family = "openai" -default = true - -[models."compat-model".limits] -context_window = 128000 - -[models."compat-model".features] -tools = true -vision = false -reasoning = false -"#, - ); - let dir = tempfile::tempdir().unwrap(); - let mut vault = Vault::load(dir.path().join("secrets.json")).unwrap(); - vault_set_token(&mut vault, "acme", "compat-key").unwrap(); - let resolver = test_resolver(vault, Arc::new(|_| None)); - let resolved = resolver - .resolve( - ProviderId::new("acme"), - CredentialUsage::ApiRequest, - &catalog, - ) - .await - .unwrap(); - - let ResolvedCredential::Api(api) = resolved; - assert_eq!( - api.auth_header, - Some(ApiKeyHeader::Bearer("compat-key".to_string())) - ); - assert_eq!( - api.base_url.as_deref(), - Some("https://default.example.com/v1") - ); - } - - #[tokio::test] - async fn with_env_lookup_overrides_vault_settings() { - let dir = tempfile::tempdir().unwrap(); - let mut vault = Vault::load(dir.path().join("secrets.json")).unwrap(); - vault_set_token(&mut vault, "OPENAI_API_KEY", "vault-key").unwrap(); - vault - .set( - "OPENAI_ORG_ID", - "vault-org", - fabro_vault::SecretType::Token, - None, - ) - .unwrap(); + async fn openai_api_key_attaches_org_and_project_from_env() { let resolver = test_resolver( - vault, + empty_vault(), Arc::new(|name| match name { - "OPENAI_API_KEY" => Some("env-key".to_string()), - "OPENAI_ORG_ID" => Some("env-org".to_string()), + "OPENAI_API_KEY" => Some("key".to_string()), + "OPENAI_ORG_ID" => Some("org".to_string()), + "OPENAI_PROJECT_ID" => Some("proj".to_string()), _ => None, }), ); - let catalog = default_catalog(); - - let resolved = resolver - .resolve(ProviderId::openai(), CredentialUsage::ApiRequest, &catalog) + let catalog = test_catalog(); + let credentials = resolver + .resolve(catalog.provider("openai").unwrap()) .await .unwrap(); - let ResolvedCredential::Api(api) = resolved; - - assert_eq!(api.org_id.as_deref(), Some("env-org")); - } - - #[tokio::test] - async fn configured_providers_returns_vault_backed_provider() { - let dir = tempfile::tempdir().unwrap(); - let mut vault = Vault::load(dir.path().join("secrets.json")).unwrap(); - vault_set_token(&mut vault, "OPENAI_API_KEY", "vault-key").unwrap(); - let resolver = test_resolver(vault, Arc::new(|_| None)); - let vault = resolver.vault.read().await; - let catalog = default_catalog(); - - assert_eq!(resolver.configured_providers(&vault, &catalog), vec![ - ProviderId::openai() - ]); - } - - #[tokio::test] - async fn resolve_uses_custom_vault_backed_provider() { - let catalog = catalog_with( - r#" -[providers.acme] -display_name = "Acme" -adapter = "openai_compatible" -agent_profile = "openai" -base_url = "https://api.acme.test/v1" - -[providers.acme.auth] -credentials = ["vault:acme"] - -[models."acme-large"] -provider = "acme" -display_name = "Acme Large" -family = "acme" -default = true - -[models."acme-large".limits] -context_window = 128000 - -[models."acme-large".features] -tools = true -vision = false -reasoning = false -"#, - ); - let dir = tempfile::tempdir().unwrap(); - let mut vault = Vault::load(dir.path().join("secrets.json")).unwrap(); - vault_set_token(&mut vault, "acme", "acme-key").unwrap(); - let resolver = test_resolver(vault, Arc::new(|_| None)); - - let resolved = resolver - .resolve( - ProviderId::new("acme"), - CredentialUsage::ApiRequest, - &catalog, - ) - .await - .unwrap(); - - let ResolvedCredential::Api(api) = resolved; - assert_eq!(api.provider, ProviderId::new("acme")); assert_eq!( - api.auth_header, - Some(ApiKeyHeader::Bearer("acme-key".to_string())) + header_value(&credentials, "OpenAI-Organization"), + Some("org") ); - assert_eq!(api.base_url.as_deref(), Some("https://api.acme.test/v1")); + assert_eq!(header_value(&credentials, "OpenAI-Project"), Some("proj")); } #[tokio::test] - async fn configured_providers_returns_env_backed_provider() { - let dir = tempfile::tempdir().unwrap(); - let vault = Vault::load(dir.path().join("secrets.json")).unwrap(); + async fn bedrock_falls_back_to_the_aws_default_chain() { + let resolver = test_resolver(empty_vault(), Arc::new(|_| None)); + let catalog = test_catalog(); + let credentials = resolver + .resolve(catalog.provider("bedrock").unwrap()) + .await + .unwrap(); + assert!(matches!(credentials, Credentials::AwsDefaultChain { .. })); + let resolver = test_resolver( - vault, - Arc::new(|name| (name == "OPENAI_API_KEY").then(|| "env-key".to_string())), + empty_vault(), + Arc::new(|name| (name == "BEDROCK_API_KEY").then(|| "bearer".to_string())), ); - let vault = resolver.vault.read().await; - let catalog = default_catalog(); - - assert_eq!(resolver.configured_providers(&vault, &catalog), vec![ - ProviderId::openai() - ]); + let credentials = resolver + .resolve(catalog.provider("bedrock").unwrap()) + .await + .unwrap(); + assert!(matches!(credentials, Credentials::BedrockBearer(_))); } #[tokio::test] - async fn vault_source_resolves_secret_header_token() { - let catalog = portkey_catalog(r#"x-team-secret = "{{ secrets.gateway_team_secret }}""#); - let dir = tempfile::tempdir().unwrap(); - let mut vault = Vault::load(dir.path().join("secrets.json")).unwrap(); - vault_set_token(&mut vault, "gateway_team_secret", "s3cr3t").unwrap(); - let resolver = test_resolver(vault, Arc::new(|_| None)); - - let resolved = resolver - .resolve( - ProviderId::new("portkey"), - CredentialUsage::ApiRequest, - &catalog, - ) + async fn missing_provider_material_is_not_configured() { + let resolver = test_resolver(empty_vault(), Arc::new(|_| None)); + let catalog = test_catalog(); + let err = resolver + .resolve(catalog.provider("anthropic").unwrap()) .await - .unwrap(); - - let ResolvedCredential::Api(api) = resolved; - assert_eq!( - api.extra_headers.get("x-team-secret"), - Some(&"s3cr3t".to_string()) - ); + .unwrap_err(); + assert!(matches!( + err, + ResolveError::NotConfigured(provider) if provider.as_str() == "anthropic" + )); } #[tokio::test] async fn modal_resolves_both_vault_proxy_headers_without_authorization() { - let catalog = modal_catalog(); - let dir = tempfile::tempdir().unwrap(); - let mut vault = Vault::load(dir.path().join("secrets.json")).unwrap(); + let mut vault = empty_vault(); vault_set_token(&mut vault, "MODAL_TOKEN_ID", "wk-test").unwrap(); vault_set_token(&mut vault, "MODAL_TOKEN_SECRET", "ws-test").unwrap(); let resolver = test_resolver(vault, Arc::new(|_| None)); - let modal = ProviderId::new("modal"); - + let catalog = test_catalog(); + let modal = catalog.provider("modal").unwrap(); { let vault = resolver.vault.read().await; - assert!( - resolver - .configured_providers(&vault, &catalog) - .contains(&modal) - ); + assert!(resolver.has_credential_material(&vault, modal)); } - - let resolved = resolver - .resolve(modal.clone(), CredentialUsage::ApiRequest, &catalog) - .await - .unwrap(); - let ResolvedCredential::Api(api) = resolved; - - assert!(api.auth_header.is_none()); - assert_eq!( - api.extra_headers, - HashMap::from([ - ("Modal-Key".to_string(), "wk-test".to_string()), - ("Modal-Secret".to_string(), "ws-test".to_string()), - ]) - ); - assert_eq!( - api.base_url.as_deref(), - Some("https://example--kimi-k3.modal.run/v1") - ); + let credentials = resolver.resolve(modal).await.unwrap(); + match &credentials { + Credentials::Http(http) => assert!(matches!(http.auth, HttpAuthentication::None)), + other => panic!("expected header-only credentials, got {other:?}"), + } + assert_eq!(header_value(&credentials, "Modal-Key"), Some("wk-test")); + assert_eq!(header_value(&credentials, "Modal-Secret"), Some("ws-test")); } #[tokio::test] async fn modal_is_not_configured_with_only_one_vault_proxy_token() { - let catalog = modal_catalog(); - let dir = tempfile::tempdir().unwrap(); - let mut vault = Vault::load(dir.path().join("secrets.json")).unwrap(); - vault_set_token(&mut vault, "MODAL_TOKEN_ID", "wk-present").unwrap(); + let mut vault = empty_vault(); + vault_set_token(&mut vault, "MODAL_TOKEN_ID", "wk-test").unwrap(); let resolver = test_resolver(vault, Arc::new(|_| None)); - let modal = ProviderId::new("modal"); - - { - let vault = resolver.vault.read().await; - assert!( - !resolver - .configured_providers(&vault, &catalog) - .contains(&modal) - ); - } - - let err = resolver - .resolve(modal.clone(), CredentialUsage::ApiRequest, &catalog) - .await - .unwrap_err(); - - assert!(matches!( - err, - ResolveError::Interpolation { ref provider, .. } if provider == &modal - )); - let message = err.to_string(); - assert!(message.contains("MODAL_TOKEN_SECRET")); - assert!(!message.contains("wk-present")); + let catalog = test_catalog(); + let modal = catalog.provider("modal").unwrap(); + let err = resolver.resolve(modal).await.unwrap_err(); + assert!(matches!(err, ResolveError::Interpolation { .. }), "{err}"); + assert!(!err.to_string().contains("wk-test")); } #[tokio::test] - async fn resolve_multi_segment_header_token() { - let catalog = portkey_catalog(r#"authorization = "Bearer {{ secrets.TOKEN }}""#); - let dir = tempfile::tempdir().unwrap(); - let mut vault = Vault::load(dir.path().join("secrets.json")).unwrap(); - vault_set_token(&mut vault, "TOKEN", "gateway-token").unwrap(); - let resolver = test_resolver(vault, Arc::new(|_| None)); - - let resolved = resolver - .resolve( - ProviderId::new("portkey"), - CredentialUsage::ApiRequest, - &catalog, - ) - .await - .unwrap(); - - let ResolvedCredential::Api(api) = resolved; - assert_eq!( - api.extra_headers.get("authorization"), - Some(&"Bearer gateway-token".to_string()) + async fn configured_providers_reads_vault_and_env_without_refreshing() { + let mut vault = empty_vault(); + vault_set_token(&mut vault, "OPENAI_API_KEY", "vault-key").unwrap(); + let resolver = test_resolver( + vault, + Arc::new(|name| (name == "ANTHROPIC_API_KEY").then(|| "env".to_string())), ); + let catalog = test_catalog(); + let vault = resolver.vault.read().await; + let configured = resolver.configured_providers(&vault, &catalog); + assert!(configured.contains(&ProviderId::new("openai"))); + assert!(configured.contains(&ProviderId::new("anthropic"))); + // Bedrock always resolves through the AWS chain but ships disabled. + assert!(!configured.contains(&ProviderId::new("bedrock"))); } #[tokio::test] - async fn missing_secret_header_fails_without_echoing_value() { - let catalog = portkey_catalog(r#"x-team-secret = "{{ secrets.MISSING }}""#); - let dir = tempfile::tempdir().unwrap(); - let mut vault = Vault::load(dir.path().join("secrets.json")).unwrap(); - vault_set_token(&mut vault, "OTHER_SECRET", "should-not-leak").unwrap(); - let resolver = test_resolver(vault, Arc::new(|_| None)); - - let err = resolver - .resolve( - ProviderId::new("portkey"), - CredentialUsage::ApiRequest, - &catalog, - ) - .await - .unwrap_err(); - - assert!(matches!( - err, - ResolveError::Interpolation { ref provider, .. } - if provider == &ProviderId::new("portkey") - )); - let source = err - .source() - .expect("interpolation errors should preserve the source error"); - assert!(source.to_string().contains("MISSING")); - let message = err.to_string(); - assert!(message.contains("MISSING")); - assert!(!message.contains("should-not-leak")); - } - - #[tokio::test] - async fn header_with_file_or_oauth_vault_entry_fails_closed() { - let catalog = portkey_catalog(r#"x-team-secret = "{{ secrets.gateway_team_secret }}""#); - let dir = tempfile::tempdir().unwrap(); - let mut vault = Vault::load(dir.path().join("secrets.json")).unwrap(); - vault_set_oauth( - &mut vault, - "gateway_team_secret", - &oauth_credential( - "https://auth.openai.com/oauth/token".to_string(), - Utc::now() + Duration::hours(1), - ), - ) - .unwrap(); - let resolver = test_resolver(vault, Arc::new(|_| None)); - - let err = resolver - .resolve( - ProviderId::new("portkey"), - CredentialUsage::ApiRequest, - &catalog, - ) - .await - .unwrap_err(); - - assert!(matches!( - err, - ResolveError::Interpolation { ref provider, .. } - if provider == &ProviderId::new("portkey") - )); - let message = err.to_string(); - assert!(message.contains("gateway_team_secret")); - assert!(!message.contains("expired-access")); - assert!(!message.contains("refresh-token")); - } - - #[tokio::test] - async fn resolve_refreshes_expired_oauth_credentials_and_persists_them() { + async fn refreshes_expired_oauth_credentials_and_persists_them() { let server = MockServer::start_async().await; let refresh_mock = server .mock_async(|when, then| { when.method(POST) .path("/oauth/token") - .header("content-type", "application/x-www-form-urlencoded") .form_urlencoded_tuple("grant_type", "refresh_token") .form_urlencoded_tuple("client_id", "test-client") .form_urlencoded_tuple("refresh_token", "refresh-token"); @@ -1173,8 +782,7 @@ reasoning = false }) .await; - let dir = tempfile::tempdir().unwrap(); - let mut vault = Vault::load(dir.path().join("secrets.json")).unwrap(); + let mut vault = empty_vault(); vault_set_oauth( &mut vault, crate::OPENAI_CODEX_VAULT_SECRET_NAME, @@ -1186,19 +794,13 @@ reasoning = false .unwrap(); let vault = Arc::new(AsyncRwLock::new(vault)); let resolver = CredentialResolver::with_env_lookup(Arc::clone(&vault), Arc::new(|_| None)); - let catalog = default_catalog(); + let catalog = test_catalog(); - let resolved = resolver - .resolve(ProviderId::openai(), CredentialUsage::ApiRequest, &catalog) + let credentials = resolver + .resolve(catalog.provider("openai-codex").unwrap()) .await .unwrap(); - let ResolvedCredential::Api(api) = resolved; - - assert_eq!( - api.auth_header, - Some(ApiKeyHeader::Bearer("new-access".to_string())) - ); - assert!(api.codex_mode); + assert_eq!(bearer_secret(&credentials), "new-access"); let stored = { let vault = vault.read().await; @@ -1213,9 +815,8 @@ reasoning = false } #[tokio::test] - async fn resolve_returns_refresh_token_missing_when_expired_oauth_has_no_refresh_token() { - let dir = tempfile::tempdir().unwrap(); - let mut vault = Vault::load(dir.path().join("secrets.json")).unwrap(); + async fn expired_oauth_without_refresh_token_requires_reauthentication() { + let mut vault = empty_vault(); let mut credential = oauth_credential( "https://auth.openai.com/oauth/token".to_string(), Utc::now() - Duration::minutes(1), @@ -1228,47 +829,43 @@ reasoning = false ) .unwrap(); let resolver = test_resolver(vault, Arc::new(|_| None)); - let catalog = default_catalog(); - + let catalog = test_catalog(); let err = resolver - .resolve(ProviderId::openai(), CredentialUsage::ApiRequest, &catalog) + .resolve(catalog.provider("openai-codex").unwrap()) .await .unwrap_err(); - - assert!(matches!( - err, - ResolveError::RefreshTokenMissing(provider) if provider == ProviderId::openai() - )); - } - - #[test] - fn auth_issue_message_formats_refresh_token_missing() { - let message = auth_issue_message( - &ProviderId::openai(), - &ResolveError::RefreshTokenMissing(ProviderId::openai()), - ); - + assert!(matches!(err, ResolveError::RefreshTokenMissing(_))); assert_eq!( - message, - "openai requires re-authentication: refresh token missing" + auth_issue_message(&ProviderId::new("openai-codex"), &err), + "openai-codex requires re-authentication: refresh token missing" ); } #[test] - fn api_credential_debug_redacts_secret_material() { - let credential = ApiCredential { - provider: ProviderId::openai(), - auth_header: Some(ApiKeyHeader::Bearer("sk-test".to_string())), - extra_headers: HashMap::new(), - base_url: None, - codex_mode: false, - org_id: None, - project_id: None, - }; - - let debug = format!("{credential:?}"); + fn api_key_credentials_follow_the_provider_scheme() { + let catalog = test_catalog(); + let vault = empty_vault(); + let openai = credentials_for_api_key( + catalog.provider("openai").unwrap(), + "sk-test".to_string(), + &vault, + ) + .unwrap(); + assert_eq!(bearer_secret(&openai), "sk-test"); + let modal = credentials_for_api_key( + catalog.provider("modal").unwrap(), + "sk-test".to_string(), + &vault, + ); + assert!(modal.is_err(), "modal has no single-key scheme"); + assert!(!accepts_api_key(catalog.provider("modal").unwrap())); + assert!(accepts_api_key(catalog.provider("openai").unwrap())); + assert!(!accepts_api_key(catalog.provider("ollama").unwrap())); + } + #[test] + fn resolved_secret_debug_redacts_material() { + let debug = format!("{:?}", ResolvedSecret::ApiKey("sk-test".to_string())); assert!(!debug.contains("sk-test")); - assert!(debug.contains("REDACTED")); } } diff --git a/lib/foundation/fabro-auth/src/sql_vault_source.rs b/lib/foundation/fabro-auth/src/sql_vault_source.rs index ca581d5b6..6f0c0803b 100644 --- a/lib/foundation/fabro-auth/src/sql_vault_source.rs +++ b/lib/foundation/fabro-auth/src/sql_vault_source.rs @@ -1,15 +1,21 @@ use std::sync::Arc; use async_trait::async_trait; -use fabro_model::{Catalog, ProviderId}; use fabro_types::SecretType; use fabro_vault::{SecretSnapshot, SecretStore, SecretStoreError, Vault}; +use lithos_llm::catalog::{Catalog, CatalogProvider, ProviderId}; +use lithos_llm::credentials::Credentials; use tokio::sync::RwLock; use tracing::error; -use crate::credential_source::{CredentialSource, ResolvedCredentials}; -use crate::{EnvLookup, VaultCredentialSource}; +use crate::credential_source::CredentialSource; +use crate::{EnvLookup, ResolveError, VaultCredentialSource}; +/// Credentials backed by the SQL secret store. +/// +/// Every lookup snapshots the store, resolves against the snapshot, and +/// writes refreshed OAuth tokens back with a revision check so two concurrent +/// refreshes cannot clobber each other. #[derive(Clone)] pub struct SqlVaultCredentialSource { store: Arc, @@ -83,6 +89,13 @@ impl SqlVaultCredentialSource { } Ok(true) } + + fn store_error(provider: &ProviderId, err: SecretStoreError) -> ResolveError { + ResolveError::RefreshFailed { + provider: provider.clone(), + source: anyhow::Error::new(err), + } + } } impl std::fmt::Debug for SqlVaultCredentialSource { @@ -94,9 +107,13 @@ impl std::fmt::Debug for SqlVaultCredentialSource { #[async_trait] impl CredentialSource for SqlVaultCredentialSource { - async fn resolve(&self, catalog: &Catalog) -> anyhow::Result { + async fn credentials(&self, provider: &CatalogProvider) -> Result { for _ in 0..2 { - let before = self.store.snapshot().await?; + let before = self + .store + .snapshot() + .await + .map_err(|err| Self::store_error(provider.id(), err))?; let has_oauth = before .entries() .values() @@ -104,16 +121,23 @@ impl CredentialSource for SqlVaultCredentialSource { if !has_oauth { // Only OAuth resolution can write back (token refresh); with no // OAuth secrets, skip the snapshot clones and CAS machinery. - return self.source_for_snapshot(before).resolve(catalog).await; + return self.source_for_snapshot(before).credentials(provider).await; } let source = self.source_for_snapshot(before.clone()); - let resolved = source.resolve(catalog).await?; + let credentials = source.credentials(provider).await?; let after = source.snapshot().await; - if self.persist_oauth_refreshes(&before, &after).await? { - return Ok(resolved); + if self + .persist_oauth_refreshes(&before, &after) + .await + .map_err(|err| Self::store_error(provider.id(), err))? + { + return Ok(credentials); } } - anyhow::bail!("OAuth credential changed concurrently during refresh") + Err(ResolveError::RefreshFailed { + provider: provider.id().clone(), + source: anyhow::anyhow!("OAuth credential changed concurrently during refresh"), + }) } async fn configured_providers(&self, catalog: &Catalog) -> Vec { diff --git a/lib/foundation/fabro-auth/src/strategies/api_key.rs b/lib/foundation/fabro-auth/src/strategies/api_key.rs index 0d594110a..225f1db34 100644 --- a/lib/foundation/fabro-auth/src/strategies/api_key.rs +++ b/lib/foundation/fabro-auth/src/strategies/api_key.rs @@ -1,6 +1,6 @@ use async_trait::async_trait; -use fabro_model::catalog::CatalogProvider; -use fabro_model::{CredentialRef, ProviderId}; +use fabro_types::catalog_policy; +use lithos_llm::catalog::{CatalogProvider, ProviderId}; use crate::context::{AuthContextRequest, AuthContextResponse}; use crate::strategy::{AuthStrategy, LoginResult}; @@ -15,24 +15,11 @@ pub struct ApiKeyStrategy { impl ApiKeyStrategy { #[must_use] pub fn new(provider: &CatalogProvider) -> Self { - let env_var_names = provider - .auth - .as_ref() - .map(|auth| { - auth.credentials - .iter() - .filter_map(|credential_ref| match credential_ref { - CredentialRef::Env(name) => Some(name.clone()), - CredentialRef::Vault(_) | CredentialRef::AwsSigv4 => None, - }) - .collect() - }) - .unwrap_or_default(); Self { - provider_id: provider.id.clone(), - display_name: provider.display_name.clone(), - env_var_names, - api_key_url: provider.api_key_url.clone(), + provider_id: provider.id().clone(), + display_name: provider.display_name().to_string(), + env_var_names: crate::env_var_names(provider), + api_key_url: catalog_policy::provider_policy(provider).api_key_url, } } } diff --git a/lib/foundation/fabro-auth/src/strategies/codex_device.rs b/lib/foundation/fabro-auth/src/strategies/codex_device.rs index 94a3399c8..47681b2e1 100644 --- a/lib/foundation/fabro-auth/src/strategies/codex_device.rs +++ b/lib/foundation/fabro-auth/src/strategies/codex_device.rs @@ -5,6 +5,7 @@ use base64::Engine; use base64::engine::general_purpose::URL_SAFE_NO_PAD; use chrono::{DateTime, Utc}; use fabro_http::HttpClient; +use fabro_types::provider_ids; use serde::{Deserialize, Serialize}; use serde_json::json; use tokio::time::sleep; @@ -298,7 +299,7 @@ impl AuthStrategy for CodexDeviceStrategy { .map_err(anyhow::Error::msg)?; Ok(LoginResult::OAuth { - provider: fabro_model::ProviderId::openai(), + provider: provider_ids::openai(), credential: OAuthCredential { tokens: OAuthTokens { access_token: token_response.access_token, diff --git a/lib/foundation/fabro-auth/src/strategy.rs b/lib/foundation/fabro-auth/src/strategy.rs index 3b0603248..60945227c 100644 --- a/lib/foundation/fabro-auth/src/strategy.rs +++ b/lib/foundation/fabro-auth/src/strategy.rs @@ -1,5 +1,6 @@ use async_trait::async_trait; -use fabro_model::{Catalog, ProviderId}; +use fabro_types::provider_ids; +use lithos_llm::catalog::{Catalog, ProviderId}; use crate::context::{AuthContextRequest, AuthContextResponse}; use crate::credential::{OAuthConfig, OAuthCredential}; @@ -60,7 +61,7 @@ pub fn strategy_for( match method { AuthMethod::ApiKey => { let provider = catalog - .provider(provider_id) + .provider(provider_id.as_str()) .expect("API key auth requires a catalog provider"); Box::new(ApiKeyStrategy::new(provider)) } @@ -73,7 +74,7 @@ pub fn strategy_for( // forgets the constraint. assert_eq!( provider_id.as_str(), - ProviderId::OPENAI, + provider_ids::OPENAI, "CodexDevice auth is only constructed by CLI code for the \ OpenAI provider; all existing call sites enforce this pairing: \ got provider_id={provider_id}" @@ -87,6 +88,7 @@ pub fn strategy_for( mod tests { use super::*; use crate::context::AuthContextRequest; + use crate::test_support::test_catalog; #[test] fn codex_oauth_config_has_expected_defaults() { @@ -100,12 +102,12 @@ mod tests { #[tokio::test] async fn api_key_strategy_uses_provider_env_names() { - let catalog = Catalog::builtin(); - let provider = catalog.provider(&ProviderId::anthropic()).unwrap(); + let catalog = test_catalog(); + let provider = catalog.provider("anthropic").unwrap(); let mut strategy = ApiKeyStrategy::new(provider); let request = strategy.init().await.unwrap(); assert_eq!(request, AuthContextRequest::ApiKey { - provider_id: ProviderId::anthropic(), + provider_id: ProviderId::new("anthropic"), display_name: "Anthropic".to_string(), env_var_names: vec!["ANTHROPIC_API_KEY".to_string()], api_key_url: Some("https://console.anthropic.com/settings/keys".to_string()), diff --git a/lib/foundation/fabro-auth/src/test_support.rs b/lib/foundation/fabro-auth/src/test_support.rs index d3bb14cde..b1cc1271c 100644 --- a/lib/foundation/fabro-auth/src/test_support.rs +++ b/lib/foundation/fabro-auth/src/test_support.rs @@ -1,4 +1,4 @@ -//! Test-only credential sources. +//! Test-only credential sources and catalogs. //! //! Feature-gated so they never link into production builds. Production code //! resolves credentials through [`VaultCredentialSource`] over a real vault; @@ -8,11 +8,28 @@ use std::collections::HashMap; use std::sync::Arc; use fabro_vault::Vault; +use lithos_llm::catalog::Catalog; use tokio::sync::RwLock as AsyncRwLock; use crate::credential_source::CredentialSource; use crate::vault_source::VaultCredentialSource; +/// Fabro's policy layer, checked in under `fabro-llm`. Tests in this crate +/// need the built-in catalog with `metadata.fabro.credentials` attached. +pub const FABRO_POLICY_TOML: &str = + include_str!("../../../components/fabro-llm/catalog/fabro-policy.toml"); + +/// The lithos built-in catalog with Fabro's policy layer applied. +#[must_use] +pub fn test_catalog() -> Catalog { + Catalog::builder() + .with_builtin() + .toml_layer("fabro-policy.toml", FABRO_POLICY_TOML) + .expect("fabro policy layer should parse") + .build() + .expect("built-in catalog with fabro policy should build") +} + /// A detached in-memory vault holding no secrets. #[must_use] pub fn empty_vault() -> Arc> { diff --git a/lib/foundation/fabro-auth/src/vault_source.rs b/lib/foundation/fabro-auth/src/vault_source.rs index aeba384a4..cbf25a66b 100644 --- a/lib/foundation/fabro-auth/src/vault_source.rs +++ b/lib/foundation/fabro-auth/src/vault_source.rs @@ -1,13 +1,15 @@ use std::sync::Arc; use async_trait::async_trait; -use fabro_model::{Catalog, ProviderId}; use fabro_vault::Vault; +use lithos_llm::catalog::{Catalog, CatalogProvider, ProviderId}; +use lithos_llm::credentials::Credentials; use tokio::sync::RwLock as AsyncRwLock; -use crate::credential_source::{CredentialSource, ResolvedCredentials}; -use crate::{CredentialResolver, CredentialUsage, EnvLookup, ResolveError, ResolvedCredential}; +use crate::credential_source::CredentialSource; +use crate::{CredentialResolver, EnvLookup, ResolveError}; +/// Credentials backed by an in-memory [`Vault`] plus an environment lookup. #[derive(Clone)] pub struct VaultCredentialSource { vault: Arc>, @@ -50,26 +52,8 @@ impl std::fmt::Debug for VaultCredentialSource { #[async_trait] impl CredentialSource for VaultCredentialSource { - async fn resolve(&self, catalog: &Catalog) -> anyhow::Result { - let mut credentials = Vec::new(); - let mut auth_issues = Vec::new(); - - for provider in catalog.providers() { - match self - .resolver - .resolve(provider.id.clone(), CredentialUsage::ApiRequest, catalog) - .await - { - Ok(ResolvedCredential::Api(credential)) => credentials.push(credential), - Err(ResolveError::NotConfigured(_)) if provider.auth.is_some() => {} - Err(err) => auth_issues.push((provider.id.clone(), err)), - } - } - - Ok(ResolvedCredentials { - credentials, - auth_issues, - }) + async fn credentials(&self, provider: &CatalogProvider) -> Result { + self.resolver.resolve(provider).await } async fn configured_providers(&self, catalog: &Catalog) -> Vec { @@ -80,15 +64,17 @@ impl CredentialSource for VaultCredentialSource { #[cfg(test)] mod tests { + use std::collections::HashMap; use std::sync::Arc; use chrono::{Duration, Utc}; - use fabro_model::{Catalog, ProviderId}; use fabro_vault::Vault; + use lithos_llm::catalog::ProviderId; use tokio::sync::RwLock as AsyncRwLock; use super::VaultCredentialSource; use crate::credential::{OAuthConfig, OAuthCredential, OAuthTokens}; + use crate::test_support::test_catalog; use crate::vault_ext::{vault_set_oauth, vault_set_token}; use crate::{CredentialSource, ResolveError}; @@ -111,14 +97,9 @@ mod tests { } } - fn default_catalog() -> Catalog { - Catalog::from_builtin().unwrap() - } - #[tokio::test] - async fn resolve_returns_credentials_and_auth_issues() { - let dir = tempfile::tempdir().unwrap(); - let mut vault = Vault::load(dir.path().join("secrets.json")).unwrap(); + async fn resolve_all_separates_ready_providers_from_auth_issues() { + let mut vault = Vault::from_entries(HashMap::new()); vault_set_oauth( &mut vault, crate::OPENAI_CODEX_VAULT_SECRET_NAME, @@ -129,63 +110,50 @@ mod tests { let source = VaultCredentialSource::with_env_lookup(Arc::new(AsyncRwLock::new(vault)), |_| None); - let catalog = default_catalog(); + let catalog = test_catalog(); - let resolved = source.resolve(&catalog).await.unwrap(); + let resolved = source.resolve_all(&catalog).await; - assert_eq!(resolved.credentials.len(), 1); - assert_eq!(resolved.credentials[0].provider, ProviderId::anthropic()); + assert_eq!(resolved.ready, vec![ProviderId::new("anthropic")]); assert_eq!(resolved.auth_issues.len(), 1); assert!(matches!( &resolved.auth_issues[0].1, - ResolveError::RefreshFailed { - provider, - .. - } if provider == &ProviderId::openai() + ResolveError::RefreshFailed { provider, .. } if provider.as_str() == "openai-codex" )); } #[tokio::test] async fn configured_providers_reads_from_vault_without_refreshing() { - let dir = tempfile::tempdir().unwrap(); - let mut vault = Vault::load(dir.path().join("secrets.json")).unwrap(); + let mut vault = Vault::from_entries(HashMap::new()); vault_set_token(&mut vault, "OPENAI_API_KEY", "openai-key").unwrap(); vault_set_token(&mut vault, "ANTHROPIC_API_KEY", "anthropic-key").unwrap(); let source = VaultCredentialSource::with_env_lookup(Arc::new(AsyncRwLock::new(vault)), |_| None); - let catalog = default_catalog(); + let catalog = test_catalog(); assert_eq!(source.configured_providers(&catalog).await, vec![ - ProviderId::anthropic(), - ProviderId::openai() + ProviderId::new("anthropic"), + ProviderId::new("openai") ]); } #[tokio::test] async fn vault_only_ignores_env_lookup_values() { - let env_dir = tempfile::tempdir().unwrap(); - let vault_only_dir = tempfile::tempdir().unwrap(); - let catalog = default_catalog(); + let catalog = test_catalog(); let env_backed = VaultCredentialSource::with_env_lookup( - Arc::new(AsyncRwLock::new( - Vault::load(env_dir.path().join("secrets.json")).unwrap(), - )), + Arc::new(AsyncRwLock::new(Vault::from_entries(HashMap::new()))), |name| (name == "OPENAI_API_KEY").then(|| "env-openai-key".to_string()), ); assert_eq!(env_backed.configured_providers(&catalog).await, vec![ - ProviderId::openai() + ProviderId::new("openai") ]); let vault_only = VaultCredentialSource::vault_only(Arc::new(AsyncRwLock::new( - Vault::load(vault_only_dir.path().join("secrets.json")).unwrap(), + Vault::from_entries(HashMap::new()), ))); - - assert!( - vault_only.configured_providers(&catalog).await.is_empty(), - "vault_only must not resolve env-backed provider keys" - ); - let resolved = vault_only.resolve(&catalog).await.unwrap(); - assert!(resolved.credentials.is_empty()); + assert!(vault_only.configured_providers(&catalog).await.is_empty()); + let resolved = vault_only.resolve_all(&catalog).await; + assert!(resolved.ready.is_empty()); assert!(resolved.auth_issues.is_empty()); } } diff --git a/lib/foundation/fabro-config/Cargo.toml b/lib/foundation/fabro-config/Cargo.toml index 8c3df637e..059ed184b 100644 --- a/lib/foundation/fabro-config/Cargo.toml +++ b/lib/foundation/fabro-config/Cargo.toml @@ -21,7 +21,6 @@ anyhow.workspace = true clap = { workspace = true, optional = true } chrono.workspace = true fabro-macros = { path = "../fabro-macros" } -fabro-model = { path = "../fabro-model" } fabro-options-metadata.workspace = true fabro-proc = { path = "../fabro-proc" } fabro-static.workspace = true diff --git a/lib/foundation/fabro-config/src/builders.rs b/lib/foundation/fabro-config/src/builders.rs index 08eb4c70b..b0a949914 100644 --- a/lib/foundation/fabro-config/src/builders.rs +++ b/lib/foundation/fabro-config/src/builders.rs @@ -1,8 +1,7 @@ -use std::collections::{BTreeMap, HashMap}; +use std::collections::HashMap; use std::fmt; use std::path::Path; -use fabro_model::catalog as model_catalog; use fabro_types::settings::run::McpServerSettings; use fabro_types::settings::{RunNamespace, WorkflowNamespace}; use fabro_types::{ServerSettings, UserSettings, WorkflowSettings}; @@ -16,9 +15,8 @@ use crate::resolve::{ }; use crate::user::load_settings_config; use crate::{ - CliLayer, Combine, CostRates, EnvironmentLayer, Error, LlmLayer, LlmModelFeatures, - LlmModelLimits, MergeMap, ModelControls, ModelCostTable, ModelSettings, ProviderSettings, - Result, RunLayer, ServerLayer, SettingsLayer, run, + CliLayer, Combine, EnvironmentLayer, Error, LlmLayer, MergeMap, Result, RunLayer, ServerLayer, + SettingsLayer, run, }; #[derive(Debug, Clone, PartialEq, Eq)] @@ -231,7 +229,8 @@ pub struct ServerRuntimeSettings { pub manifest_run_defaults: RunLayer, pub manifest_environment_defaults: crate::MergeMap, pub manifest_run_settings: std::result::Result, - pub llm_catalog_settings: model_catalog::LlmCatalogSettings, + /// Operator catalog overlay, applied above the built-in and policy layers. + pub llm_overlay: LlmLayer, } pub fn load_server_runtime_settings( @@ -246,12 +245,12 @@ pub fn load_server_runtime_settings( resolve_server_runtime_settings(layer, run_overrides, server_overrides) } -pub fn load_llm_catalog_settings(path: Option<&Path>) -> Result { +pub fn load_llm_overlay(path: Option<&Path>) -> Result { let layer = match path { Some(path) => load_settings_path(path, SettingsSource::ActiveSettings)?, None => load_settings_config(None)?, }; - Ok(llm_catalog_settings_from_layer(&layer)) + Ok(llm_overlay_from_layer(&layer)) } #[cfg(test)] @@ -286,7 +285,7 @@ fn resolve_server_runtime_settings( let manifest_run_defaults = layer.run.clone().unwrap_or_default(); let manifest_environment_defaults = layer.environments.clone(); - let llm_catalog_settings = llm_catalog_settings_from_layer(&layer); + let llm_overlay = llm_overlay_from_layer(&layer); Ok(ServerRuntimeSettings { server_settings: ServerSettingsBuilder::from_layer(&layer)?, manifest_run_settings: RunSettingsBuilder::from_layer(&SettingsLayer { @@ -297,162 +296,13 @@ fn resolve_server_runtime_settings( .map_err(|err| SharedError::new(anyhow::Error::new(err))), manifest_run_defaults, manifest_environment_defaults, - llm_catalog_settings, + llm_overlay, }) } -fn llm_catalog_settings_from_layer(layer: &SettingsLayer) -> model_catalog::LlmCatalogSettings { +fn llm_overlay_from_layer(layer: &SettingsLayer) -> LlmLayer { let layer = layer.clone().combine(DEFAULTS_LAYER.clone()); - layer - .llm - .map(llm_layer_to_catalog_settings) - .unwrap_or_default() -} - -fn llm_layer_to_catalog_settings(llm: LlmLayer) -> model_catalog::LlmCatalogSettings { - model_catalog::LlmCatalogSettings { - providers: llm - .providers - .into_inner() - .into_iter() - .map(|(id, settings)| (id, provider_settings_to_catalog(settings))) - .collect(), - models: llm - .models - .into_inner() - .into_iter() - .map(|(id, settings)| (id, model_settings_to_catalog(settings))) - .collect(), - } -} - -fn provider_settings_to_catalog( - settings: ProviderSettings, -) -> model_catalog::ProviderCatalogSettings { - #[expect( - clippy::disallowed_methods, - reason = "collapse the authoring InterpString header values to their catalog source \ - strings; they are re-parsed and resolved at the credential boundary" - )] - let extra_headers = settings.extra_headers.map(|headers| { - headers - .into_iter() - .map(|(name, value)| (name, value.as_source())) - .collect() - }); - let models = settings - .models - .into_inner() - .into_iter() - .map(|(id, settings)| (id, model_settings_to_catalog(settings))) - .collect(); - model_catalog::ProviderCatalogSettings { - display_name: settings.display_name, - adapter: settings.adapter, - codec: settings.codec, - agent_profile: settings.agent_profile, - auth: settings.auth, - billing_policy: settings.billing_policy, - api_key_url: settings.api_key_url, - base_url: settings.base_url, - extra_headers, - priority: settings.priority, - enabled: settings.enabled, - aliases: settings.aliases, - models, - } -} - -fn model_settings_to_catalog(settings: ModelSettings) -> model_catalog::ModelCatalogSettings { - let ModelSettings { - provider, - api_id, - codec, - billing_policy, - agent_profile, - display_name, - family, - training, - knowledge_cutoff, - default, - small_default, - probe, - enabled, - aliases, - estimated_output_tps, - limits, - features, - controls, - costs, - } = settings; - model_catalog::ModelCatalogSettings { - provider, - api_id, - codec, - billing_policy, - agent_profile, - display_name, - family, - training, - knowledge_cutoff, - default, - small_default, - probe, - enabled, - aliases, - estimated_output_tps, - limits: limits.as_ref().map(model_limits_to_catalog), - features: features.as_ref().map(model_features_to_catalog), - controls: controls.map(model_controls_to_catalog), - costs: costs.as_ref().map(model_cost_table_to_catalog), - } -} - -fn model_limits_to_catalog(limits: &LlmModelLimits) -> model_catalog::SettingsModelLimits { - model_catalog::SettingsModelLimits { - context_window: limits.context_window, - max_output: limits.max_output, - } -} - -fn model_features_to_catalog(features: &LlmModelFeatures) -> model_catalog::SettingsModelFeatures { - model_catalog::SettingsModelFeatures { - tools: features.tools, - vision: features.vision, - reasoning: features.reasoning, - reasoning_by_default: features.reasoning_by_default, - reasoning_effort: features.reasoning_effort, - prompt_cache: features.prompt_cache, - cache_control_breakpoints: features.cache_control_breakpoints, - sampling_params: features.sampling_params, - } -} - -fn model_controls_to_catalog(controls: ModelControls) -> model_catalog::SettingsModelControls { - model_catalog::SettingsModelControls { - reasoning_effort: controls.reasoning_effort, - speed: controls.speed, - } -} - -fn model_cost_table_to_catalog(costs: &ModelCostTable) -> model_catalog::SettingsModelCostTable { - model_catalog::SettingsModelCostTable { - base: cost_rates_to_catalog(&costs.base), - speed: costs.speed.as_ref().map(|speed| { - speed - .iter() - .map(|(key, rates)| (key.clone(), cost_rates_to_catalog(rates))) - .collect::>() - }), - } -} - -fn cost_rates_to_catalog(rates: &CostRates) -> model_catalog::CostRates { - model_catalog::CostRates { - input_cost_per_mtok: rates.input_cost_per_mtok, - output_cost_per_mtok: rates.output_cost_per_mtok, - cache_input_cost_per_mtok: rates.cache_input_cost_per_mtok, - } + layer.llm.unwrap_or_default() } fn parse_settings_toml(source: &str, kind: SettingsSource) -> Result { @@ -829,7 +679,7 @@ provider = "docker" } #[test] - fn server_runtime_settings_preserves_llm_catalog_overrides() { + fn server_runtime_settings_preserves_llm_overlay() { let settings = server_runtime_settings_from_toml( r#" _version = 1 @@ -839,92 +689,31 @@ methods = ["dev-token"] [llm.providers.acme] display_name = "Acme" -adapter = "openai_compatible" +adapter = "openai-compatible" +codec = "openai-chat" base_url = "https://api.acme.test/v1" -agent_profile = "anthropic" +auth = { type = "bearer" } -[llm.providers.acme.auth] +[llm.providers.acme.metadata.fabro] +enabled = true credentials = ["env:ACME_API_KEY"] -[llm.models."acme-large"] -provider = "acme" +[llm.providers.acme.models."acme-large"] display_name = "Acme Large" -family = "acme" -default = true -agent_profile = "gemini" - -[llm.models."acme-large".limits] -context_window = 128000 - -[llm.models."acme-large".features] -tools = true -vision = false -reasoning = false +api_model = "acme-large" "#, None, None, ) .expect("server runtime settings should resolve"); - let catalog = - fabro_model::Catalog::from_builtin_with_overrides(&settings.llm_catalog_settings) - .expect("catalog overrides should build"); - + let overlay = settings.llm_overlay.0; + let acme = &overlay["providers"]["acme"]; + assert_eq!(acme["display_name"].as_str(), Some("Acme")); + assert_eq!(acme["metadata"]["fabro"]["enabled"].as_bool(), Some(true)); assert_eq!( - catalog - .get_on_provider(&fabro_model::ProviderId::new("acme"), "acme-large") - .map(|model| model.provider.clone()), - Some(fabro_model::ProviderId::new("acme")) - ); - assert_eq!( - catalog - .effective_agent_profile(&fabro_model::ProviderId::new("acme"), Some("acme-large")), - Some(fabro_model::AgentProfileKind::Gemini) - ); - } - - #[test] - fn server_runtime_settings_preserves_extra_header_sources() { - let settings = server_runtime_settings_from_toml( - r#" -_version = 1 - -[server.auth] -methods = ["dev-token"] - -[llm.providers.acme] -display_name = "Acme" -adapter = "openai_compatible" -base_url = "https://api.acme.test/v1" - -[llm.providers.acme.extra_headers] -x-title = "My App" -x-api-key = "{{ env.ACME_GATEWAY_API_KEY }}" -x-team-secret = "Bearer {{ secrets.ACME_GATEWAY_TOKEN }}" -"#, - None, - None, - ) - .expect("server runtime settings should resolve"); - - let provider = settings - .llm_catalog_settings - .providers - .get("acme") - .expect("provider settings should be present"); - let headers = provider - .extra_headers - .as_ref() - .expect("extra header settings should be present"); - - assert_eq!(headers.get("x-title").map(String::as_str), Some("My App")); - assert_eq!( - headers.get("x-api-key").map(String::as_str), - Some("{{ env.ACME_GATEWAY_API_KEY }}") - ); - assert_eq!( - headers.get("x-team-secret").map(String::as_str), - Some("Bearer {{ secrets.ACME_GATEWAY_TOKEN }}") + acme["models"]["acme-large"]["api_model"].as_str(), + Some("acme-large") ); } } diff --git a/lib/foundation/fabro-config/src/layers/combine.rs b/lib/foundation/fabro-config/src/layers/combine.rs index aec393777..d058c2977 100644 --- a/lib/foundation/fabro-config/src/layers/combine.rs +++ b/lib/foundation/fabro-config/src/layers/combine.rs @@ -1,6 +1,5 @@ -use std::collections::{BTreeMap, HashMap}; +use std::collections::HashMap; -use fabro_model::{AgentProfileKind, BillingPolicy, CodecKind, ProviderAuthConfig}; use fabro_types::PermissionLevel; use fabro_types::settings::cli::{CliAuthStrategy, OutputFormat, OutputVerbosity}; use fabro_types::settings::run::{ @@ -15,7 +14,6 @@ use fabro_types::settings::{Duration, InterpString, Size}; use super::LogFilter; use super::cli::{CliAuthLayer, CliLoggingLayer, CliTargetLayer}; use super::environment::EnvironmentDockerfileLayer; -use super::llm::{CostRates, CredentialRef, ReasoningEffortFeature}; use super::run::{ HookAgentMarker, HookEntry, HookTlsMode, InterviewProviderLayer, ModelRefOrSplice, NotificationProviderLayer, RunArtifactsLayer, RunCheckpointLayer, RunGoalLayer, @@ -85,11 +83,6 @@ impl_combine_or_option!( ServerAuthMethod, WebhookStrategy, LogFilter, - AgentProfileKind, - BillingPolicy, - CodecKind, - ProviderAuthConfig, - ReasoningEffortFeature, ); impl Combine for Option> { @@ -98,24 +91,12 @@ impl Combine for Option> { } } -impl Combine for Option> { - fn combine(self, other: Self) -> Self { - self.or(other) - } -} - impl Combine for Option> { fn combine(self, other: Self) -> Self { self.or(other) } } -impl Combine for Option> { - fn combine(self, other: Self) -> Self { - self.or(other) - } -} - impl Combine for Option> { fn combine(self, other: Self) -> Self { self.or(other) diff --git a/lib/foundation/fabro-config/src/layers/llm.rs b/lib/foundation/fabro-config/src/layers/llm.rs index 298968a62..17df191ab 100644 --- a/lib/foundation/fabro-config/src/layers/llm.rs +++ b/lib/foundation/fabro-config/src/layers/llm.rs @@ -1,982 +1,130 @@ //! `[llm]` settings layer. //! -//! Holds the trusted, mergeable LLM provider/model catalog data: +//! Operator model catalog overrides. The table uses the lithos-llm catalog +//! schema verbatim, minus `schema_version`, and is applied as an overlay layer +//! on top of the lithos built-in catalog and Fabro's policy layer: //! //! ```toml //! [llm.providers.moonshot] -//! display_name = "Moonshot AI" -//! adapter = "openai_compatible" -//! base_url = "https://api.moonshot.ai/v1" //! priority = 60 +//! +//! [llm.providers.moonshot.metadata.fabro] //! enabled = true -//! aliases = ["moonshot-ai"] +//! credentials = ["env:MOONSHOT_API_KEY", "vault:MOONSHOT_API_KEY"] //! -//! [llm.providers.moonshot.auth] -//! credentials = [ -//! "env:MOONSHOT_API_KEY", -//! "env:KIMI_API_KEY", -//! "vault:MOONSHOT_API_KEY", -//! "vault:KIMI_API_KEY", -//! ] -//! -//! [llm.providers.moonshot.models."kimi-k2.5"] -//! ... +//! [llm.providers.moonshot.models."kimi-k2.5".metadata.fabro] +//! small_default = true //! ``` //! -//! Per-provider and per-model entries field-merge across layers (default → -//! user → server → project → workflow/run). Inner arrays such as -//! `auth.credentials`, `aliases`, `controls.reasoning_effort`, and -//! `controls.speed` replace as whole arrays. -//! -//! Adapter keys (`adapter = "..."`) are parsed as plain strings here. -//! Resolution against the static adapter registry happens in `fabro-model` -//! when the resolved [`Catalog`](fabro_model::Catalog) is built. +//! Layers merge the same way lithos merges overlays: tables merge key by key +//! and every other value replaces. Fabro never interprets the table; lithos +//! validates it when the catalog is built. -use std::collections::{BTreeMap, HashMap}; - -use fabro_model::catalog::deserialize_knowledge_cutoff; -use fabro_model::{ - AgentProfileKind, BillingPolicy, CodecKind, ModelId, ProviderAuthConfig, ProviderId, catalog, -}; -pub use fabro_model::{CredentialRef, CredentialRefParseError, ReasoningEffortFeature}; -use fabro_types::settings::InterpString; use serde::{Deserialize, Serialize}; +use toml::Table; -use super::maps::MergeMap; +use super::combine::Combine; -/// Top-level `[llm]` settings layer. -#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize, fabro_macros::Combine)] -#[serde(deny_unknown_fields)] -pub struct LlmLayer { - /// Provider definitions keyed by provider ID. - #[serde(default, skip_serializing_if = "MergeMap::is_empty")] - pub providers: MergeMap, - /// Legacy top-level model definitions. New settings put models below - /// their provider; parsing normalizes this map before layers combine. - #[serde(default, skip_serializing_if = "MergeMap::is_empty")] - pub models: MergeMap, -} - -/// One entry in `[llm.providers.]`. -#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize, fabro_macros::Combine)] -#[serde(deny_unknown_fields)] -pub struct ProviderSettings { - #[serde(default, skip_serializing_if = "Option::is_none")] - pub display_name: Option, - /// Adapter registry key (e.g. `"openai_compatible"`). - #[serde(default, skip_serializing_if = "Option::is_none")] - pub adapter: Option, - /// Wire dialect for this provider's routes (e.g. `"anthropic_messages"`). - /// Defaults to the adapter's codec; only the default pairing is accepted - /// today — validated at catalog build. - #[serde(default, skip_serializing_if = "Option::is_none")] - pub codec: Option, - /// Agent profile used for routing/profile-specific behavior. - #[serde(default, skip_serializing_if = "Option::is_none")] - pub agent_profile: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub auth: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub billing_policy: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub api_key_url: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub base_url: Option, - /// Extra HTTP headers attached to every outgoing provider request after - /// credential resolution. Values are literal text or - /// `{{ secrets.NAME }}` interpolation strings. Put credentials in a secret - /// and reference them with a token, not a bare literal. - #[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, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub enabled: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub aliases: Option>, - /// Model offerings served by this provider, keyed by canonical model ID. - #[serde(default, skip_serializing_if = "MergeMap::is_empty")] - pub models: MergeMap, -} - -/// One entry in `[llm.providers..models.]`. -#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize, fabro_macros::Combine)] -#[serde(deny_unknown_fields)] -pub struct ModelSettings { - /// Compatibility-only provider for legacy `[llm.models.]` rows. - #[serde(default, skip_serializing_if = "Option::is_none")] - pub provider: Option, - /// Identifier sent to the provider API. Defaults to the catalog model ID - /// when omitted. - #[serde(default, skip_serializing_if = "Option::is_none")] - pub api_id: Option, - /// Wire dialect for this model's route, overriding the provider's codec. - /// Only the adapter's default pairing is accepted today — validated at - /// catalog build. - #[serde(default, skip_serializing_if = "Option::is_none")] - pub codec: Option, - /// Billing family for this model, overriding the provider's policy - /// (e.g. Anthropic cache billing for a Claude model served through an - /// aggregator). - #[serde(default, skip_serializing_if = "Option::is_none")] - pub billing_policy: Option, - /// Agent profile used for routing/profile-specific behavior. - #[serde(default, skip_serializing_if = "Option::is_none")] - pub agent_profile: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub display_name: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub family: Option, - /// Training data cutoff label. Built-ins keep the exact public string - /// already exposed by the model API. - #[serde(default, skip_serializing_if = "Option::is_none")] - pub training: Option, - /// Public knowledge cutoff label. Built-ins keep values such as - /// `"May 2025"` exactly; bare TOML dates are normalized to `YYYY-MM-DD`. - #[serde( - default, - deserialize_with = "deserialize_knowledge_cutoff", - skip_serializing_if = "Option::is_none" - )] - pub knowledge_cutoff: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub default: Option, - /// Whether this model should be preferred for small/cheap utility tasks. - /// Missing or false falls back to the provider default model. - #[serde(default, skip_serializing_if = "Option::is_none")] - pub small_default: Option, - /// Whether this model should be preferred for provider connectivity - /// probes. Missing or false falls back to the provider default model. - #[serde(default, skip_serializing_if = "Option::is_none")] - pub probe: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub enabled: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub aliases: Option>, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub estimated_output_tps: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub limits: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub features: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub controls: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub costs: Option, -} +/// Top-level `[llm]` settings layer: a raw lithos catalog overlay. +#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)] +#[serde(transparent)] +pub struct LlmLayer(pub Table); impl LlmLayer { - /// Normalize the temporary legacy model table before this source is - /// combined with any other settings source. Resolution runs against this - /// layer's own providers plus the built-in catalog, because a single - /// source may reference built-in offerings that merge in later. - pub(crate) fn normalize_legacy_models(&mut self) -> Result<(), catalog::LegacyModelError> { - for (provider, settings) in self.providers.iter() { - for (model, settings) in settings.models.iter() { - if settings.provider.is_some() { - return Err(catalog::LegacyModelError::ScopedModelDeclaresProvider { - provider: ProviderId::new(provider.clone()), - model: ModelId::new(model.clone()), - }); - } - } - } + #[must_use] + pub fn is_empty(&self) -> bool { + self.0.is_empty() + } - let legacy_models = std::mem::take(&mut self.models.0); - if legacy_models.is_empty() { - return Ok(()); - } - let mut legacy_models = legacy_models.into_iter().collect::>(); - legacy_models.sort_by(|(left, _), (right, _)| left.cmp(right)); - - let mut index = catalog::LegacyModelIndex::default(); - let mut provider_ids = self.providers.keys().cloned().collect::>(); - provider_ids.sort_unstable(); - for provider_id in &provider_ids { - let settings = self - .providers - .get(provider_id) - .expect("provider ID came from provider map keys"); - let mut model_ids = settings.models.keys().cloned().collect::>(); - model_ids.sort_unstable(); - index.add_provider( - ProviderId::new(provider_id.clone()), - settings.aliases.clone().unwrap_or_default(), - model_ids.into_iter().map(|model_id| { - let model = settings - .models - .get(&model_id) - .expect("model ID came from model map keys"); - let aliases = model.aliases.clone().unwrap_or_default(); - (ModelId::new(model_id), aliases) - }), - ); - } - let index = index.with_builtin()?; - - for (legacy_id, mut settings) in legacy_models { - let explicit_provider = settings.provider.take(); - let (provider, model) = index.resolve(&legacy_id, explicit_provider.as_deref())?; - - let provider_settings = self.providers.entry(provider.to_string()).or_default(); - if provider_settings.models.contains_key(model.as_str()) { - return Err(catalog::LegacyModelError::DuplicateModel { provider, model }); - } - provider_settings - .models - .insert(model.into_inner(), settings); - } - Ok(()) + /// Render this layer as a lithos catalog overlay document. + #[must_use] + pub fn to_overlay_toml(&self) -> String { + toml::to_string(&self.0).expect("a TOML table always serializes") } } -#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize, fabro_macros::Combine)] -#[serde(deny_unknown_fields)] -pub struct ModelLimits { - #[serde(default, skip_serializing_if = "Option::is_none")] - pub context_window: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub max_output: Option, +impl Combine for LlmLayer { + fn combine(self, other: Self) -> Self { + let mut base = toml::Value::Table(other.0); + merge(&mut base, toml::Value::Table(self.0)); + match base { + toml::Value::Table(table) => Self(table), + _ => unreachable!("merging two tables yields a table"), + } + } } -#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize, fabro_macros::Combine)] -#[serde(deny_unknown_fields)] -pub struct ModelFeatures { - #[serde(default, skip_serializing_if = "Option::is_none")] - pub tools: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub vision: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub reasoning: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub reasoning_by_default: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub reasoning_effort: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub prompt_cache: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub cache_control_breakpoints: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub sampling_params: Option, -} - -/// User-facing allow-list for native control values Fabro accepts on this -/// model. Whole-array replacement on merge. -#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize, fabro_macros::Combine)] -#[serde(deny_unknown_fields)] -pub struct ModelControls { - /// Allowed reasoning-effort values. Strings (e.g. `"low"`, `"high"`, - /// `"xhigh"`) — validated as `ReasoningEffort` at catalog build. - #[serde(default, skip_serializing_if = "Option::is_none")] - pub reasoning_effort: Option>, - /// Additional speeds beyond `Speed::Standard`. Strings — validated as - /// `Speed` at catalog build. `Speed::Standard` is implicit and must not - /// appear here. - #[serde(default, skip_serializing_if = "Option::is_none")] - pub speed: Option>, -} - -/// Pricing table. Base [`CostRates`] always apply; per-speed overrides -/// substitute when the request specifies a non-standard speed. -#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize, fabro_macros::Combine)] -#[serde(deny_unknown_fields)] -pub struct ModelCostTable { - #[serde(flatten)] - pub base: CostRates, - /// Per-speed cost overrides (e.g. `costs.speed.fast = { ... }`). Keys - /// must reference a speed declared in `controls.speed`. `standard` is - /// not a valid override key — base rates serve standard speed. - #[serde(default, skip_serializing_if = "Option::is_none")] - pub speed: Option>, -} - -#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize, fabro_macros::Combine)] -#[serde(deny_unknown_fields)] -pub struct CostRates { - #[serde(default, skip_serializing_if = "Option::is_none")] - pub input_cost_per_mtok: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub output_cost_per_mtok: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub cache_input_cost_per_mtok: Option, +fn merge(base: &mut toml::Value, overlay: toml::Value) { + match (base, overlay) { + (toml::Value::Table(base), toml::Value::Table(overlay)) => { + for (key, value) in overlay { + if let Some(existing) = base.get_mut(&key) { + merge(existing, value); + } else { + base.insert(key, value); + } + } + } + (base, overlay) => *base = overlay, + } } #[cfg(test)] mod tests { - use std::str::FromStr; - - use fabro_model::ApiKeyHeaderPolicy; - use super::*; - use crate::layers::Combine; - // ---- CredentialRef ---------------------------------------------------- - - #[test] - fn credential_ref_parses_vault_form() { - let r = CredentialRef::from_str("vault:OPENAI_CODEX").unwrap(); - assert_eq!(r, CredentialRef::Vault("OPENAI_CODEX".to_string())); + fn layer(source: &str) -> LlmLayer { + LlmLayer(toml::from_str(source).unwrap()) } #[test] - fn credential_ref_parses_env_form() { - let r = CredentialRef::from_str("env:KIMI_API_KEY").unwrap(); - assert_eq!(r, CredentialRef::Env("KIMI_API_KEY".to_string())); - } - - #[test] - fn credential_ref_rejects_literal_secret() { - // A literal API key contains no `vault:` or `env:` prefix. - let err = CredentialRef::from_str("sk-ant-1234").unwrap_err(); - assert!(err.to_string().contains("must be")); - assert!( - !err.to_string().contains("sk-ant-1234"), - "error must not echo the literal secret string back to the user", + fn higher_layer_wins_scalars_and_merges_tables() { + let higher = layer( + r" +[providers.acme] +priority = 10 +[providers.acme.metadata.fabro] +enabled = false +", ); - } - - #[test] - fn credential_ref_rejects_empty_vault_name() { - let err = CredentialRef::from_str("vault:").unwrap_err(); - assert!(err.to_string().contains("missing")); - } - - #[test] - fn credential_ref_rejects_empty_env_name() { - let err = CredentialRef::from_str("env:").unwrap_err(); - assert!(err.to_string().contains("missing")); - } - - #[test] - fn credential_ref_round_trips_through_string() { - let r = CredentialRef::Vault("kimi".to_string()); - assert_eq!(r.to_string(), "vault:kimi"); - let back: CredentialRef = r.to_string().parse().unwrap(); - assert_eq!(back, r); - } - - #[test] - fn credential_ref_serializes_as_string_in_toml() { - let r = CredentialRef::Env("KIMI_API_KEY".to_string()); - let s = toml::Value::try_from(&r).unwrap(); - assert_eq!(s.as_str(), Some("env:KIMI_API_KEY")); - } - - #[test] - fn credential_ref_deserializes_from_toml_string() { - let parsed: CredentialRef = toml::from_str(r#"v = "vault:foo""#) - .map(|v: toml::Value| { - v.as_table() - .unwrap() - .get("v") - .unwrap() - .clone() - .try_into() - .unwrap() - }) - .unwrap(); - assert_eq!(parsed, CredentialRef::Vault("foo".to_string())); - } - - #[test] - fn credential_ref_in_array_rejects_literal_secret() { - // serde rejects literal secrets when parsed inside an array of - // CredentialRef. The error bubbles up as a TOML deserialization - // failure. - #[derive(Deserialize)] - #[expect( - dead_code, - reason = "field exists only to drive the deserializer; we assert on the parse error" - )] - struct Wrap { - v: Vec, - } - let err: Result = toml::from_str(r#"v = ["sk-literal-secret"]"#); - assert!(err.is_err(), "literal secret strings must fail to parse"); - } - - #[test] - fn provider_agent_profile_parses_from_toml() { - let parsed: LlmLayer = toml::from_str( + let lower = layer( r#" [providers.acme] -adapter = "openai_compatible" -agent_profile = "anthropic" -"#, - ) - .unwrap(); - - assert_eq!( - parsed.providers.get("acme").unwrap().agent_profile, - Some(fabro_model::AgentProfileKind::Anthropic) - ); - } - - #[test] - fn provider_codec_parses_from_toml() { - let parsed: LlmLayer = toml::from_str( - r#" -[providers.acme] -adapter = "openai_compatible" -codec = "openai_compatible" -"#, - ) - .unwrap(); - - assert_eq!( - parsed.providers.get("acme").unwrap().codec, - Some(fabro_model::CodecKind::OpenAiCompatible) - ); - } - - #[test] - fn model_codec_parses_from_toml() { - let parsed: LlmLayer = toml::from_str( - r#" -[models.acme_large] -provider = "acme" -codec = "anthropic_messages" -"#, - ) - .unwrap(); - - assert_eq!( - parsed.models.get("acme_large").unwrap().codec, - Some(fabro_model::CodecKind::AnthropicMessages) - ); - } - - #[test] - fn model_billing_policy_parses_from_toml() { - let parsed: LlmLayer = toml::from_str( - r#" -[models.acme_claude] -provider = "acme" -billing_policy = "anthropic" -"#, - ) - .unwrap(); - - assert_eq!( - parsed.models.get("acme_claude").unwrap().billing_policy, - Some(fabro_model::BillingPolicy::Anthropic) - ); - } - - #[test] - fn model_agent_profile_parses_from_toml() { - let parsed: LlmLayer = toml::from_str( - r#" -[models.acme_large] -provider = "acme" -agent_profile = "gemini" -"#, - ) - .unwrap(); - - assert_eq!( - parsed.models.get("acme_large").unwrap().agent_profile, - Some(fabro_model::AgentProfileKind::Gemini) - ); - } - - // ---- Provider extra headers ------------------------------------------ - - #[expect( - clippy::disallowed_methods, - reason = "tests assert unresolved interpolation header source round-trips" - )] - fn interp_source(value: &InterpString) -> String { - value.as_source() - } - - // ---- LlmLayer parsing ------------------------------------------------- - - #[test] - fn parses_minimal_provider_entry() { - let toml = r#" -[providers.moonshot] -display_name = "Moonshot AI" -adapter = "openai_compatible" -agent_profile = "openai" -base_url = "https://api.moonshot.ai/v1" -priority = 60 +priority = 5 +base_url = "https://acme.test" +[providers.acme.metadata.fabro] enabled = true -aliases = ["moonshot-ai"] - -[providers.moonshot.auth] -credentials = [ - "env:MOONSHOT_API_KEY", - "env:KIMI_API_KEY", - "vault:MOONSHOT_API_KEY", - "vault:KIMI_API_KEY", -] -"#; - let layer: LlmLayer = toml::from_str(toml).unwrap(); - let moonshot = layer.providers.get("moonshot").unwrap(); - assert_eq!(moonshot.display_name.as_deref(), Some("Moonshot AI")); - assert_eq!(moonshot.adapter.as_deref(), Some("openai_compatible")); - assert_eq!(moonshot.agent_profile, Some(AgentProfileKind::OpenAi)); - let auth = moonshot.auth.as_ref().expect("expected api_key auth"); - assert_eq!(auth.header, ApiKeyHeaderPolicy::Bearer); - assert_eq!(auth.credentials, vec![ - CredentialRef::Env("MOONSHOT_API_KEY".to_string()), - CredentialRef::Env("KIMI_API_KEY".to_string()), - CredentialRef::Vault("MOONSHOT_API_KEY".to_string()), - CredentialRef::Vault("KIMI_API_KEY".to_string()), - ]); - assert_eq!( - moonshot.base_url.as_deref(), - Some("https://api.moonshot.ai/v1") +small_default = true +"#, ); - assert_eq!(moonshot.priority, Some(60)); - assert_eq!(moonshot.enabled, Some(true)); + let merged = higher.combine(lower).0; + let acme = &merged["providers"]["acme"]; + assert_eq!(acme["priority"].as_integer(), Some(10)); + assert_eq!(acme["base_url"].as_str(), Some("https://acme.test")); + assert_eq!(acme["metadata"]["fabro"]["enabled"].as_bool(), Some(false)); assert_eq!( - moonshot.aliases.as_deref(), - Some(&["moonshot-ai".to_string()][..]) + acme["metadata"]["fabro"]["small_default"].as_bool(), + Some(true) ); } #[test] - fn provider_extra_headers_parse_interp_tokens() { - let toml = r#" -[providers.portkey] -display_name = "Portkey Bedrock" -adapter = "anthropic" -base_url = "https://api.portkey.ai/v1" - -[providers.portkey.extra_headers] -x-title = "My App" -x-portkey-api-key = "{{ env.PORTKEY_API_KEY }}" -x-team-secret = "{{ secrets.gateway_team_secret }}" -"#; - - let layer: LlmLayer = toml::from_str(toml).unwrap(); - let portkey = layer.providers.get("portkey").unwrap(); - - assert!(portkey.auth.is_none()); - let headers = portkey.extra_headers.as_ref().unwrap(); + fn arrays_replace_whole() { + let higher = layer("[providers.acme]\naliases = [\"a\"]\n"); + let lower = layer("[providers.acme]\naliases = [\"b\", \"c\"]\n"); + let merged = higher.combine(lower).0; assert_eq!( - interp_source(headers.get("x-title").expect("x-title header should parse")), - "My App", - ); - assert_eq!( - interp_source( - headers - .get("x-portkey-api-key") - .expect("x-portkey-api-key header should parse") - ), - "{{ env.PORTKEY_API_KEY }}", - ); - assert_eq!( - interp_source( - headers - .get("x-team-secret") - .expect("x-team-secret header should parse") - ), - "{{ secrets.gateway_team_secret }}", + merged["providers"]["acme"]["aliases"] + .as_array() + .map(Vec::len), + Some(1) ); } #[test] - fn provider_extra_headers_accepts_bare_string_literal() { - let toml = r#" -[providers.portkey.extra_headers] -x-portkey-api-key = "sk-portkey-literal" -"#; - - let layer: LlmLayer = toml::from_str(toml).unwrap(); - let headers = layer - .providers - .get("portkey") - .unwrap() - .extra_headers - .as_ref() - .unwrap(); - let header = headers.get("x-portkey-api-key").unwrap(); - - assert!(header.is_literal()); - assert_eq!(interp_source(header), "sk-portkey-literal"); - } - - #[test] - fn parses_full_model_entry() { - let toml = r#" -[models."kimi-k2.5"] -provider = "moonshot" -api_id = "kimi-k2.5" -display_name = "Kimi K2.5" -family = "kimi" -training = "2025-01-01" -knowledge_cutoff = 2025-01-01 -default = true -enabled = true -aliases = ["kimi"] -estimated_output_tps = 50 - -[models."kimi-k2.5".limits] -context_window = 262144 -max_output = 32768 - -[models."kimi-k2.5".features] -tools = true -vision = false -reasoning = true - -[models."kimi-k2.5".costs] -input_cost_per_mtok = 0.60 -output_cost_per_mtok = 2.50 -cache_input_cost_per_mtok = 0.15 -"#; - let layer: LlmLayer = toml::from_str(toml).unwrap(); - let m = layer.models.get("kimi-k2.5").unwrap(); - assert_eq!(m.provider.as_deref(), Some("moonshot")); - assert_eq!(m.api_id.as_deref(), Some("kimi-k2.5")); - assert_eq!(m.display_name.as_deref(), Some("Kimi K2.5")); - assert_eq!(m.family.as_deref(), Some("kimi")); - assert_eq!(m.training.as_deref(), Some("2025-01-01")); - assert_eq!(m.knowledge_cutoff.as_deref(), Some("2025-01-01")); - assert_eq!(m.default, Some(true)); - assert_eq!(m.enabled, Some(true)); - assert_eq!(m.aliases.as_deref(), Some(&["kimi".to_string()][..])); - assert_eq!(m.estimated_output_tps, Some(50.0)); - - let limits = m.limits.as_ref().unwrap(); - assert_eq!(limits.context_window, Some(262_144)); - assert_eq!(limits.max_output, Some(32_768)); - - let features = m.features.as_ref().unwrap(); - assert_eq!(features.tools, Some(true)); - assert_eq!(features.vision, Some(false)); - assert_eq!(features.reasoning, Some(true)); - - let costs = m.costs.as_ref().unwrap(); - assert_eq!(costs.base.input_cost_per_mtok, Some(0.60)); - assert_eq!(costs.base.output_cost_per_mtok, Some(2.50)); - assert_eq!(costs.base.cache_input_cost_per_mtok, Some(0.15)); - assert!(costs.speed.is_none()); - } - - #[test] - fn parses_model_reasoning_effort_and_prompt_cache_features() { - let toml = r#" -[models."claude-bedrock"] -provider = "bedrock" - -[models."claude-bedrock".features] -tools = true -vision = true -reasoning = true -reasoning_by_default = false -reasoning_effort = "levels" -prompt_cache = false -"#; - let layer: LlmLayer = toml::from_str(toml).unwrap(); - let features = layer - .models - .get("claude-bedrock") - .unwrap() - .features - .as_ref() - .unwrap(); - - assert_eq!( - features.reasoning_effort, - Some(fabro_model::ReasoningEffortFeature::Levels) - ); - assert_eq!(features.reasoning_by_default, Some(false)); - assert_eq!(features.prompt_cache, Some(false)); - } - - #[test] - fn parses_knowledge_cutoff_display_label() { - let toml = r#" -[models."claude-opus-4-7"] -provider = "anthropic" -knowledge_cutoff = "May 2025" -"#; - let layer: LlmLayer = toml::from_str(toml).unwrap(); - let m = layer.models.get("claude-opus-4-7").unwrap(); - - assert_eq!(m.knowledge_cutoff.as_deref(), Some("May 2025")); - } - - #[test] - fn parses_controls_and_per_speed_costs() { - let toml = r#" -[models."claude-opus-4-6".controls] -reasoning_effort = ["low", "medium", "high"] -speed = ["fast"] - -[models."claude-opus-4-6".costs.speed.fast] -input_cost_per_mtok = 90.0 -output_cost_per_mtok = 450.0 -cache_input_cost_per_mtok = 9.0 -"#; - let layer: LlmLayer = toml::from_str(toml).unwrap(); - let m = layer.models.get("claude-opus-4-6").unwrap(); - - let controls = m.controls.as_ref().unwrap(); - assert_eq!( - controls.reasoning_effort.as_deref(), - Some(&["low".to_string(), "medium".to_string(), "high".to_string()][..]) - ); - assert_eq!(controls.speed.as_deref(), Some(&["fast".to_string()][..])); - - let costs = m.costs.as_ref().unwrap(); - let fast = costs.speed.as_ref().unwrap().get("fast").unwrap(); - assert_eq!(fast.input_cost_per_mtok, Some(90.0)); - assert_eq!(fast.output_cost_per_mtok, Some(450.0)); - assert_eq!(fast.cache_input_cost_per_mtok, Some(9.0)); - } - - #[test] - fn rejects_unknown_provider_field() { - let toml = r#" -[providers.moonshot] -adapter = "openai_compatible" -unknown_field = true -"#; - let err = toml::from_str::(toml).unwrap_err(); - assert!(err.to_string().contains("unknown_field")); - } - - #[test] - fn rejects_removed_provider_base_url_env_field() { - let toml = r#" -[providers.moonshot] -adapter = "openai_compatible" -base_url_env = "KIMI_BASE_URL" -"#; - let err = toml::from_str::(toml).unwrap_err(); - assert!(err.to_string().contains("base_url_env")); - } - - #[test] - fn rejects_unknown_model_field() { - let toml = r#" -[models.foo] -provider = "x" -mystery = 1 -"#; - let err = toml::from_str::(toml).unwrap_err(); - assert!(err.to_string().contains("mystery")); - } - - // ---- Combine / merge -------------------------------------------------- - - #[test] - fn provider_field_merge_keeps_self_values_and_fills_holes() { - let high = ProviderSettings { - adapter: Some("openai_compatible".to_string()), - base_url: Some("https://override.example".to_string()), - agent_profile: Some(fabro_model::AgentProfileKind::Anthropic), - ..ProviderSettings::default() - }; - let low = ProviderSettings { - adapter: Some("anthropic".to_string()), - base_url: Some("https://defaults.example".to_string()), - display_name: Some("Default".to_string()), - priority: Some(10), - agent_profile: Some(fabro_model::AgentProfileKind::OpenAi), - ..ProviderSettings::default() - }; - let merged = high.combine(low); - assert_eq!(merged.adapter.as_deref(), Some("openai_compatible")); - assert_eq!(merged.base_url.as_deref(), Some("https://override.example")); - assert_eq!(merged.display_name.as_deref(), Some("Default")); - assert_eq!(merged.priority, Some(10)); - assert_eq!( - merged.agent_profile, - Some(fabro_model::AgentProfileKind::Anthropic) - ); - } - - #[test] - fn provider_auth_replaces_wholesale() { - // Higher layer redeclares auth, so the low layer's auth table is - // dropped entirely (whole-value replacement). - let high = ProviderSettings { - auth: Some(ProviderAuthConfig { - credentials: vec![CredentialRef::Env("FOO".to_string())], - header: ApiKeyHeaderPolicy::Bearer, - }), - ..ProviderSettings::default() - }; - let low = ProviderSettings { - auth: Some(ProviderAuthConfig { - credentials: vec![ - CredentialRef::Vault("bar".to_string()), - CredentialRef::Env("BAZ".to_string()), - ], - header: ApiKeyHeaderPolicy::Custom { - name: "x-api-key".to_string(), - }, - }), - ..ProviderSettings::default() - }; - let merged = high.combine(low); - assert_eq!( - merged.auth, - Some(ProviderAuthConfig { - credentials: vec![CredentialRef::Env("FOO".to_string())], - header: ApiKeyHeaderPolicy::Bearer, - }) - ); - } - - #[test] - fn provider_auth_inherits_when_unset_in_higher_layer() { - let high = ProviderSettings::default(); - let low = ProviderSettings { - auth: Some(ProviderAuthConfig { - credentials: vec![CredentialRef::Env("FOO".to_string())], - header: ApiKeyHeaderPolicy::Bearer, - }), - ..ProviderSettings::default() - }; - let merged = high.combine(low); - assert_eq!( - merged.auth, - Some(ProviderAuthConfig { - credentials: vec![CredentialRef::Env("FOO".to_string())], - header: ApiKeyHeaderPolicy::Bearer, - }) - ); - } - - #[test] - fn provider_extra_headers_map_replaces_wholesale() { - let high = ProviderSettings { - extra_headers: Some(HashMap::from([( - "x-portkey-provider".to_string(), - InterpString::from("@bedrock-prod"), - )])), - ..ProviderSettings::default() - }; - let low = ProviderSettings { - extra_headers: Some(HashMap::from([ - ( - "x-portkey-api-key".to_string(), - InterpString::from("{{ env.PORTKEY_API_KEY }}"), - ), - ( - "x-portkey-provider".to_string(), - InterpString::from("@bedrock-default"), - ), - ])), - ..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(&InterpString::from("@bedrock-prod")), - ); - 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(), - InterpString::from("{{ env.PORTKEY_API_KEY }}"), - )])), - ..ProviderSettings::default() - }; - - let merged = high.combine(low); - - assert_eq!( - merged.extra_headers.unwrap().get("x-portkey-api-key"), - Some(&InterpString::from("{{ env.PORTKEY_API_KEY }}")), - ); - } - - #[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(), - InterpString::from("{{ env.PORTKEY_API_KEY }}"), - )])), - ..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 = - std::collections::HashMap::new(); - high_map.insert("moonshot".to_string(), ProviderSettings { - base_url: Some("https://override".to_string()), - ..ProviderSettings::default() - }); - let high: MergeMap = MergeMap::from(high_map); - - let mut low_map: std::collections::HashMap = - std::collections::HashMap::new(); - low_map.insert("moonshot".to_string(), ProviderSettings { - adapter: Some("openai_compatible".to_string()), - base_url: Some("https://defaults".to_string()), - ..ProviderSettings::default() - }); - let low: MergeMap = MergeMap::from(low_map); - - let merged = high.combine(low); - let moonshot = merged.get("moonshot").unwrap(); - assert_eq!(moonshot.adapter.as_deref(), Some("openai_compatible")); - assert_eq!(moonshot.base_url.as_deref(), Some("https://override")); - } - - #[test] - fn model_controls_replace_wholesale() { - // Whole-array replacement: high layer's `reasoning_effort` shadows - // the low layer's list completely. - let high = ModelControls { - reasoning_effort: Some(vec!["high".to_string()]), - ..ModelControls::default() - }; - let low = ModelControls { - reasoning_effort: Some(vec!["low".to_string(), "high".to_string()]), - speed: Some(vec!["fast".to_string()]), - }; - let merged = high.combine(low); - assert_eq!( - merged.reasoning_effort.as_deref(), - Some(&["high".to_string()][..]) - ); - assert_eq!(merged.speed.as_deref(), Some(&["fast".to_string()][..])); - } - - #[test] - fn model_agent_profile_merges_as_scalar() { - let high = ModelSettings { - agent_profile: Some(fabro_model::AgentProfileKind::Gemini), - ..ModelSettings::default() - }; - let low = ModelSettings { - agent_profile: Some(fabro_model::AgentProfileKind::Anthropic), - ..ModelSettings::default() - }; - - assert_eq!( - high.combine(low).agent_profile, - Some(fabro_model::AgentProfileKind::Gemini) - ); + fn overlay_toml_round_trips() { + let source = layer("[providers.acme]\npriority = 3\n"); + let rendered = source.to_overlay_toml(); + assert_eq!(layer(&rendered), source); } } diff --git a/lib/foundation/fabro-config/src/layers/mod.rs b/lib/foundation/fabro-config/src/layers/mod.rs index c3fa1632c..f8b625c1d 100644 --- a/lib/foundation/fabro-config/src/layers/mod.rs +++ b/lib/foundation/fabro-config/src/layers/mod.rs @@ -20,11 +20,7 @@ pub use environment::{ EnvironmentDockerfileLayer, EnvironmentImageLayer, EnvironmentLayer, EnvironmentLifecycleLayer, EnvironmentNetworkLayer, EnvironmentResourcesLayer, RunEnvironmentLayer, }; -pub use llm::{ - CostRates, CredentialRef, CredentialRefParseError, LlmLayer, ModelControls, ModelCostTable, - ModelFeatures as LlmModelFeatures, ModelLimits as LlmModelLimits, ModelSettings, - ProviderSettings, ReasoningEffortFeature, -}; +pub use llm::LlmLayer; pub use log_filter::LogFilter; pub use maps::{MergeMap, ReplaceMap, StickyMap}; pub use project::ProjectLayer; diff --git a/lib/foundation/fabro-config/src/lib.rs b/lib/foundation/fabro-config/src/lib.rs index f6097b9be..33eeed889 100644 --- a/lib/foundation/fabro-config/src/lib.rs +++ b/lib/foundation/fabro-config/src/lib.rs @@ -32,8 +32,7 @@ use std::path::Path; pub use builders::{ ResolveErrors, RunSettingsBuilder, ServerRuntimeSettings, ServerSettingsBuilder, - UserSettingsBuilder, WorkflowSettingsBuilder, load_llm_catalog_settings, - load_server_runtime_settings, + UserSettingsBuilder, WorkflowSettingsBuilder, load_llm_overlay, load_server_runtime_settings, }; pub use error::{Error, Result}; pub use fabro_util::path::expand_tilde; @@ -42,22 +41,21 @@ pub use input_overrides::{InputOverrideParseError, parse_input_overrides, parse_ pub(crate) use layers::Combine; pub use layers::{ CliAuthLayer, CliExecAgentLayer, CliExecLayer, CliExecModelLayer, CliLayer, CliLoggingLayer, - CliOutputLayer, CliTargetLayer, CliUpdatesLayer, CostRates, CredentialRef, - CredentialRefParseError, EnvironmentDockerfileLayer, EnvironmentImageLayer, EnvironmentLayer, - EnvironmentLifecycleLayer, EnvironmentNetworkLayer, EnvironmentResourcesLayer, 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, ReasoningEffortFeature, ReplaceMap, RunAgentLayer, - RunArtifactsLayer, RunCheckpointLayer, RunCloneLayer, RunEnvironmentLayer, RunExecutionLayer, - RunGitLayer, RunGoalLayer, RunIntegrationsGithubLayer, RunIntegrationsLayer, RunLayer, - RunMetaBranchLayer, RunModelControlsLayer, RunModelLayer, RunPrepareLayer, RunPullRequestLayer, - RunRunBranchLayer, RunScmLayer, ScmGitHubLayer, ServerApiLayer, ServerArtifactsLayer, - ServerAuthGithubLayer, ServerAuthLayer, ServerIntegrationsLayer, ServerLayer, - ServerListenLayer, ServerLoggingLayer, ServerSandboxLayer, ServerSandboxProviderLayer, - ServerSandboxProvidersLayer, ServerSchedulerLayer, ServerSlateDbLayer, ServerStorageLayer, - ServerWebLayer, SettingsLayer, SlackIntegrationLayer, StickyMap, StringOrSplice, WorkflowLayer, + CliOutputLayer, CliTargetLayer, CliUpdatesLayer, EnvironmentDockerfileLayer, + EnvironmentImageLayer, EnvironmentLayer, EnvironmentLifecycleLayer, EnvironmentNetworkLayer, + EnvironmentResourcesLayer, GitAuthorLayer, GithubIntegrationLayer, HookAgentMarker, HookEntry, + HookTlsMode, IntegrationWebhooksLayer, InterviewProviderLayer, InterviewsLayer, LlmLayer, + LogFilter, McpEntryLayer, MergeMap, ModelRefOrSplice, NotificationProviderLayer, + NotificationRouteLayer, ObjectStoreLocalLayer, ObjectStoreS3Layer, PrepareStep, ProjectLayer, + ReplaceMap, RunAgentLayer, RunArtifactsLayer, RunCheckpointLayer, RunCloneLayer, + RunEnvironmentLayer, RunExecutionLayer, RunGitLayer, RunGoalLayer, RunIntegrationsGithubLayer, + RunIntegrationsLayer, RunLayer, RunMetaBranchLayer, RunModelControlsLayer, RunModelLayer, + RunPrepareLayer, RunPullRequestLayer, RunRunBranchLayer, RunScmLayer, ScmGitHubLayer, + ServerApiLayer, ServerArtifactsLayer, ServerAuthGithubLayer, ServerAuthLayer, + ServerIntegrationsLayer, ServerLayer, ServerListenLayer, ServerLoggingLayer, + ServerSandboxLayer, ServerSandboxProviderLayer, ServerSandboxProvidersLayer, + ServerSchedulerLayer, ServerSlateDbLayer, ServerStorageLayer, ServerWebLayer, SettingsLayer, + SlackIntegrationLayer, StickyMap, StringOrSplice, WorkflowLayer, }; pub use logging::{resolve_log_destination, resolve_log_destination_with_env}; pub use parse::ParseError; diff --git a/lib/foundation/fabro-config/src/parse.rs b/lib/foundation/fabro-config/src/parse.rs index 8ed2273bd..8051b1286 100644 --- a/lib/foundation/fabro-config/src/parse.rs +++ b/lib/foundation/fabro-config/src/parse.rs @@ -1,7 +1,5 @@ use std::fmt; -use fabro_model::catalog::LegacyModelError; - use crate::SettingsLayer; const CURRENT_VERSION: u32 = 1; @@ -32,7 +30,6 @@ const LEGACY_LLM_KEYS: &[&str] = &[ #[derive(Debug, Clone, PartialEq, Eq)] pub enum ParseError { Toml(String), - LlmCatalog(LegacyModelError), Version(VersionError), UnknownTopLevelKey { key: String, @@ -48,7 +45,6 @@ impl fmt::Display for ParseError { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { match self { Self::Toml(msg) => write!(f, "settings file is not valid TOML: {msg}"), - Self::LlmCatalog(err) => fmt::Display::fmt(err, f), Self::Version(err) => fmt::Display::fmt(err, f), Self::UnknownTopLevelKey { key, hint } => { if let Some(hint) = hint { @@ -122,14 +118,8 @@ pub(crate) fn parse_settings(input: &str) -> Result { } } - let mut layer = raw - .try_into::() - .map_err(|e| ParseError::Toml(e.to_string()))?; - if let Some(llm) = layer.llm.as_mut() { - llm.normalize_legacy_models() - .map_err(ParseError::LlmCatalog)?; - } - Ok(layer) + raw.try_into::() + .map_err(|e| ParseError::Toml(e.to_string())) } #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -291,227 +281,36 @@ mod tests { } #[test] - fn accepts_new_llm_providers_subtree() { - let parsed = "[llm.providers.moonshot]\nadapter = \"openai_compatible\"\n" + fn accepts_llm_overlay_subtree() { + let parsed = "[llm.providers.moonshot]\npriority = 60\n" .parse::() .unwrap(); - assert!(parsed.llm.unwrap().providers.contains_key("moonshot")); - } - - #[test] - fn accepts_new_llm_models_subtree() { - let parsed = "[llm.providers.moonshot.models.\"foo\"]\n" - .parse::() - .unwrap(); - assert!( - parsed - .llm - .unwrap() - .providers - .get("moonshot") - .unwrap() - .models - .contains_key("foo") + let llm = parsed.llm.unwrap(); + assert_eq!( + llm.0["providers"]["moonshot"]["priority"].as_integer(), + Some(60) ); } #[test] - fn provider_scoped_model_rejects_redundant_provider_field() { - let error = r#" -[llm.providers.openai.models."gpt-5.4"] -provider = "openai" -"# - .parse::() - .unwrap_err(); - - assert!(matches!( - error, - ParseError::LlmCatalog(LegacyModelError::ScopedModelDeclaresProvider { - provider, - model, - }) if provider.as_str() == "openai" && model.as_str() == "gpt-5.4" - )); - } - - #[test] - fn legacy_model_row_with_provider_normalizes_before_merge() { - use crate::layers::Combine as _; - + fn llm_overlay_merges_across_layers() { let higher = r#" -[llm.models."gpt-5.4"] -provider = "openai" -display_name = "Configured display name" +[llm.providers.openai.models."gpt-5.4".metadata.fabro] +small_default = true "# .parse::() .unwrap(); - let fallback = r#" -[llm.providers.openai.models."gpt-5.4"] -family = "gpt-5" + let lower = r#" +[llm.providers.openai.models."gpt-5.4".metadata.fabro] +probe = true "# .parse::() .unwrap(); - - let merged = higher.combine(fallback); - let llm = merged.llm.unwrap(); - assert!(llm.models.is_empty()); - let model = llm - .providers - .get("openai") - .unwrap() - .models - .get("gpt-5.4") - .unwrap(); - assert_eq!( - model.display_name.as_deref(), - Some("Configured display name") - ); - assert_eq!(model.family.as_deref(), Some("gpt-5")); - } - - #[test] - fn provider_less_legacy_row_adopts_unique_builtin_offering() { - let parsed = r#" -[llm.models.mercury] -display_name = "Configured Mercury" -"# - .parse::() - .unwrap(); - let llm = parsed.llm.unwrap(); - - assert!(llm.models.is_empty()); - assert!( - llm.providers - .get("inception") - .unwrap() - .models - .contains_key("mercury-2") - ); - } - - #[test] - fn provider_less_legacy_row_rejects_ambiguous_builtin_offering() { - let error = r#" -[llm.models."gpt-5.6-sol"] -display_name = "Ambiguous" -"# - .parse::() - .unwrap_err(); - - assert!(matches!( - error, - ParseError::LlmCatalog(LegacyModelError::AmbiguousModel { - model, - candidates, - }) if model == "gpt-5.6-sol" && candidates.len() >= 2 - )); - } - - #[test] - fn same_source_legacy_and_provider_scoped_rows_conflict() { - let error = r#" -[llm.providers.openai.models."gpt-5.4"] -display_name = "Canonical" - -[llm.models."gpt-5.4"] -provider = "openai" -display_name = "Legacy" -"# - .parse::() - .unwrap_err(); - - assert!(matches!( - error, - ParseError::LlmCatalog(LegacyModelError::DuplicateModel { - provider, - model, - }) if provider.as_str() == "openai" && model.as_str() == "gpt-5.4" - )); - } - - #[test] - fn legacy_builtin_model_id_normalizes_with_explicit_provider() { - let parsed = r#" -[llm.models."openai/gpt-5.6-sol"] -provider = "openrouter" -display_name = "Configured Sol" -"# - .parse::() - .unwrap(); - let llm = parsed.llm.unwrap(); - - assert!(llm.models.is_empty()); - let model = llm - .providers - .get("openrouter") - .unwrap() - .models - .get("gpt-5.6-sol") - .unwrap(); - assert_eq!(model.display_name.as_deref(), Some("Configured Sol")); - assert!(model.provider.is_none()); - } - - #[test] - fn legacy_builtin_model_id_without_provider_uses_historical_catalog_provider() { - let parsed = r#" -[llm.models."anthropic/claude-fable-5"] -display_name = "Configured Fable" -"# - .parse::() - .unwrap(); - let llm = parsed.llm.unwrap(); - - assert!(llm.models.is_empty()); - assert!( - llm.providers - .get("openrouter") - .unwrap() - .models - .contains_key("claude-fable-5") - ); - } - - #[test] - fn same_model_slug_on_different_providers_merges_independently() { - use crate::layers::Combine as _; - - let direct = r#" -[llm.providers.openai.models.shared] -display_name = "Direct" -"# - .parse::() - .unwrap(); - let aggregator = r#" -[llm.providers.openrouter.models.shared] -display_name = "Aggregator" -"# - .parse::() - .unwrap(); - - let merged = direct.combine(aggregator); - let providers = merged.llm.unwrap().providers; - assert_eq!( - providers - .get("openai") - .unwrap() - .models - .get("shared") - .unwrap() - .display_name - .as_deref(), - Some("Direct") - ); - assert_eq!( - providers - .get("openrouter") - .unwrap() - .models - .get("shared") - .unwrap() - .display_name - .as_deref(), - Some("Aggregator") - ); + let merged = crate::Combine::combine(higher, lower); + let fabro = + &merged.llm.unwrap().0["providers"]["openai"]["models"]["gpt-5.4"]["metadata"]["fabro"]; + assert_eq!(fabro["small_default"].as_bool(), Some(true)); + assert_eq!(fabro["probe"].as_bool(), Some(true)); } #[test] diff --git a/lib/foundation/fabro-types/Cargo.toml b/lib/foundation/fabro-types/Cargo.toml index 273309c73..4cd59e7bb 100644 --- a/lib/foundation/fabro-types/Cargo.toml +++ b/lib/foundation/fabro-types/Cargo.toml @@ -21,9 +21,9 @@ workspace = true chrono = { workspace = true, features = ["serde"] } clap = { workspace = true, optional = true } dirs.workspace = true -fabro-model = { path = "../fabro-model" } fabro-util = { path = "../fabro-util" } hex.workspace = true +lithos-llm = { workspace = true, features = ["runtime"] } serde.workspace = true serde_json.workspace = true sha2.workspace = true diff --git a/lib/foundation/fabro-types/src/agent_profile.rs b/lib/foundation/fabro-types/src/agent_profile.rs new file mode 100644 index 000000000..7dbe74fa2 --- /dev/null +++ b/lib/foundation/fabro-types/src/agent_profile.rs @@ -0,0 +1,81 @@ +//! Agent profile vocabulary shared by the catalog policy and the agent. +//! +//! The catalog records which profile a model should run under in its +//! `metadata.fabro.agent_profile` entry. This enum is the Rust spelling of +//! that value. + +use serde::{Deserialize, Serialize}; +use strum::{Display, EnumString, IntoStaticStr, VariantArray}; + +#[derive( + Debug, + Clone, + Copy, + PartialEq, + Eq, + Hash, + Serialize, + Deserialize, + Display, + EnumString, + IntoStaticStr, + VariantArray, +)] +#[serde(rename_all = "snake_case")] +#[strum(serialize_all = "snake_case")] +pub enum AgentProfileKind { + Anthropic, + /// Claude 5 models trained against Anthropic's current coding-agent + /// harness. This remains model-scoped so older Claude models keep the + /// established Anthropic profile. + #[serde(rename = "claude-5")] + #[strum(to_string = "claude-5")] + Claude5, + #[serde(rename = "openai")] + #[strum(to_string = "openai")] + OpenAi, + Gemini, + /// Kimi (Moonshot) models, wherever they are served from. Selected per + /// model rather than per provider, so a Kimi model reached through a + /// gateway such as OpenRouter gets the same profile as one reached + /// directly at `api.moonshot.ai`. + Kimi, + /// GPT-5.6 models (Sol, Terra, Luna), which Codex drives with a narrower + /// core tool set than earlier GPT models: a shell, a file editor, and + /// `update_plan`, plus optional web search. The profile omits dedicated + /// file-read, discovery, and fetch tools. Selected per model rather than + /// per provider, so other models on the `openai` provider keep + /// [`Self::OpenAi`]. + Gpt56, +} + +impl AgentProfileKind { + #[must_use] + pub fn as_str(self) -> &'static str { + self.into() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn agent_profile_kind_round_trips_as_settings_strings() { + for kind in AgentProfileKind::VARIANTS { + let expected = kind.to_string(); + let json = serde_json::to_string(&kind).unwrap(); + assert_eq!(json, format!("\"{expected}\"")); + let parsed: AgentProfileKind = serde_json::from_str(&json).unwrap(); + assert_eq!(parsed, *kind); + assert_eq!(expected.parse::().unwrap(), *kind); + } + } + + #[test] + fn claude5_and_gpt56_use_their_catalog_spellings() { + assert_eq!(AgentProfileKind::Claude5.as_str(), "claude-5"); + assert_eq!(AgentProfileKind::Gpt56.as_str(), "gpt56"); + assert_eq!(AgentProfileKind::OpenAi.as_str(), "openai"); + } +} diff --git a/lib/foundation/fabro-types/src/billing.rs b/lib/foundation/fabro-types/src/billing.rs index d994df7ea..88b7b27ef 100644 --- a/lib/foundation/fabro-types/src/billing.rs +++ b/lib/foundation/fabro-types/src/billing.rs @@ -1,6 +1,466 @@ -pub use fabro_model::{ - AnthropicBillingFacts, AnthropicModelPricing, BilledModelUsage, BilledTokenCounts, - GeminiBillingFacts, GeminiModelPricing, GeminiStoragePricing, GeminiStorageSegment, - ModelBillingFacts, ModelBillingInput, ModelPricing, ModelPricingPolicy, ModelRef, ModelUsage, - OpenAiBillingFacts, OpenAiModelPricing, PricePerMTok, Speed, TokenCounts, UsdMicros, -}; +//! Billing rollup vocabulary. +//! +//! Per-response token usage and cost come from lithos: [`TokenCounts`] holds +//! the five disjoint buckets and [`CostSource`] says where a cost came from. +//! Fabro sums that usage across responses, stages, and runs. The types here +//! are those sums, plus [`ModelRef`], the identity a billed response is +//! grouped under. + +use lithos_llm::catalog::{ModelHandle, ModelId, ProviderId}; +pub use lithos_llm::types::{Cost, CostSource, Speed, TokenCounts}; +use serde::{Deserialize, Serialize}; + +use crate::controls; + +const USD_MICROS_PER_USD_F64: f64 = 1_000_000.0; + +#[allow( + clippy::cast_possible_truncation, + clippy::cast_precision_loss, + reason = "Billing rounds bounded finite floats into i64 counters by design." +)] +fn saturating_rounded_f64_to_i64(value: f64) -> i64 { + if !value.is_finite() { + return if value.is_sign_negative() { + i64::MIN + } else { + i64::MAX + }; + } + + if value <= i64::MIN as f64 { + i64::MIN + } else if value >= i64::MAX as f64 { + i64::MAX + } else { + value as i64 + } +} + +fn saturating_u64_to_i64(value: u64) -> i64 { + i64::try_from(value).unwrap_or(i64::MAX) +} + +fn saturating_i64_to_u64(value: i64) -> u64 { + u64::try_from(value).unwrap_or_default() +} + +/// A USD amount in micros (one millionth of a dollar). +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Default, Serialize, Deserialize)] +pub struct UsdMicros(pub i64); + +impl UsdMicros { + #[must_use] + pub fn from_usd(usd: f64) -> Self { + Self(saturating_rounded_f64_to_i64( + (usd * USD_MICROS_PER_USD_F64).round(), + )) + } + + /// Converts a lithos cost into Fabro's signed micros. + #[must_use] + pub fn from_cost(cost: &Cost) -> Self { + Self(saturating_u64_to_i64(cost.usd_micros)) + } + + /// Folds a cost into a running total that stays `None` until a cost is + /// observed (`None` means "no provider data", not $0). + pub fn accumulate(total: &mut Option, cost: Option) { + if let Some(cost) = cost { + *total.get_or_insert_default() += cost; + } + } +} + +impl std::ops::Add for UsdMicros { + type Output = Self; + + fn add(self, rhs: Self) -> Self::Output { + Self(self.0.saturating_add(rhs.0)) + } +} + +impl std::ops::AddAssign for UsdMicros { + fn add_assign(&mut self, rhs: Self) { + *self = *self + rhs; + } +} + +impl std::iter::Sum for UsdMicros { + fn sum>(iter: I) -> Self { + iter.fold(Self::default(), |acc, value| acc + value) + } +} + +/// Adds `rhs` into `total` bucket by bucket with saturation. +pub fn add_usage(total: &mut TokenCounts, rhs: TokenCounts) { + total.input = total.input.saturating_add(rhs.input); + total.output = total.output.saturating_add(rhs.output); + total.reasoning = total.reasoning.saturating_add(rhs.reasoning); + total.cache_read = total.cache_read.saturating_add(rhs.cache_read); + total.cache_write = total.cache_write.saturating_add(rhs.cache_write); +} + +fn accumulate_optional_usd_micros(total: &mut Option, cost: Option) { + let mut typed_total = (*total).map(UsdMicros); + UsdMicros::accumulate(&mut typed_total, cost.map(UsdMicros)); + *total = typed_total.map(|value| value.0); +} + +/// Provider-qualified model identity a billed response is grouped under. +/// +/// Carries the requested speed tier because providers price tiers +/// differently, so two responses from the same model at different speeds are +/// separate billing rows. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ModelRef { + pub provider: ProviderId, + pub model_id: ModelId, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub speed: Option, +} + +impl ModelRef { + #[must_use] + pub fn new(provider: ProviderId, model_id: ModelId) -> Self { + Self { + provider, + model_id, + speed: None, + } + } + + #[must_use] + pub fn from_handle(handle: &ModelHandle, speed: Option) -> Self { + Self { + provider: handle.provider().clone(), + model_id: handle.model().clone(), + speed, + } + } + + #[must_use] + pub fn with_speed(mut self, speed: Option) -> Self { + self.speed = speed; + self + } + + #[must_use] + pub fn handle(&self) -> ModelHandle { + ModelHandle::new(self.provider.clone(), self.model_id.clone()) + } + + /// Stable ordering key: provider, then model, then speed label. + #[must_use] + pub fn sort_key(&self) -> (&str, &str, &'static str) { + ( + self.provider.as_str(), + self.model_id.as_str(), + self.speed.map_or("", controls::speed_name), + ) + } +} + +impl std::hash::Hash for ModelRef { + fn hash(&self, state: &mut H) { + self.provider.hash(state); + self.model_id.hash(state); + self.speed.map(controls::speed_name).hash(state); + } +} + +impl std::fmt::Display for ModelRef { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{}/{}", self.provider, self.model_id)?; + if let Some(speed) = self.speed { + write!(f, " ({})", controls::speed_name(speed))?; + } + Ok(()) + } +} + +/// Usage and cost of one billed model response. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct BilledModelUsage { + pub model: ModelRef, + pub tokens: TokenCounts, + /// Cost for `tokens`, when the provider reported one or the catalog could + /// price them. `None` means no cost data, not zero. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub total_usd_micros: Option, +} + +impl BilledModelUsage { + #[must_use] + pub fn new(model: ModelRef, tokens: TokenCounts, cost: Option) -> Self { + Self { + model, + tokens, + total_usd_micros: cost.map(|cost| UsdMicros::from_cost(&cost).0), + } + } + + #[must_use] + pub fn model(&self) -> &ModelRef { + &self.model + } + + #[must_use] + pub fn model_id(&self) -> &str { + self.model.model_id.as_str() + } + + #[must_use] + pub fn tokens(&self) -> TokenCounts { + self.tokens + } + + /// Overrides the billed total with a reported cost; `None` leaves the + /// existing value in place. + #[must_use] + pub fn with_reported_cost(mut self, cost: Option) -> Self { + if let Some(cost) = cost { + self.total_usd_micros = Some(cost.0); + } + self + } +} + +/// Token counts summed across one or more responses, with the summed cost. +/// +/// `total_tokens` is the sum of the five buckets. `total_usd_micros` stays +/// `None` until at least one summed response carried a cost. +#[derive(Debug, Clone, PartialEq, Eq, Default, Serialize, Deserialize)] +pub struct BilledTokenCounts { + pub input_tokens: i64, + pub output_tokens: i64, + pub total_tokens: i64, + #[serde(default)] + pub reasoning_tokens: i64, + #[serde(default)] + pub cache_read_tokens: i64, + #[serde(default)] + pub cache_write_tokens: i64, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub total_usd_micros: Option, +} + +impl BilledTokenCounts { + #[must_use] + pub fn from_token_counts(tokens: TokenCounts, total_usd_micros: Option) -> Self { + Self { + input_tokens: saturating_u64_to_i64(tokens.input), + output_tokens: saturating_u64_to_i64(tokens.output), + total_tokens: saturating_u64_to_i64(tokens.total()), + reasoning_tokens: saturating_u64_to_i64(tokens.reasoning), + cache_read_tokens: saturating_u64_to_i64(tokens.cache_read), + cache_write_tokens: saturating_u64_to_i64(tokens.cache_write), + total_usd_micros, + } + } + + #[must_use] + pub fn from_billed_usage(billed: &[BilledModelUsage]) -> Self { + let mut counts = Self::default(); + for entry in billed { + counts.add_billed_usage(entry); + } + counts + } + + /// Returns the five disjoint per-call token buckets, dropping the derived + /// `total_tokens` sum and the optional `total_usd_micros` cost. + #[must_use] + pub fn token_counts(&self) -> TokenCounts { + TokenCounts { + input: saturating_i64_to_u64(self.input_tokens), + output: saturating_i64_to_u64(self.output_tokens), + reasoning: saturating_i64_to_u64(self.reasoning_tokens), + cache_read: saturating_i64_to_u64(self.cache_read_tokens), + cache_write: saturating_i64_to_u64(self.cache_write_tokens), + } + } + + pub fn add_counts(&mut self, source: &Self) { + self.input_tokens = self.input_tokens.saturating_add(source.input_tokens); + self.output_tokens = self.output_tokens.saturating_add(source.output_tokens); + self.total_tokens = self.total_tokens.saturating_add(source.total_tokens); + self.reasoning_tokens = self + .reasoning_tokens + .saturating_add(source.reasoning_tokens); + self.cache_read_tokens = self + .cache_read_tokens + .saturating_add(source.cache_read_tokens); + self.cache_write_tokens = self + .cache_write_tokens + .saturating_add(source.cache_write_tokens); + accumulate_optional_usd_micros(&mut self.total_usd_micros, source.total_usd_micros); + } + + pub fn add_billed_usage(&mut self, usage: &BilledModelUsage) { + self.add_counts(&Self::from_token_counts( + usage.tokens, + usage.total_usd_micros, + )); + } + + pub fn replace_with_billed_usage(&mut self, usage: &BilledModelUsage) { + *self = Self::from_billed_usage(std::slice::from_ref(usage)); + } + + /// Overrides the billed total with a reported cost; `None` leaves any + /// existing value in place. + #[must_use] + pub fn with_reported_cost(mut self, cost: Option) -> Self { + if let Some(cost) = cost { + self.total_usd_micros = Some(cost.0); + } + self + } + + #[must_use] + pub fn is_zero(&self) -> bool { + self.input_tokens == 0 + && self.output_tokens == 0 + && self.total_tokens == 0 + && self.reasoning_tokens == 0 + && self.cache_read_tokens == 0 + && self.cache_write_tokens == 0 + && self.total_usd_micros.unwrap_or(0) == 0 + } +} + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::*; + + fn tokens() -> TokenCounts { + TokenCounts { + input: 100, + output: 20, + reasoning: 5, + cache_read: 7, + cache_write: 3, + } + } + + fn model() -> ModelRef { + ModelRef::new( + ProviderId::new("anthropic"), + ModelId::new("claude-sonnet-5"), + ) + } + + #[test] + fn usd_micros_from_usd_rounds_to_nearest_micro() { + assert_eq!(UsdMicros::from_usd(0.012_345), UsdMicros(12_345)); + assert_eq!(UsdMicros::from_usd(1.0), UsdMicros(1_000_000)); + assert_eq!(UsdMicros::from_usd(f64::INFINITY), UsdMicros(i64::MAX)); + } + + #[test] + fn usd_micros_from_cost_saturates() { + let cost = Cost { + usd_micros: u64::MAX, + source: CostSource::Provider, + }; + assert_eq!(UsdMicros::from_cost(&cost), UsdMicros(i64::MAX)); + } + + #[test] + fn accumulate_stays_none_without_costs() { + let mut total = None; + UsdMicros::accumulate(&mut total, None); + assert_eq!(total, None); + UsdMicros::accumulate(&mut total, Some(UsdMicros(5))); + UsdMicros::accumulate(&mut total, None); + UsdMicros::accumulate(&mut total, Some(UsdMicros(7))); + assert_eq!(total, Some(UsdMicros(12))); + } + + #[test] + fn billed_token_counts_from_token_counts_sums_total() { + let counts = BilledTokenCounts::from_token_counts(tokens(), Some(42)); + assert_eq!(counts.input_tokens, 100); + assert_eq!(counts.output_tokens, 20); + assert_eq!(counts.reasoning_tokens, 5); + assert_eq!(counts.cache_read_tokens, 7); + assert_eq!(counts.cache_write_tokens, 3); + assert_eq!(counts.total_tokens, 135); + assert_eq!(counts.total_usd_micros, Some(42)); + assert_eq!(counts.token_counts(), tokens()); + } + + #[test] + fn billed_token_counts_sum_billed_usage_and_costs() { + let priced = BilledModelUsage::new( + model(), + tokens(), + Some(Cost { + usd_micros: 10, + source: CostSource::Catalog, + }), + ); + let unpriced = BilledModelUsage::new(model(), tokens(), None); + let counts = BilledTokenCounts::from_billed_usage(&[priced, unpriced]); + assert_eq!(counts.input_tokens, 200); + assert_eq!(counts.total_tokens, 270); + assert_eq!(counts.total_usd_micros, Some(10)); + } + + #[test] + fn billed_token_counts_without_costs_report_none() { + let counts = + BilledTokenCounts::from_billed_usage(&[BilledModelUsage::new(model(), tokens(), None)]); + assert_eq!(counts.total_usd_micros, None); + assert!(!counts.is_zero()); + assert!(BilledTokenCounts::default().is_zero()); + } + + #[test] + fn billed_model_usage_serializes_lithos_token_buckets() { + let usage = BilledModelUsage::new(model().with_speed(Some(Speed::Fast)), tokens(), None); + let value = serde_json::to_value(&usage).unwrap(); + assert_eq!( + value, + json!({ + "model": { + "provider": "anthropic", + "model_id": "claude-sonnet-5", + "speed": "fast", + }, + "tokens": { + "input": 100, + "output": 20, + "reasoning": 5, + "cache_read": 7, + "cache_write": 3, + }, + }) + ); + let back: BilledModelUsage = serde_json::from_value(value).unwrap(); + assert_eq!(back, usage); + } + + #[test] + fn model_ref_hash_distinguishes_speed_tiers() { + use std::collections::HashSet; + + let mut set = HashSet::new(); + set.insert(model()); + set.insert(model().with_speed(Some(Speed::Fast))); + set.insert(model().with_speed(Some(Speed::Fast))); + assert_eq!(set.len(), 2); + } + + #[test] + fn model_ref_display_names_the_route_and_speed() { + assert_eq!(model().to_string(), "anthropic/claude-sonnet-5"); + assert_eq!( + model().with_speed(Some(Speed::Fast)).to_string(), + "anthropic/claude-sonnet-5 (fast)" + ); + } +} diff --git a/lib/foundation/fabro-types/src/billing_rollup.rs b/lib/foundation/fabro-types/src/billing_rollup.rs index acdefa407..9ec5aee23 100644 --- a/lib/foundation/fabro-types/src/billing_rollup.rs +++ b/lib/foundation/fabro-types/src/billing_rollup.rs @@ -1,7 +1,5 @@ use std::collections::HashMap; -use fabro_model::Catalog; - use crate::{BilledTokenCounts, ModelRef, RunProjection, RunTiming, StageSummary, StageTiming}; #[derive(Debug, Clone, PartialEq)] @@ -119,10 +117,7 @@ impl ProjectionBillingRollup { } #[must_use] -pub fn billing_rollup_from_projection( - projection: &RunProjection, - catalog: Option<&Catalog>, -) -> ProjectionBillingRollup { +pub fn billing_rollup_from_projection(projection: &RunProjection) -> ProjectionBillingRollup { let mut stage_indices = HashMap::::new(); let mut stages = Vec::::new(); let mut by_model = HashMap::::new(); @@ -134,8 +129,7 @@ pub fn billing_rollup_from_projection( if projection.is_boundary_stage(stage_id.node_id()) { continue; } - let usage = stage.billed_usage(catalog); - let usage = usage.as_ref(); + let usage = &stage.usage; if stage.completion.is_none() && stage.timing.is_none() && usage.is_zero() { continue; } @@ -180,19 +174,7 @@ pub fn billing_rollup_from_projection( } let mut by_model = by_model.into_values().collect::>(); - by_model.sort_by(|left, right| { - let left_provider = left.model.provider.to_string(); - let right_provider = right.model.provider.to_string(); - left_provider - .cmp(&right_provider) - .then_with(|| left.model.model_id.cmp(&right.model.model_id)) - .then_with(|| { - left.model - .speed - .map(<&'static str>::from) - .cmp(&right.model.speed.map(<&'static str>::from)) - }) - }); + by_model.sort_by(|left, right| left.model.sort_key().cmp(&right.model.sort_key())); ProjectionBillingRollup { stages, diff --git a/lib/foundation/fabro-types/src/catalog_api.rs b/lib/foundation/fabro-types/src/catalog_api.rs new file mode 100644 index 000000000..f304dc0e7 --- /dev/null +++ b/lib/foundation/fabro-types/src/catalog_api.rs @@ -0,0 +1,95 @@ +//! API projections of the model catalog. +//! +//! `GET /models` and `GET /providers` return these. They are views over the +//! lithos catalog plus Fabro policy, stamped per request with whether the +//! server holds credentials for each provider. + +use serde::{Deserialize, Serialize}; + +use crate::{ModelId, ProviderId, ReasoningEffort}; + +/// Token limits for a model. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub struct ModelLimits { + pub context_window: i64, + pub max_output: Option, +} + +/// Capability flags for a model. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub struct ModelFeatures { + pub tools: bool, + pub vision: bool, + pub reasoning: bool, + pub prompt_cache: bool, + /// Whether the model accepts classic sampling parameters + /// (`temperature`, `top_p`). + pub sampling: bool, +} + +/// Request-control values a model accepts. +#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)] +pub struct ModelControls { + /// Reasoning-effort values accepted by this offering. Empty means the + /// control is unsupported. + #[serde(default)] + pub reasoning_effort: Vec, +} + +/// Pricing per million tokens in USD. +#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)] +pub struct ModelCosts { + pub input_cost_per_mtok: Option, + pub output_cost_per_mtok: Option, + pub cache_input_cost_per_mtok: Option, +} + +/// One provider's offering of a model. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct Model { + pub id: ModelId, + pub provider: ProviderId, + pub family: String, + pub display_name: String, + pub limits: ModelLimits, + pub training: Option, + pub knowledge_cutoff: Option, + pub features: ModelFeatures, + #[serde(default)] + pub controls: ModelControls, + pub costs: ModelCosts, + pub estimated_output_tps: Option, + pub aliases: Vec, + #[serde(default)] + pub default: bool, + #[serde(default)] + pub small_default: bool, + /// Whether the server holds credential material for this model's + /// provider. Stamped per request; never implies the credential works. + #[serde(default)] + pub configured: bool, +} + +/// An LLM provider with effective configuration and configured status. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct Provider { + pub id: ProviderId, + pub display_name: String, + /// lithos adapter id, such as `openai` or `openai-compatible`. + pub adapter: String, + pub base_url: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub api_key_url: Option, + pub priority: i32, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub aliases: Vec, + pub model_count: u32, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub default_model: Option, + #[serde(default)] + pub configured: bool, + /// Vault secret an operator creates to configure this provider, when the + /// provider reads one. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub expected_secret_name: Option, +} diff --git a/lib/foundation/fabro-types/src/catalog_policy.rs b/lib/foundation/fabro-types/src/catalog_policy.rs new file mode 100644 index 000000000..57f127fb9 --- /dev/null +++ b/lib/foundation/fabro-types/src/catalog_policy.rs @@ -0,0 +1,229 @@ +//! Fabro's `metadata.fabro` catalog namespace. +//! +//! lithos-llm owns provider and model facts. Fabro attaches its own policy to +//! each entry under `metadata.fabro`, which lithos carries verbatim and never +//! interprets. These types are the typed view of that namespace. Every field +//! is optional in the TOML; the accessors here apply Fabro's defaults. + +use lithos_llm::catalog::{CatalogModel, CatalogProvider}; +use serde::{Deserialize, Serialize}; + +use crate::AgentProfileKind; + +/// Name of the metadata namespace Fabro owns on catalog entries. +pub const FABRO_METADATA_NAMESPACE: &str = "fabro"; + +/// Provider-level Fabro policy. +#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)] +#[serde(default)] +pub struct ProviderPolicy { + /// Whether Fabro offers this provider at all. Missing means enabled. + pub enabled: Option, + /// Default agent profile for models on this provider. + pub agent_profile: Option, + /// Where an operator obtains an API key. + pub api_key_url: Option, + /// Ordered credential references (`env:NAME`, `vault:NAME`, `aws_sigv4`). + /// The first that resolves wins. + pub credentials: Vec, + /// Extra request headers. Values are literal text or `{{ secrets.NAME }}` + /// interpolation strings resolved against the vault. + #[serde(skip_serializing_if = "std::collections::BTreeMap::is_empty")] + pub extra_headers: std::collections::BTreeMap, + /// Another provider this one serves requests for when that provider has + /// no credentials of its own. Used by `openai-codex`, which answers + /// `openai` requests with a ChatGPT OAuth credential. + pub stands_in_for: Option, +} + +impl ProviderPolicy { + #[must_use] + pub fn is_enabled(&self) -> bool { + self.enabled.unwrap_or(true) + } +} + +/// Model-level Fabro policy. +#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)] +#[serde(default)] +pub struct ModelPolicy { + /// Whether Fabro offers this model. Missing means enabled. + pub enabled: Option, + /// Agent profile override for this model. + pub agent_profile: Option, + /// Model family label for display and grouping. + pub family: Option, + /// Training data cutoff label. + pub training: Option, + /// Public knowledge cutoff label. + pub knowledge_cutoff: Option, + /// Estimated output tokens per second. + pub estimated_output_tps: Option, + /// Preferred for small utility calls such as title generation. + pub small_default: bool, + /// Preferred for provider connectivity probes. + pub probe: bool, + /// Whether requests reason when no effort is requested. Missing means + /// "reasons when the model supports reasoning". + pub reasoning_by_default: Option, +} + +impl ModelPolicy { + #[must_use] + pub fn is_enabled(&self) -> bool { + self.enabled.unwrap_or(true) + } +} + +/// Reads a provider's Fabro policy. Malformed metadata falls back to the +/// defaults; the catalog build is the place to validate shape, and Fabro's +/// own policy file is checked in tests. +#[must_use] +pub fn provider_policy(provider: &CatalogProvider) -> ProviderPolicy { + provider + .metadata() + .namespace::(FABRO_METADATA_NAMESPACE) + .ok() + .flatten() + .unwrap_or_default() +} + +/// Reads a model's Fabro policy. +#[must_use] +pub fn model_policy(model: &CatalogModel) -> ModelPolicy { + model + .metadata() + .namespace::(FABRO_METADATA_NAMESPACE) + .ok() + .flatten() + .unwrap_or_default() +} + +/// The agent profile a model runs under: the model override, then the +/// provider default, then the profile implied by the provider's adapter. +#[must_use] +pub fn effective_agent_profile( + provider: &CatalogProvider, + model: &CatalogModel, +) -> AgentProfileKind { + model_policy(model) + .agent_profile + .or(provider_policy(provider).agent_profile) + .unwrap_or_else(|| default_agent_profile(provider)) +} + +/// The agent profile implied by a provider's wire protocol. +#[must_use] +pub fn default_agent_profile(provider: &CatalogProvider) -> AgentProfileKind { + match provider.adapter().as_str() { + "anthropic" | "bedrock" => AgentProfileKind::Anthropic, + "gemini" => AgentProfileKind::Gemini, + _ => AgentProfileKind::OpenAi, + } +} + +#[cfg(test)] +mod tests { + use lithos_llm::catalog::Catalog; + + use super::*; + + fn catalog() -> Catalog { + Catalog::builder() + .toml_layer( + "test", + r#" +schema_version = 1 + +[providers.acme] +display_name = "Acme" +adapter = "openai-compatible" +codec = "openai-chat" +base_url = "https://acme.test/v1" +auth = { type = "bearer" } +default_model = "large" + +[providers.acme.metadata.fabro] +enabled = false +credentials = ["env:ACME_API_KEY"] +agent_profile = "kimi" + +[providers.acme.models.large] +display_name = "Large" +api_model = "large" + +[providers.acme.models.large.metadata.fabro] +small_default = true +probe = true +family = "acme" + +[providers.acme.models.small] +display_name = "Small" +api_model = "small" +[providers.acme.models.small.metadata.fabro] +agent_profile = "openai" +enabled = false +"#, + ) + .unwrap() + .build() + .unwrap() + } + + #[test] + fn reads_provider_and_model_policy() { + let catalog = catalog(); + let provider = catalog.provider("acme").unwrap(); + let policy = provider_policy(provider); + assert!(!policy.is_enabled()); + assert_eq!(policy.credentials, vec!["env:ACME_API_KEY"]); + assert_eq!(policy.agent_profile, Some(AgentProfileKind::Kimi)); + + let large = provider.model("large").unwrap(); + let policy = model_policy(large); + assert!(policy.small_default && policy.probe && policy.is_enabled()); + assert_eq!(policy.family.as_deref(), Some("acme")); + assert_eq!( + effective_agent_profile(provider, large), + AgentProfileKind::Kimi + ); + + let small = provider.model("small").unwrap(); + assert!(!model_policy(small).is_enabled()); + assert_eq!( + effective_agent_profile(provider, small), + AgentProfileKind::OpenAi + ); + } + + #[test] + fn missing_namespace_yields_defaults() { + let catalog = Catalog::builder() + .toml_layer( + "test", + r#" +schema_version = 1 +[providers.bare] +display_name = "Bare" +adapter = "anthropic" +codec = "anthropic-messages" +base_url = "https://bare.test" +auth = { type = "none" } +[providers.bare.models.m] +display_name = "M" +api_model = "m" +"#, + ) + .unwrap() + .build() + .unwrap(); + let provider = catalog.provider("bare").unwrap(); + assert!(provider_policy(provider).is_enabled()); + let model = provider.model("m").unwrap(); + assert!(model_policy(model).is_enabled()); + assert_eq!( + effective_agent_profile(provider, model), + AgentProfileKind::Anthropic + ); + } +} diff --git a/lib/foundation/fabro-types/src/controls.rs b/lib/foundation/fabro-types/src/controls.rs new file mode 100644 index 000000000..9b3a54e6d --- /dev/null +++ b/lib/foundation/fabro-types/src/controls.rs @@ -0,0 +1,137 @@ +//! Helpers over the lithos request-control enums. +//! +//! lithos owns [`ReasoningEffort`] and [`Speed`] and marks both +//! `#[non_exhaustive]`. Fabro needs to list, name, and parse them for +//! settings, graph attributes, and CLI flags, so the spellings live here in +//! one place. The names match the lithos serde form. + +pub use lithos_llm::types::{ReasoningEffort, Speed}; + +/// Every reasoning effort, least to most. +pub const REASONING_EFFORTS: &[ReasoningEffort] = &[ + ReasoningEffort::Minimal, + ReasoningEffort::Low, + ReasoningEffort::Medium, + ReasoningEffort::High, + ReasoningEffort::Xhigh, + ReasoningEffort::Max, +]; + +/// Every speed tier. +pub const SPEEDS: &[Speed] = &[Speed::Fast, Speed::Balanced, Speed::Economical]; + +/// The wire spelling of a reasoning effort. +#[must_use] +pub fn reasoning_effort_name(effort: ReasoningEffort) -> &'static str { + match effort { + ReasoningEffort::Minimal => "minimal", + ReasoningEffort::Low => "low", + ReasoningEffort::Medium => "medium", + ReasoningEffort::High => "high", + ReasoningEffort::Xhigh => "xhigh", + ReasoningEffort::Max => "max", + _ => "unknown", + } +} + +/// The wire spelling of a speed tier. +#[must_use] +pub fn speed_name(speed: Speed) -> &'static str { + match speed { + Speed::Fast => "fast", + Speed::Balanced => "balanced", + Speed::Economical => "economical", + _ => "unknown", + } +} + +/// Parses a reasoning effort from its wire spelling. +#[must_use] +pub fn parse_reasoning_effort(value: &str) -> Option { + REASONING_EFFORTS + .iter() + .copied() + .find(|effort| reasoning_effort_name(*effort) == value) +} + +/// Parses a speed tier from its wire spelling. +#[must_use] +pub fn parse_speed(value: &str) -> Option { + SPEEDS + .iter() + .copied() + .find(|speed| speed_name(*speed) == value) +} + +/// Position of an effort in the least-to-most ordering. +fn effort_rank(effort: ReasoningEffort) -> usize { + REASONING_EFFORTS + .iter() + .position(|candidate| *candidate == effort) + .unwrap_or(REASONING_EFFORTS.len()) +} + +/// Selects the supported effort nearest to `requested`. +/// +/// When two supported values are equally distant, the higher effort wins. +/// Returns `None` when nothing is supported. +#[must_use] +pub fn closest_supported_effort( + requested: ReasoningEffort, + supported: impl Fn(ReasoningEffort) -> bool, +) -> Option { + let target = effort_rank(requested); + REASONING_EFFORTS + .iter() + .copied() + .filter(|effort| supported(*effort)) + .min_by_key(|effort| { + let rank = effort_rank(*effort); + (rank.abs_diff(target), std::cmp::Reverse(rank)) + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn names_round_trip_through_serde() { + for effort in REASONING_EFFORTS { + let json = serde_json::to_string(effort).unwrap(); + assert_eq!(json, format!("\"{}\"", reasoning_effort_name(*effort))); + assert_eq!( + parse_reasoning_effort(reasoning_effort_name(*effort)), + Some(*effort) + ); + } + for speed in SPEEDS { + let json = serde_json::to_string(speed).unwrap(); + assert_eq!(json, format!("\"{}\"", speed_name(*speed))); + assert_eq!(parse_speed(speed_name(*speed)), Some(*speed)); + } + assert_eq!(parse_reasoning_effort("standard"), None); + assert_eq!(parse_speed("standard"), None); + } + + #[test] + fn closest_supported_prefers_the_higher_neighbor_on_ties() { + let supported = |effort| matches!(effort, ReasoningEffort::Low | ReasoningEffort::High); + assert_eq!( + closest_supported_effort(ReasoningEffort::Medium, supported), + Some(ReasoningEffort::High) + ); + assert_eq!( + closest_supported_effort(ReasoningEffort::Max, supported), + Some(ReasoningEffort::High) + ); + assert_eq!( + closest_supported_effort(ReasoningEffort::Minimal, supported), + Some(ReasoningEffort::Low) + ); + assert_eq!( + closest_supported_effort(ReasoningEffort::Medium, |_| false), + None + ); + } +} diff --git a/lib/foundation/fabro-types/src/lib.rs b/lib/foundation/fabro-types/src/lib.rs index e5639e4ee..60fd05c9e 100644 --- a/lib/foundation/fabro-types/src/lib.rs +++ b/lib/foundation/fabro-types/src/lib.rs @@ -1,14 +1,18 @@ extern crate self as fabro_types; +pub mod agent_profile; pub mod artifact; pub mod auth; pub mod billing; pub mod billing_rollup; pub mod blob_hash; pub mod blob_ref; +pub mod catalog_api; +pub mod catalog_policy; pub mod checkpoint; pub mod command_output; pub mod conclusion; +pub mod controls; pub mod dense; pub mod diff; pub mod event_envelope; @@ -20,10 +24,12 @@ pub mod interview; pub mod llm_backend; pub mod manifest_path; pub mod mcp_store; +pub mod model_test; pub mod outcome; pub mod pair; pub mod parallel; pub mod principal; +pub mod provider_ids; pub mod pull_request; pub mod reasoning; pub mod repository; @@ -60,23 +66,23 @@ pub mod workflow_path; pub mod workflow_version; pub mod workflow_version_id; +pub use agent_profile::AgentProfileKind; pub use artifact::ArtifactUpload; pub use auth::{IdpIdentity, IdpIdentityError}; pub use billing::{ - AnthropicBillingFacts, AnthropicModelPricing, BilledModelUsage, BilledTokenCounts, - GeminiBillingFacts, GeminiModelPricing, GeminiStoragePricing, GeminiStorageSegment, - ModelBillingFacts, ModelBillingInput, ModelPricing, ModelPricingPolicy, ModelRef, ModelUsage, - OpenAiBillingFacts, OpenAiModelPricing, PricePerMTok, Speed, TokenCounts, UsdMicros, + BilledModelUsage, BilledTokenCounts, Cost, CostSource, ModelRef, Speed, TokenCounts, UsdMicros, }; pub use blob_hash::BlobHash; pub use blob_ref::{format_blob_ref, parse_blob_ref, parse_managed_blob_file_ref}; +pub use catalog_api::{Model, ModelControls, ModelCosts, ModelFeatures, ModelLimits, Provider}; +pub use catalog_policy::{ModelPolicy, ProviderPolicy}; pub use checkpoint::Checkpoint; pub use command_output::{CommandOutputStream, CommandTermination}; pub use conclusion::{Conclusion, StageSummary}; +pub use controls::ReasoningEffort; pub use dense::{ServerSettings, UserSettings, WorkflowSettings}; pub use diff::{DiffStats, DiffSummary, RunDiff}; pub use event_envelope::EventEnvelope; -pub use fabro_model::ReasoningEffort; pub use failure_signature::FailureSignature; pub use graph::{ AttrValue, AttributeScope, ContextKeyAttr, Edge, Graph, KNOWN_HANDLER_TYPES, Node, OnFailure, @@ -89,6 +95,10 @@ pub use input_scalar::{ pub use interview::{ InterviewQuestionRecord, QuestionType, ReviewTarget, ReviewTargetError, ReviewTargetKind, }; +pub use lithos_llm::catalog::{ModelHandle, ModelId, ProviderId}; +pub use lithos_llm::types::{ + FinishReason, Request, RequestBuildError, RequestBuilder, Response, ResponseFormat, StreamEvent, +}; pub use llm_backend::AgentBackend; pub use manifest_path::{ManifestPath, ManifestPathParseError}; pub use mcp_store::{ @@ -96,6 +106,7 @@ pub use mcp_store::{ McpServerRevisionParseError, McpServerValidationError, McpServerView, McpTransportView, validate_mcp_server_fields, }; +pub use model_test::ModelTestMode; pub use outcome::{ FailureCategory, FailureDetail, NodeResult, Outcome, OutcomeMeta, StageOutcome, StageState, }; @@ -192,8 +203,10 @@ pub use system_integrations::{ pub use timing::{RunTiming, StageTiming}; pub use todo::{TodoListKind, TodoListProjection, TodoPatch, TodoProjection, TodoStatus}; pub use transcript::{ - AudioData, ContentPart, DocumentData, ImageData, Message, MessageId, MessageKind, - MessageSource, PairMessageRef, Role, ThinkingData, ToolCall, ToolResult, TranscriptMessage, + AudioContent, ContentPart, DocumentContent, ImageContent, MediaSource, Message, MessageId, + MessageKind, MessageSource, PairMessageRef, ReasoningContent, Role, ToolCall, ToolCallKind, + ToolChoice, ToolDefinition, ToolDefinitionKind, ToolInput, ToolResult, TranscriptMessage, + text_of, tool_call_arguments, tool_result_from_json, tool_result_to_json, }; pub use variable::{ CreateVariableRequest, UpdateVariableRequest, Variable, VariableListResponse, is_env_style_name, diff --git a/lib/foundation/fabro-types/src/model_test.rs b/lib/foundation/fabro-types/src/model_test.rs new file mode 100644 index 000000000..610d51b91 --- /dev/null +++ b/lib/foundation/fabro-types/src/model_test.rs @@ -0,0 +1,35 @@ +//! Model probe modes exposed by `POST /models/{id}/test`. + +use serde::{Deserialize, Serialize}; +use strum::{Display, EnumString, IntoStaticStr}; + +#[derive( + Debug, + Clone, + Copy, + PartialEq, + Eq, + Default, + Serialize, + Deserialize, + Display, + EnumString, + IntoStaticStr, +)] +#[serde(rename_all = "lowercase")] +#[strum(serialize_all = "lowercase")] +pub enum ModelTestMode { + #[default] + Basic, + Deep, +} + +impl ModelTestMode { + #[must_use] + pub const fn timeout_secs(self) -> u64 { + match self { + Self::Basic => 30, + Self::Deep => 90, + } + } +} diff --git a/lib/foundation/fabro-types/src/provider_ids.rs b/lib/foundation/fabro-types/src/provider_ids.rs new file mode 100644 index 000000000..f58d0d7d6 --- /dev/null +++ b/lib/foundation/fabro-types/src/provider_ids.rs @@ -0,0 +1,27 @@ +//! Well-known provider identifiers. +//! +//! Provider identity is open-ended catalog data, so [`ProviderId`] is a plain +//! string newtype. The three first-party providers are named here because +//! code paths such as Codex login and the install flow refer to them +//! directly. + +use lithos_llm::catalog::ProviderId; + +pub const ANTHROPIC: &str = "anthropic"; +pub const OPENAI: &str = "openai"; +pub const GEMINI: &str = "gemini"; + +#[must_use] +pub fn anthropic() -> ProviderId { + ProviderId::new(ANTHROPIC) +} + +#[must_use] +pub fn openai() -> ProviderId { + ProviderId::new(OPENAI) +} + +#[must_use] +pub fn gemini() -> ProviderId { + ProviderId::new(GEMINI) +} diff --git a/lib/foundation/fabro-types/src/run_event/agent.rs b/lib/foundation/fabro-types/src/run_event/agent.rs index 171e69aa3..b06f21e08 100644 --- a/lib/foundation/fabro-types/src/run_event/agent.rs +++ b/lib/foundation/fabro-types/src/run_event/agent.rs @@ -1,4 +1,3 @@ -use fabro_model::{CostSource, ReasoningEffort, Speed}; use serde::{Deserialize, Serialize}; use serde_json::Value; use strum::{Display, EnumString, IntoStaticStr}; @@ -6,8 +5,9 @@ use strum::{Display, EnumString, IntoStaticStr}; use super::{BilledTokenCounts, ExecOutputTail}; use crate::transcript::{ToolCall, ToolResult, TranscriptMessage}; use crate::{ - CommandTermination, MessageId, ModelRef, PairId, PairMessageId, PairSystemMessageKind, - PermissionLevel, ReasoningOutput, StageContextWindowProjection, TurnId, + CommandTermination, CostSource, MessageId, ModelRef, PairId, PairMessageId, + PairSystemMessageKind, PermissionLevel, ReasoningEffort, ReasoningOutput, Speed, + StageContextWindowProjection, TurnId, }; #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] @@ -149,7 +149,7 @@ pub struct AgentToolStartedProps { pub tool_call_id: String, pub arguments: Value, pub visit: u32, - /// Canonical tool call payload. Carries `tool_type`, `raw_arguments`, and + /// Canonical tool call payload. Carries the typed input and /// `provider_metadata` (e.g. Gemini `thought_signature`) so tool actions /// can be replayed against the originating provider. #[serde(default, skip_serializing_if = "Option::is_none")] @@ -526,14 +526,13 @@ mod tests { use serde_json::json; use super::*; - use crate::transcript::{ContentPart, MessageKind, MessageSource, TranscriptMessage}; + use crate::provider_ids; + use crate::transcript::{ + ContentPart, MessageKind, MessageSource, TranscriptMessage, tool_result_from_json, + }; fn sample_model_ref() -> ModelRef { - ModelRef { - provider: fabro_model::ProviderId::openai(), - model_id: "gpt-5".into(), - speed: None, - } + ModelRef::new(provider_ids::openai(), "gpt-5".into()) } #[test] @@ -587,7 +586,9 @@ mod tests { #[test] fn agent_message_props_carries_canonical_transcript_message() { let msg = TranscriptMessage::new(MessageKind::Agent, MessageSource::ProviderAnswer, vec![ - ContentPart::text("ok"), + ContentPart::Text { + text: "ok".to_string(), + }, ]); let props = AgentMessageProps { text: "ok".to_string(), @@ -624,8 +625,9 @@ mod tests { #[test] fn agent_tool_started_props_carries_canonical_tool_call_and_linkage() { - let mut tc = ToolCall::new("call_1", "Bash", json!({"cmd": "ls"})); - tc.provider_metadata = Some(json!({"thought_signature": "sig"})); + let mut tc = ToolCall::function("call_1", "Bash", json!({"cmd": "ls"})); + tc.provider_metadata + .insert("gemini".to_string(), json!({"thought_signature": "sig"})); let parent = MessageId::new(); let turn = TurnId::new(); let props = AgentToolStartedProps { @@ -639,7 +641,7 @@ mod tests { }; let v = serde_json::to_value(&props).unwrap(); assert_eq!( - v["tool_call"]["provider_metadata"]["thought_signature"], + v["tool_call"]["provider_metadata"]["gemini"]["thought_signature"], "sig" ); assert_eq!(v["turn_id"], turn.to_string()); @@ -667,7 +669,7 @@ mod tests { #[test] fn agent_tool_completed_props_carries_canonical_tool_result() { - let tr = ToolResult::success("call_1", json!({"stdout": "ok"})); + let tr = tool_result_from_json("call_1", json!({"stdout": "ok"}), false); let turn = TurnId::new(); let props = AgentToolCompletedProps { tool_name: "Bash".to_string(), @@ -682,7 +684,7 @@ mod tests { turn_id: Some(turn), }; let v = serde_json::to_value(&props).unwrap(); - assert_eq!(v["tool_result"]["content"]["stdout"], "ok"); + assert_eq!(v["tool_result"]["content"][0]["value"]["stdout"], "ok"); assert_eq!(v["output_bytes_observed"], 120); assert_eq!(v["output_bytes_retained"], 100); assert_eq!(v["output_bytes_omitted"], 20); diff --git a/lib/foundation/fabro-types/src/run_event/misc.rs b/lib/foundation/fabro-types/src/run_event/misc.rs index 9e5600287..b9206ac33 100644 --- a/lib/foundation/fabro-types/src/run_event/misc.rs +++ b/lib/foundation/fabro-types/src/run_event/misc.rs @@ -1,10 +1,9 @@ -use fabro_model::ReasoningEffort; use serde::{Deserialize, Serialize}; use super::ExecOutputTail; use crate::{ - CommandTermination, ParallelBranchResult, PullRequestCreationId, PullRequestLink, ReviewTarget, - StageId, StageOutcome, + CommandTermination, ParallelBranchResult, PullRequestCreationId, PullRequestLink, + ReasoningEffort, ReviewTarget, StageId, StageOutcome, }; #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] diff --git a/lib/foundation/fabro-types/src/run_event/mod.rs b/lib/foundation/fabro-types/src/run_event/mod.rs index a0450dea6..fc2e6c9c5 100644 --- a/lib/foundation/fabro-types/src/run_event/mod.rs +++ b/lib/foundation/fabro-types/src/run_event/mod.rs @@ -8,7 +8,6 @@ pub mod todo; pub use agent::*; use chrono::{DateTime, Utc}; -pub use fabro_model::BilledTokenCounts; pub use infra::*; pub use misc::*; pub use run::*; @@ -20,7 +19,7 @@ pub use session::*; pub use stage::*; pub use todo::*; -use crate::{ParallelBranchId, Principal, RunId, StageId, UsdMicros}; +use crate::{BilledTokenCounts, ParallelBranchId, Principal, RunId, StageId}; /// Maximum accepted body size for `POST /runs/{id}/events`. /// @@ -920,8 +919,10 @@ impl RunEvent { } } -/// Upgrades historical wire shapes only in the value being decoded. Legacy -/// importers still retain and compare the original stored JSON. +/// Upgrades historical envelope shapes only in the value being decoded. +/// +/// Event bodies carry no compatibility rewrites: Fabro is greenfield, so a +/// stored body either matches the current schema or fails to decode. fn normalize_legacy_event(value: &mut Value) { let Some(event) = value .get("event") @@ -941,81 +942,17 @@ fn normalize_legacy_event_properties(event: &str, properties: &mut Value) { return; }; match event { - "agent.message" => normalize_legacy_agent_message(object), "run.completed" => normalize_legacy_timing(object, false), "run.failed" => { normalize_legacy_run_failure(object); normalize_legacy_timing(object, false); } - "stage.completed" => { - normalize_legacy_usage_field(object); - normalize_legacy_billing_field(object, "billing"); - normalize_legacy_timing(object, true); - } - "stage.failed" => normalize_legacy_billing_field(object, "billing"), - "prompt.completed" => { - normalize_legacy_usage_field(object); - normalize_legacy_billing_field(object, "billing"); - } - "checkpoint.completed" => normalize_legacy_checkpoint_billing(object), + "stage.completed" => normalize_legacy_timing(object, true), "sandbox.initialized" => normalize_legacy_sandbox_id(object), _ => {} } } -fn normalize_legacy_agent_message(properties: &mut Map) { - let speed = properties - .get("usage") - .and_then(Value::as_object) - .and_then(|usage| usage.get("speed")) - .and_then(Value::as_str) - .map(str::to_owned); - if let Some(model_id) = properties.get("model").and_then(Value::as_str) { - let mut model = Map::from_iter([ - ( - "provider".to_owned(), - Value::String(legacy_provider_for_model(model_id).to_owned()), - ), - ("model_id".to_owned(), Value::String(model_id.to_owned())), - ]); - if let Some(speed @ ("standard" | "fast")) = speed.as_deref() { - model.insert("speed".to_owned(), Value::String(speed.to_owned())); - } - properties.insert("model".to_owned(), Value::Object(model)); - } - if !properties.contains_key("billing") { - if let Some(usage) = properties.remove("usage") { - properties.insert("billing".to_owned(), usage); - } - } -} - -fn normalize_legacy_usage_field(properties: &mut Map) { - if !properties.contains_key("billing") { - if let Some(usage) = properties.remove("usage") { - properties.insert("billing".to_owned(), usage); - } - } -} - -fn normalize_legacy_billing_field(properties: &mut Map, field: &str) { - if let Some(billing) = properties.get_mut(field) { - normalize_legacy_billing_values(billing); - } -} - -fn normalize_legacy_checkpoint_billing(properties: &mut Map) { - let Some(outcomes) = properties - .get_mut("node_outcomes") - .and_then(Value::as_object_mut) - else { - return; - }; - for outcome in outcomes.values_mut().filter_map(Value::as_object_mut) { - normalize_legacy_billing_field(outcome, "usage"); - } -} - fn normalize_legacy_timing(properties: &mut Map, stage: bool) { if properties.contains_key("timing") { return; @@ -1076,135 +1013,6 @@ fn normalize_legacy_sandbox_id(properties: &mut Map) { properties.insert("id".to_owned(), Value::String(id.to_owned())); } -fn normalize_legacy_billing_values(value: &mut Value) { - if legacy_stage_usage(value) { - let legacy = std::mem::take(value); - *value = normalized_legacy_stage_usage(&legacy); - return; - } - match value { - Value::Array(values) => { - for value in values { - normalize_legacy_billing_values(value); - } - } - Value::Object(object) => { - if let Some(facts) = object.get_mut("facts").and_then(Value::as_object_mut) { - if !facts.contains_key("algorithm") { - let provider = facts - .get("provider") - .and_then(Value::as_str) - .map(str::to_owned); - if let Some(provider) = provider { - facts.remove("provider"); - facts.insert( - "algorithm".to_owned(), - Value::String(legacy_billing_algorithm(&provider).to_owned()), - ); - } - } - } - for value in object.values_mut() { - normalize_legacy_billing_values(value); - } - } - _ => {} - } -} - -fn legacy_stage_usage(value: &Value) -> bool { - let Some(object) = value.as_object() else { - return false; - }; - object.get("model").is_some_and(Value::is_string) - && object.get("input_tokens").is_some_and(Value::is_number) - && object.get("output_tokens").is_some_and(Value::is_number) -} - -fn normalized_legacy_stage_usage(legacy: &Value) -> Value { - let object = legacy - .as_object() - .expect("legacy stage usage was validated as an object"); - let model_id = object - .get("model") - .and_then(Value::as_str) - .expect("legacy stage usage was validated with a string model"); - let provider = legacy_provider_for_model(model_id); - let mut model = json!({ - "provider": provider, - "model_id": model_id, - }); - if let Some(speed @ ("standard" | "fast")) = object.get("speed").and_then(Value::as_str) { - model["speed"] = Value::String(speed.to_owned()); - } - let input_tokens = object - .get("input_tokens") - .and_then(Value::as_i64) - .unwrap_or(0); - let output_tokens = object - .get("output_tokens") - .and_then(Value::as_i64) - .unwrap_or(0); - let reasoning_tokens = object - .get("reasoning_tokens") - .and_then(Value::as_i64) - .unwrap_or(0); - let cache_read_tokens = object - .get("cache_read_tokens") - .and_then(Value::as_i64) - .unwrap_or(0); - let cache_write_tokens = object - .get("cache_write_tokens") - .and_then(Value::as_i64) - .unwrap_or(0); - let mut normalized = json!({ - "input": { - "usage": { - "model": model, - "tokens": { - "input_tokens": input_tokens, - "output_tokens": output_tokens, - "reasoning_tokens": reasoning_tokens, - "cache_read_tokens": cache_read_tokens, - "cache_write_tokens": cache_write_tokens, - } - }, - "facts": { - "algorithm": legacy_billing_algorithm(provider), - } - } - }); - if let Some(cost) = object.get("cost").and_then(Value::as_f64) { - normalized["total_usd_micros"] = Value::from(UsdMicros::from_usd(cost).0); - } - normalized -} - -fn legacy_provider_for_model(model_id: &str) -> &'static str { - if model_id.starts_with("claude-") { - "anthropic" - } else if model_id.starts_with("gemini-") { - "gemini" - } else if model_id.starts_with("gpt-") - || model_id.starts_with("chatgpt-") - || model_id.starts_with("o1") - || model_id.starts_with("o3") - || model_id.starts_with("o4") - { - "openai" - } else { - "legacy" - } -} - -fn legacy_billing_algorithm(provider: &str) -> &'static str { - match provider { - "anthropic" => "anthropic", - "gemini" => "gemini", - _ => "openai", - } -} - impl Serialize for RunEvent { fn serialize(&self, serializer: S) -> Result where @@ -1232,8 +1040,8 @@ mod tests { use super::*; use crate::{ - AuthMethod, BlobHash, CommandTermination, Edge, Graph, IdpIdentity, Node, PendingReason, - WorkflowSettings, fixtures, test_support, + AuthMethod, BlobHash, CommandTermination, Edge, Graph, IdpIdentity, ModelRef, Node, + PendingReason, WorkflowSettings, fixtures, provider_ids, test_support, }; fn user_principal(login: &str) -> Principal { @@ -1377,138 +1185,6 @@ mod tests { assert_eq!(props.settings.run, WorkflowSettings::default().run); } - #[test] - fn historical_agent_message_accepts_string_model() { - let line = stored_event( - "agent.message", - &json!({ - "text": "done", - "model": "gemini-3.1-pro-preview", - "billing": { - "input_tokens": 10, - "output_tokens": 5, - "total_tokens": 15 - }, - "tool_call_count": 0, - "visit": 1 - }), - ); - - let parsed = RunEvent::from_value(line).unwrap(); - let normalized = parsed.to_value().unwrap(); - - assert_eq!(normalized["properties"]["model"]["provider"], "gemini"); - assert_eq!( - normalized["properties"]["model"]["model_id"], - "gemini-3.1-pro-preview" - ); - } - - #[test] - fn historical_stage_usage_and_duration_are_upgraded() { - let line = stored_event( - "stage.completed", - &json!({ - "index": 0, - "duration_ms": 42, - "status": "succeeded", - "usage": { - "model": "claude-sonnet-4-6", - "input_tokens": 100, - "output_tokens": 20, - "cache_read_tokens": 7, - "cache_write_tokens": 3, - "reasoning_tokens": 2, - "speed": "fast", - "cost": 0.012_345 - }, - "attempt": 1, - "max_attempts": 1 - }), - ); - - let parsed = RunEvent::from_value(line).unwrap(); - let normalized = parsed.to_value().unwrap(); - let properties = &normalized["properties"]; - - assert_eq!(properties["timing"]["wall_time_ms"], 42); - assert_eq!( - properties["billing"]["input"]["facts"]["algorithm"], - "anthropic" - ); - assert_eq!( - properties["billing"]["input"]["usage"]["model"]["speed"], - "fast" - ); - assert_eq!(properties["billing"]["total_usd_micros"], 12_345); - assert!(properties.get("duration_ms").is_none()); - assert!(properties.get("usage").is_none()); - } - - #[test] - fn historical_billing_provider_tags_are_upgraded() { - let legacy_billing = json!({ - "input": { - "usage": { - "model": { - "provider": "anthropic", - "model_id": "claude-sonnet-4-6" - }, - "tokens": { - "input_tokens": 100, - "output_tokens": 20, - "reasoning_tokens": 0, - "cache_read_tokens": 7, - "cache_write_tokens": 3 - } - }, - "facts": { - "provider": "anthropic", - "cache_write_5m_tokens": 3, - "cache_write_1h_tokens": 0 - } - }, - "total_usd_micros": 123 - }); - let prompt = stored_event( - "prompt.completed", - &json!({ - "response": "done", - "model": "claude-sonnet-4-6", - "provider": "anthropic", - "billing": legacy_billing.clone() - }), - ); - let checkpoint = stored_event( - "checkpoint.completed", - &json!({ - "status": "succeeded", - "current_node": "build", - "node_outcomes": { - "build": { - "status": "succeeded", - "usage": legacy_billing - } - } - }), - ); - - let prompt = RunEvent::from_value(prompt).unwrap().to_value().unwrap(); - let checkpoint = RunEvent::from_value(checkpoint) - .unwrap() - .to_value() - .unwrap(); - - assert_eq!( - prompt["properties"]["billing"]["input"]["facts"]["algorithm"], - "anthropic" - ); - assert_eq!( - checkpoint["properties"]["node_outcomes"]["build"]["usage"]["input"]["facts"]["algorithm"], - "anthropic" - ); - } - #[test] fn historical_terminal_and_sandbox_events_are_upgraded() { let completed = stored_event( @@ -2688,11 +2364,7 @@ mod tests { fn agent_message_omits_context_window_when_absent() { let body = EventBody::AgentMessage(AgentMessageProps { text: "ok".to_string(), - model: crate::ModelRef { - provider: fabro_model::ProviderId::openai(), - model_id: "gpt-5.4".into(), - speed: None, - }, + model: ModelRef::new(provider_ids::openai(), "gpt-5.4".into()), billing: BilledTokenCounts::default(), cost_source: None, tool_call_count: 0, @@ -2719,11 +2391,7 @@ mod tests { fn agent_message_omits_reasoning_when_absent() { let body = EventBody::AgentMessage(AgentMessageProps { text: "ok".to_string(), - model: crate::ModelRef { - provider: fabro_model::ProviderId::openai(), - model_id: "gpt-5.4".into(), - speed: None, - }, + model: ModelRef::new(provider_ids::openai(), "gpt-5.4".into()), billing: BilledTokenCounts::default(), cost_source: None, tool_call_count: 0, @@ -2747,11 +2415,7 @@ mod tests { fn agent_message_carries_reasoning_through_canonical_json() { let body = EventBody::AgentMessage(AgentMessageProps { text: String::new(), - model: crate::ModelRef { - provider: fabro_model::ProviderId::openai(), - model_id: "gpt-5.4".into(), - speed: None, - }, + model: ModelRef::new(provider_ids::openai(), "gpt-5.4".into()), billing: BilledTokenCounts::default(), cost_source: None, tool_call_count: 1, @@ -2804,11 +2468,7 @@ mod tests { }; let body = EventBody::AgentMessage(AgentMessageProps { text: "ok".to_string(), - model: crate::ModelRef { - provider: fabro_model::ProviderId::openai(), - model_id: "gpt-5.4".into(), - speed: None, - }, + model: ModelRef::new(provider_ids::openai(), "gpt-5.4".into()), billing: BilledTokenCounts::default(), cost_source: None, tool_call_count: 0, diff --git a/lib/foundation/fabro-types/src/run_event/session.rs b/lib/foundation/fabro-types/src/run_event/session.rs index 02dd98ea1..6f50a88f2 100644 --- a/lib/foundation/fabro-types/src/run_event/session.rs +++ b/lib/foundation/fabro-types/src/run_event/session.rs @@ -1,8 +1,7 @@ -use fabro_model::ProviderId; use serde::{Deserialize, Serialize}; use serde_json::Value; -use crate::TurnId; +use crate::{ProviderId, TurnId}; #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct RunSessionCreatedProps { diff --git a/lib/foundation/fabro-types/src/run_event/stage.rs b/lib/foundation/fabro-types/src/run_event/stage.rs index 18a5319b0..cc6aa69ab 100644 --- a/lib/foundation/fabro-types/src/run_event/stage.rs +++ b/lib/foundation/fabro-types/src/run_event/stage.rs @@ -1,12 +1,12 @@ use std::collections::BTreeMap; -use fabro_model::{ReasoningEffort, Speed}; use serde::{Deserialize, Serialize}; use serde_json::Value; use super::ExecOutputTail; use crate::{ - BilledModelUsage, DiffSummary, FailureDetail, Outcome, StageId, StageOutcome, StageTiming, + BilledModelUsage, DiffSummary, FailureDetail, Outcome, ReasoningEffort, Speed, StageId, + StageOutcome, StageTiming, }; #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] diff --git a/lib/foundation/fabro-types/src/run_projection.rs b/lib/foundation/fabro-types/src/run_projection.rs index d41258298..b05d249e3 100644 --- a/lib/foundation/fabro-types/src/run_projection.rs +++ b/lib/foundation/fabro-types/src/run_projection.rs @@ -3,7 +3,6 @@ use std::collections::{BTreeMap, BTreeSet, HashMap}; use std::num::NonZeroU32; use chrono::{DateTime, Utc}; -use fabro_model::{Catalog, ReasoningEffort, Speed}; use strum::{Display, EnumString, IntoStaticStr}; use crate::run_event::{AgentSessionActivatedProps, StagePromptProps}; @@ -11,9 +10,9 @@ use crate::{ AgentBackend, AgentMcpToolSummary, AgentSkillActivationSource, AgentSkillSummary, AgentToolSummary, BilledTokenCounts, Checkpoint, Conclusion, InterviewQuestionRecord, InvalidTransition, LlmOutputKind, ModelRef, ParallelBranchId, PermissionLevel, - PullRequestCreation, PullRequestLink, RunApproval, RunControlAction, RunDiff, RunId, - RunSandbox, RunSpec, RunStatus, RunTiming, StageCompletion, StageHandler, StageId, StageState, - StageTiming, StartRecord, TodoListProjection, timing, + PullRequestCreation, PullRequestLink, ReasoningEffort, RunApproval, RunControlAction, RunDiff, + RunId, RunSandbox, RunSpec, RunStatus, RunTiming, Speed, StageCompletion, StageHandler, + StageId, StageState, StageTiming, StartRecord, TodoListProjection, timing, }; #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] @@ -605,29 +604,6 @@ impl StageProjection { self.state } - /// This stage's token counts with a cost attached. - /// - /// A provider-reported cost always wins. Otherwise the catalog prices the - /// recorded tokens for the stage's model. The stored counts pass through - /// untouched when there is no catalog, no model, or no price for that - /// model. Empty usage also passes through untouched. These cases leave - /// `total_usd_micros` as `None` rather than zero. - #[must_use] - pub fn billed_usage(&self, catalog: Option<&Catalog>) -> Cow<'_, BilledTokenCounts> { - if self.usage.total_usd_micros.is_some() || self.usage.is_zero() { - return Cow::Borrowed(&self.usage); - } - let (Some(catalog), Some(model)) = (catalog, self.model.as_ref()) else { - return Cow::Borrowed(&self.usage); - }; - let Some(total_usd_micros) = catalog.price_tokens(model, &self.usage.token_counts()) else { - return Cow::Borrowed(&self.usage); - }; - let mut usage = self.usage.clone(); - usage.total_usd_micros = Some(total_usd_micros); - Cow::Owned(usage) - } - /// Live wall-clock time in milliseconds. /// /// While the stage is non-terminal (`Pending`, `Running`, or `Retrying`), @@ -1119,11 +1095,10 @@ mod iter_stages_tests { use std::num::NonZeroU32; use chrono::Utc; - use fabro_model::{Catalog, ModelRef, ProviderId}; use serde_json::json; use super::RunProjection; - use crate::{AgentControlState, BilledTokenCounts, StageProjection, test_support}; + use crate::{AgentControlState, StageProjection, test_support}; fn seq(n: u32) -> NonZeroU32 { NonZeroU32::new(n).unwrap() @@ -1230,77 +1205,6 @@ mod iter_stages_tests { assert_eq!(order, vec!["build@1", "verify@1", "verify@2"]); } } - - fn priced_stage(total_usd_micros: Option) -> StageProjection { - let mut stage = StageProjection::new(seq(1)); - stage.usage = BilledTokenCounts { - input_tokens: 500_000, - output_tokens: 125_000, - total_tokens: 625_000, - total_usd_micros, - ..BilledTokenCounts::default() - }; - stage.model = Some(ModelRef { - provider: ProviderId::openai(), - model_id: "gpt-5.4".into(), - speed: None, - }); - stage - } - - #[test] - fn billed_usage_prices_uncosted_tokens_from_the_catalog() { - let stage = priced_stage(None); - - assert_eq!(stage.billed_usage(None).total_usd_micros, None); - let priced = stage.billed_usage(Some(Catalog::builtin())); - assert!( - priced.total_usd_micros.is_some_and(|cost| cost > 0), - "expected a catalog price, got {:?}", - priced.total_usd_micros - ); - // Pricing only fills in the cost; the token buckets pass through. - assert_eq!(priced.input_tokens, 500_000); - assert_eq!(priced.output_tokens, 125_000); - } - - #[test] - fn billed_usage_keeps_a_provider_reported_cost_over_the_catalog_estimate() { - let stage = priced_stage(Some(42)); - - assert_eq!( - stage - .billed_usage(Some(Catalog::builtin())) - .total_usd_micros, - Some(42) - ); - } - - #[test] - fn billed_usage_leaves_a_modelless_stage_uncosted() { - let mut stage = priced_stage(None); - stage.model = None; - - assert_eq!( - stage - .billed_usage(Some(Catalog::builtin())) - .total_usd_micros, - None - ); - } - - #[test] - fn billed_usage_leaves_zero_tokens_uncosted() { - let mut stage = priced_stage(None); - stage.usage = BilledTokenCounts::default(); - - assert_eq!( - stage - .billed_usage(Some(Catalog::builtin())) - .total_usd_micros, - None - ); - } } #[cfg(test)] @@ -1334,11 +1238,7 @@ mod live_timing_tests { StageInferenceProjection { session_id: "session-1".to_string(), started_at, - requested_model: ModelRef { - provider: "anthropic".parse().unwrap(), - model_id: "claude-sonnet-5".into(), - speed: None, - }, + requested_model: ModelRef::new("anthropic".into(), "claude-sonnet-5".into()), first_output_at: None, first_output_kind: None, retries: 0, diff --git a/lib/foundation/fabro-types/src/session.rs b/lib/foundation/fabro-types/src/session.rs index b825338d9..5387a673a 100644 --- a/lib/foundation/fabro-types/src/session.rs +++ b/lib/foundation/fabro-types/src/session.rs @@ -1,5 +1,5 @@ use chrono::{DateTime, Utc}; -use fabro_model::ProviderId; +use lithos_llm::catalog::ProviderId; use serde::{Deserialize, Serialize}; use strum::{Display, EnumString, IntoStaticStr}; diff --git a/lib/foundation/fabro-types/src/settings/model_ref.rs b/lib/foundation/fabro-types/src/settings/model_ref.rs index b551684b7..941fda119 100644 --- a/lib/foundation/fabro-types/src/settings/model_ref.rs +++ b/lib/foundation/fabro-types/src/settings/model_ref.rs @@ -24,6 +24,7 @@ use std::fmt; use std::str::FromStr; +use lithos_llm::catalog::Catalog; use serde::de::{self, Visitor}; use serde::{Deserialize, Deserializer, Serialize, Serializer}; @@ -228,14 +229,14 @@ impl ModelRef { } } -impl ModelRegistry for fabro_model::Catalog { +impl ModelRegistry for Catalog { fn is_provider(&self, token: &str) -> bool { - self.provider(&fabro_model::ProviderId::from(token)) - .is_some() + self.provider(token).is_ok() } fn is_model(&self, token: &str) -> bool { - self.is_model_selector(token) + self.providers() + .any(|provider| provider.model(token).is_some()) } } diff --git a/lib/foundation/fabro-types/src/transcript.rs b/lib/foundation/fabro-types/src/transcript.rs index fce181dba..32efda94f 100644 --- a/lib/foundation/fabro-types/src/transcript.rs +++ b/lib/foundation/fabro-types/src/transcript.rs @@ -1,16 +1,21 @@ -//! Canonical provider-neutral transcript primitives. +//! Canonical transcript primitives. //! -//! These types are the durable replay shapes for agent sessions. They were -//! promoted from `fabro-llm` so the Fabro event stream, API responses, and -//! runtime history can share one canonical Rust model rather than ferrying -//! parallel DTOs between layers. `fabro-llm::types` re-exports these so -//! existing imports keep working. +//! The message vocabulary (`Message`, `ContentPart`, `ToolCall`, `ToolResult`, +//! and friends) is lithos's, re-exported here so the event stream, API +//! responses, and runtime history share one Rust model. [`TranscriptMessage`] +//! is Fabro's durable replay record: identity, provenance, and usage wrapped +//! around lithos content parts. use chrono::{DateTime, Utc}; -use fabro_model::{ModelRef, TokenCounts}; -use serde::{Deserialize, Serialize, de}; +pub use lithos_llm::types::{ + AudioContent, ContentPart, DocumentContent, ImageContent, MediaSource, Message, + ReasoningContent, Role, TokenCounts, ToolArgumentError, ToolArguments, ToolCall, ToolCallKind, + ToolChoice, ToolDefinition, ToolDefinitionKind, ToolInput, ToolResult, UnknownContent, +}; +use serde::{Deserialize, Serialize}; use strum::{Display, EnumString, IntoStaticStr}; +use crate::billing::ModelRef; use crate::id::ulid_id; use crate::pair::{PairId, PairMessageId}; use crate::principal::Principal; @@ -18,336 +23,68 @@ use crate::session::TurnId; ulid_id!(MessageId); -// --- Content data structures ------------------------------------------------- - -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -pub struct ImageData { - pub url: Option, - pub data: Option>, - pub media_type: Option, - pub detail: Option, +/// Concatenates the text parts of a message or response. +#[must_use] +pub fn text_of(parts: &[ContentPart]) -> String { + parts + .iter() + .filter_map(|part| match part { + ContentPart::Text { text } => Some(text.as_str()), + _ => None, + }) + .collect() } -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -pub struct AudioData { - pub url: Option, - pub data: Option>, - pub media_type: Option, -} - -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -pub struct DocumentData { - pub url: Option, - pub data: Option>, - pub media_type: Option, - pub file_name: Option, -} - -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -pub struct ThinkingData { - pub text: String, - pub signature: Option, - pub redacted: bool, -} - -// --- Tool call / tool result ------------------------------------------------- - -fn default_tool_type() -> String { - "function".to_string() -} - -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -pub struct ToolCall { - pub id: String, - pub name: String, - #[serde(rename = "type", default = "default_tool_type")] - pub tool_type: String, - pub arguments: serde_json::Value, - pub raw_arguments: Option, - /// Opaque provider-specific metadata (e.g. Gemini `thought_signature`). - /// Preserved across round-trips so the provider can include it when - /// sending conversation history back to the API. - #[serde(skip_serializing_if = "Option::is_none")] - pub provider_metadata: Option, -} - -impl ToolCall { - pub fn new( - id: impl Into, - name: impl Into, - arguments: serde_json::Value, - ) -> Self { - Self { - id: id.into(), - name: name.into(), - tool_type: "function".to_string(), - arguments, - raw_arguments: None, - provider_metadata: None, - } - } -} - -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -pub struct ToolResult { - pub tool_call_id: String, - pub content: serde_json::Value, - pub is_error: bool, - #[serde(skip_serializing_if = "Option::is_none")] - pub image_data: Option>, - #[serde(skip_serializing_if = "Option::is_none")] - pub image_media_type: Option, -} - -impl ToolResult { - pub fn success(id: impl Into, content: serde_json::Value) -> Self { - Self { - tool_call_id: id.into(), - content, - is_error: false, - image_data: None, - image_media_type: None, - } - } - - pub fn error(id: impl Into, message: impl Into) -> Self { - Self { - tool_call_id: id.into(), - content: serde_json::Value::String(message.into()), - is_error: true, - image_data: None, - image_media_type: None, - } - } -} - -// --- ContentPart ------------------------------------------------------------- - -#[derive(Debug, Clone, PartialEq, Eq)] -pub enum ContentPart { - Text(String), - Image(ImageData), - Audio(AudioData), - Document(DocumentData), - ToolCall(ToolCall), - ToolResult(ToolResult), - Thinking(ThinkingData), - Other { - kind: String, - data: serde_json::Value, - }, -} - -impl Serialize for ContentPart { - fn serialize(&self, serializer: S) -> Result { - use serde::ser::SerializeMap; - let mut map = serializer.serialize_map(Some(2))?; - match self { - Self::Text(v) => { - map.serialize_entry("kind", "text")?; - map.serialize_entry("data", v)?; - } - Self::Image(v) => { - map.serialize_entry("kind", "image")?; - map.serialize_entry("data", v)?; - } - Self::Audio(v) => { - map.serialize_entry("kind", "audio")?; - map.serialize_entry("data", v)?; - } - Self::Document(v) => { - map.serialize_entry("kind", "document")?; - map.serialize_entry("data", v)?; - } - Self::ToolCall(v) => { - map.serialize_entry("kind", "tool_call")?; - map.serialize_entry("data", v)?; - } - Self::ToolResult(v) => { - map.serialize_entry("kind", "tool_result")?; - map.serialize_entry("data", v)?; - } - Self::Thinking(v) => { - let kind = if v.redacted { - "redacted_thinking" - } else { - "thinking" - }; - map.serialize_entry("kind", kind)?; - map.serialize_entry("data", v)?; - } - Self::Other { kind, data } => { - map.serialize_entry("kind", kind)?; - map.serialize_entry("data", data)?; - } - } - map.end() - } -} - -impl<'de> Deserialize<'de> for ContentPart { - fn deserialize>(deserializer: D) -> Result { - let value = serde_json::Value::deserialize(deserializer)?; - let kind = value - .get("kind") - .and_then(serde_json::Value::as_str) - .ok_or_else(|| de::Error::missing_field("kind"))?; - let data = value - .get("data") - .cloned() - .unwrap_or(serde_json::Value::Null); - match kind { - "text" => serde_json::from_value(data) - .map(Self::Text) - .map_err(de::Error::custom), - "image" => serde_json::from_value(data) - .map(Self::Image) - .map_err(de::Error::custom), - "audio" => serde_json::from_value(data) - .map(Self::Audio) - .map_err(de::Error::custom), - "document" => serde_json::from_value(data) - .map(Self::Document) - .map_err(de::Error::custom), - "tool_call" => serde_json::from_value(data) - .map(Self::ToolCall) - .map_err(de::Error::custom), - "tool_result" => serde_json::from_value(data) - .map(Self::ToolResult) - .map_err(de::Error::custom), - "thinking" => serde_json::from_value(data) - .map(Self::Thinking) - .map_err(de::Error::custom), - "redacted_thinking" => serde_json::from_value::(data) - .map(|mut td| { - td.redacted = true; - Self::Thinking(td) - }) - .map_err(de::Error::custom), - other => Ok(Self::Other { - kind: other.to_string(), - data, - }), - } - } -} - -impl ContentPart { - /// Kind string for opaque OpenAI reasoning output items. - pub const OPENAI_REASONING: &str = "openai_reasoning"; - /// Kind string for opaque OpenAI message output items. - pub const OPENAI_MESSAGE: &str = "openai_message"; - /// Kind string for opaque OpenAI-compatible `reasoning_details` entries. - /// The data is the received array of detail objects, preserved verbatim - /// so encrypted entries survive for future provider-aware replay. Only - /// known readable members are ever normalized out of it. - pub const OPENAI_COMPAT_REASONING_DETAILS: &str = "openai_compat_reasoning_details"; - - pub fn text(text: impl Into) -> Self { - Self::Text(text.into()) - } - - /// Returns `true` if this is an opaque OpenAI item (reasoning or message) - /// that should be round-tripped verbatim through the API. - pub fn is_opaque_openai(&self) -> bool { - matches!( - self, - Self::Other { kind, .. } - if kind == Self::OPENAI_REASONING || kind == Self::OPENAI_MESSAGE - ) - } -} - -// --- Role / Message -// ----------------------------------------------------------- - -/// Author role of a chat [`Message`]. -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] -#[serde(rename_all = "snake_case")] -pub enum Role { - System, - User, - Assistant, - Tool, - Developer, -} - -/// Provider-neutral chat message exchanged with an LLM. +/// Builds a tool result whose content is one JSON value. /// -/// This is the request/response message shape shared by `fabro-llm` -/// requests and the completions API wire contract. The durable -/// session-transcript record is [`TranscriptMessage`], which carries -/// identity, provenance, and usage on top of the same [`ContentPart`] -/// vocabulary. -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -pub struct Message { - pub role: Role, - pub content: Vec, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub name: Option, - #[serde(default, skip_serializing_if = "Option::is_none")] - pub tool_call_id: Option, +/// Plain strings become a text part so providers render them as text; every +/// other value is carried as structured JSON. +#[must_use] +pub fn tool_result_from_json( + tool_call_id: impl Into, + content: serde_json::Value, + is_error: bool, +) -> ToolResult { + let part = match content { + serde_json::Value::String(text) => ContentPart::Text { text }, + value => ContentPart::Json { value }, + }; + ToolResult { + tool_call_id: tool_call_id.into(), + name: None, + content: vec![part], + is_error, + } } -impl Message { - pub fn system(text: impl Into) -> Self { - Self { - role: Role::System, - content: vec![ContentPart::text(text)], - name: None, - tool_call_id: None, - } +/// Projects a tool result back to one JSON value, the inverse of +/// [`tool_result_from_json`]. +/// +/// A lone text part becomes a string and a lone JSON part its value. Any +/// other shape is carried as the array of serialized parts. +#[must_use] +pub fn tool_result_to_json(result: &ToolResult) -> serde_json::Value { + match result.content.as_slice() { + [ContentPart::Text { text }] => serde_json::Value::String(text.clone()), + [ContentPart::Json { value }] => value.clone(), + parts => serde_json::Value::Array( + parts + .iter() + .map(|part| serde_json::to_value(part).unwrap_or(serde_json::Value::Null)) + .collect(), + ), } +} - pub fn user(text: impl Into) -> Self { - Self { - role: Role::User, - content: vec![ContentPart::text(text)], - name: None, - tool_call_id: None, - } - } - - pub fn assistant(text: impl Into) -> Self { - Self { - role: Role::Assistant, - content: vec![ContentPart::text(text)], - name: None, - tool_call_id: None, - } - } - - pub fn tool_result( - tool_call_id: impl Into, - content: serde_json::Value, - is_error: bool, - ) -> Self { - let id = tool_call_id.into(); - Self { - role: Role::Tool, - content: vec![ContentPart::ToolResult(ToolResult { - tool_call_id: id.clone(), - content, - is_error, - image_data: None, - image_media_type: None, - })], - name: None, - tool_call_id: Some(id), - } - } - - /// Concatenates text from all text content parts. - #[must_use] - pub fn text(&self) -> String { - self.content - .iter() - .filter_map(|part| match part { - ContentPart::Text(text) => Some(text.as_str()), - _ => None, - }) - .collect() - } +/// The arguments of a tool call as one JSON value. +/// +/// Function arguments are the parsed JSON object; malformed arguments and +/// custom free-form input are carried as their raw text. +#[must_use] +pub fn tool_call_arguments(call: &ToolCall) -> serde_json::Value { + call.input + .to_value() + .unwrap_or_else(|_| serde_json::Value::String(call.input.raw().to_string())) } // --- TranscriptMessage ------------------------------------------------------ @@ -423,7 +160,7 @@ pub struct PairMessageRef { /// Canonical durable transcript message. /// /// Named `TranscriptMessage` rather than `Message` to avoid import ambiguity -/// with `fabro_agent::Message` and `fabro_llm::types::Message`. +/// with `fabro_agent::Message` and the lithos request [`Message`]. /// /// `kind` captures provider/model-role semantics for replay; `source` /// captures audit/UI provenance. Both are required to faithfully reconstruct @@ -480,59 +217,30 @@ mod tests { use super::*; #[test] - fn content_part_text_roundtrips() { - let part = ContentPart::text("hello"); - let v = serde_json::to_value(&part).unwrap(); - assert_eq!(v, json!({"kind": "text", "data": "hello"})); - let back: ContentPart = serde_json::from_value(v).unwrap(); - assert_eq!(back, part); + fn text_of_concatenates_text_parts_only() { + let parts = vec![ + ContentPart::Text { + text: "hello ".to_string(), + }, + ContentPart::Json { value: json!(1) }, + ContentPart::Text { + text: "world".to_string(), + }, + ]; + assert_eq!(text_of(&parts), "hello world"); } #[test] - fn content_part_thinking_preserves_signature_and_redaction() { - let part = ContentPart::Thinking(ThinkingData { - text: "private thought".to_string(), - signature: Some("sig_abc".to_string()), - redacted: true, - }); - let v = serde_json::to_value(&part).unwrap(); - assert_eq!(v["kind"], "redacted_thinking"); - assert_eq!(v["data"]["signature"], "sig_abc"); - let back: ContentPart = serde_json::from_value(v).unwrap(); - assert_eq!(back, part); - } - - #[test] - fn content_part_other_preserves_provider_kind() { - let part = ContentPart::Other { - kind: ContentPart::OPENAI_REASONING.to_string(), - data: json!({"item_id": "rs_1", "encrypted": "x"}), - }; - assert!(part.is_opaque_openai()); - let v = serde_json::to_value(&part).unwrap(); - let back: ContentPart = serde_json::from_value(v).unwrap(); - assert_eq!(back, part); - } - - #[test] - fn tool_call_preserves_provider_metadata() { - let mut tc = ToolCall::new("call_1", "Bash", json!({"cmd": "ls"})); - tc.provider_metadata = Some(json!({"thought_signature": "sig"})); - tc.raw_arguments = Some("{\"cmd\":\"ls\"}".to_string()); - let v = serde_json::to_value(&tc).unwrap(); - assert_eq!(v["provider_metadata"]["thought_signature"], "sig"); - let back: ToolCall = serde_json::from_value(v).unwrap(); - assert_eq!(back, tc); - } - - #[test] - fn tool_result_round_trips_with_default_image_fields() { - let tr = ToolResult::success("call_1", json!({"ok": true})); - let v = serde_json::to_value(&tr).unwrap(); - // Optional image fields are omitted on serialize. - assert!(v.get("image_data").is_none()); - let back: ToolResult = serde_json::from_value(v).unwrap(); - assert_eq!(back, tr); + fn tool_result_from_json_keeps_strings_as_text() { + let result = tool_result_from_json("call_1", json!("ok"), false); + assert_eq!(result.content, vec![ContentPart::Text { + text: "ok".to_string(), + }]); + let result = tool_result_from_json("call_1", json!({"ok": true}), true); + assert!(result.is_error); + assert_eq!(result.content, vec![ContentPart::Json { + value: json!({"ok": true}), + }]); } #[test] @@ -544,7 +252,9 @@ mod tests { source: MessageSource::Steer, actor: None, pair: None, - content: vec![ContentPart::text("please continue")], + content: vec![ContentPart::Text { + text: "please continue".to_string(), + }], model: None, response_id: None, usage: None, @@ -553,6 +263,10 @@ mod tests { let v = serde_json::to_value(&msg).unwrap(); assert_eq!(v["kind"], "user"); assert_eq!(v["source"], "steer"); + assert_eq!( + v["content"][0], + json!({"type": "text", "text": "please continue"}) + ); let back: TranscriptMessage = serde_json::from_value(v).unwrap(); assert_eq!(back, msg); } @@ -560,7 +274,9 @@ mod tests { #[test] fn transcript_message_drops_optional_fields_on_serialize() { let msg = TranscriptMessage::new(MessageKind::Agent, MessageSource::ProviderAnswer, vec![ - ContentPart::text("done"), + ContentPart::Text { + text: "done".to_string(), + }, ]); let v = serde_json::to_value(&msg).unwrap(); let obj = v.as_object().unwrap(); @@ -574,6 +290,22 @@ mod tests { assert!(!obj.contains_key("created_at")); } + #[test] + fn transcript_message_usage_uses_lithos_buckets() { + let mut msg = + TranscriptMessage::new(MessageKind::Agent, MessageSource::ProviderAnswer, vec![]); + msg.usage = Some(TokenCounts { + input: 10, + output: 2, + ..TokenCounts::default() + }); + let v = serde_json::to_value(&msg).unwrap(); + assert_eq!( + v["usage"], + json!({"input": 10, "output": 2, "reasoning": 0, "cache_read": 0, "cache_write": 0}) + ); + } + #[test] fn pair_message_ref_skips_empty_client_id() { let r = PairMessageRef {