diff --git a/bindings/python/tests/test_schema_sync.py b/bindings/python/tests/test_schema_sync.py index 23882878..f25df8a9 100644 --- a/bindings/python/tests/test_schema_sync.py +++ b/bindings/python/tests/test_schema_sync.py @@ -148,3 +148,22 @@ def test_hook_result_json_roundtrip(): assert restored.context_injection == "Lint error on line 42" assert restored.suppress_output is True assert restored.user_message == "Found 1 issue" + + +def test_model_info_pricing_field(): + """ModelInfo.pricing and Pricing are importable and round-trip via JSON.""" + from amplifier_core import ModelInfo, Pricing + + info = ModelInfo( + id="test-model", + display_name="Test", + context_window=1000, + max_output_tokens=100, + pricing=Pricing(input_per_million=3.0, output_per_million=15.0), + ) + json_str = info.model_dump_json() + parsed = json.loads(json_str) + assert parsed["pricing"]["input_per_million"] == 3.0 + + restored = ModelInfo.model_validate(parsed) + assert restored.pricing.output_per_million == 15.0 diff --git a/crates/amplifier-core/src/bridges/grpc_provider.rs b/crates/amplifier-core/src/bridges/grpc_provider.rs index 609dff93..7f51c98f 100644 --- a/crates/amplifier-core/src/bridges/grpc_provider.rs +++ b/crates/amplifier-core/src/bridges/grpc_provider.rs @@ -47,6 +47,23 @@ fn parse_defaults_json(json_str: &str, id: &str) -> HashMap { }) } +/// Parse a JSON string into an optional Pricing, logging a warning on +/// non-empty parse failures. Empty string means "no pricing available". +fn parse_pricing_json(json_str: &str, id: &str) -> Option { + if json_str.is_empty() { + return None; + } + serde_json::from_str(json_str) + .map_err(|e| { + log::warn!( + "Failed to parse model '{}' pricing_json: {e} — pricing unavailable", + id + ); + e + }) + .ok() +} + /// A bridge that wraps a remote gRPC `ProviderService` as a native [`Provider`]. /// /// The client is held behind a [`tokio::sync::Mutex`] because @@ -125,6 +142,7 @@ impl Provider for GrpcProviderBridge { .into_iter() .map(|m| { let defaults = parse_defaults_json(&m.defaults_json, &m.id); + let pricing = parse_pricing_json(&m.pricing_json, &m.id); ModelInfo { id: m.id, display_name: m.display_name, @@ -132,6 +150,7 @@ impl Provider for GrpcProviderBridge { max_output_tokens: m.max_output_tokens as i64, capabilities: m.capabilities, defaults, + pricing, } }) .collect(); diff --git a/crates/amplifier-core/src/generated/amplifier.module.rs b/crates/amplifier-core/src/generated/amplifier.module.rs index d510a11f..5f76cc5f 100644 --- a/crates/amplifier-core/src/generated/amplifier.module.rs +++ b/crates/amplifier-core/src/generated/amplifier.module.rs @@ -452,6 +452,10 @@ pub struct ModelInfo { pub capabilities: ::prost::alloc::vec::Vec<::prost::alloc::string::String>, #[prost(string, tag = "6")] pub defaults_json: ::prost::alloc::string::String, + /// JSON-encoded Pricing, or empty string if pricing is unavailable + /// (e.g., local providers, self-hosted backends). + #[prost(string, tag = "7")] + pub pricing_json: ::prost::alloc::string::String, } #[derive(Clone, PartialEq, ::prost::Message)] pub struct ProviderInfo { diff --git a/crates/amplifier-core/src/generated/conversions.rs b/crates/amplifier-core/src/generated/conversions.rs index 39b06cfa..6599f8ed 100644 --- a/crates/amplifier-core/src/generated/conversions.rs +++ b/crates/amplifier-core/src/generated/conversions.rs @@ -95,6 +95,11 @@ impl From for super::amplifier_module::ModelInfo { }), capabilities: native.capabilities, defaults_json: to_json_or_warn(&native.defaults, "ModelInfo defaults"), + pricing_json: native + .pricing + .as_ref() + .map(|p| to_json_or_warn(p, "ModelInfo pricing")) + .unwrap_or_default(), } } } @@ -112,6 +117,19 @@ impl From for crate::models::ModelInfo { } else { from_json_or_default(&proto.defaults_json, "ModelInfo defaults_json") }, + pricing: if proto.pricing_json.is_empty() { + None + } else { + match serde_json::from_str::(&proto.pricing_json) { + Ok(p) => Some(p), + Err(e) => { + log::warn!( + "Failed to parse ModelInfo pricing_json: {e} — pricing unavailable" + ); + None + } + } + }, } } } @@ -1029,10 +1047,54 @@ mod tests { max_output_tokens: 8192, capabilities: vec!["tools".into(), "vision".into()], defaults: HashMap::from([("temperature".to_string(), serde_json::json!(0.7))]), + pricing: Some(crate::models::Pricing { + input_per_million: 30.0, + output_per_million: 60.0, + cache_read_per_million: None, + cache_write_per_million: None, + currency: "USD".into(), + }), + }; + let proto: super::super::amplifier_module::ModelInfo = original.clone().into(); + let restored: crate::models::ModelInfo = proto.into(); + assert_eq!(original, restored); + } + + #[test] + fn model_info_pricing_none_roundtrips_to_empty_json() { + let original = crate::models::ModelInfo { + id: "local-model".into(), + display_name: "Local Model".into(), + context_window: 8192, + max_output_tokens: 4096, + capabilities: vec![], + defaults: HashMap::new(), + pricing: None, }; let proto: super::super::amplifier_module::ModelInfo = original.clone().into(); + assert!(proto.pricing_json.is_empty()); let restored: crate::models::ModelInfo = proto.into(); assert_eq!(original, restored); + assert!(restored.pricing.is_none()); + } + + #[test] + fn model_info_pricing_invalid_json_becomes_none() { + let mut proto = super::super::amplifier_module::ModelInfo { + id: "broken-model".into(), + display_name: "Broken".into(), + context_window: 1000, + max_output_tokens: 100, + capabilities: vec![], + defaults_json: String::new(), + pricing_json: "not-valid-json".into(), + }; + let restored: crate::models::ModelInfo = proto.clone().into(); + assert!(restored.pricing.is_none()); + + proto.pricing_json = String::new(); + let restored_empty: crate::models::ModelInfo = proto.into(); + assert!(restored_empty.pricing.is_none()); } #[test] @@ -1122,6 +1184,7 @@ mod tests { max_output_tokens: 100, capabilities: vec![], defaults: HashMap::new(), + pricing: None, }; let proto: super::super::amplifier_module::ModelInfo = original.into(); assert_eq!(proto.context_window, i32::MAX); @@ -1136,6 +1199,7 @@ mod tests { max_output_tokens: i64::from(i32::MAX) + 500, capabilities: vec![], defaults: HashMap::new(), + pricing: None, }; let proto: super::super::amplifier_module::ModelInfo = original.into(); assert_eq!(proto.max_output_tokens, i32::MAX); diff --git a/crates/amplifier-core/src/generated/equivalence_tests.rs b/crates/amplifier-core/src/generated/equivalence_tests.rs index 87b4775f..098c6f3c 100644 --- a/crates/amplifier-core/src/generated/equivalence_tests.rs +++ b/crates/amplifier-core/src/generated/equivalence_tests.rs @@ -194,6 +194,8 @@ mod tests { max_output_tokens: 4096, capabilities: vec!["vision".into(), "tools".into(), "streaming".into()], defaults_json: r#"{"temperature":0.7}"#.into(), + pricing_json: + r#"{"input_per_million":15.0,"output_per_million":75.0,"currency":"USD"}"#.into(), }; assert_eq!(info.id, "claude-3-opus"); assert_eq!(info.display_name, "Claude 3 Opus"); @@ -201,6 +203,10 @@ mod tests { assert_eq!(info.max_output_tokens, 4096); assert_eq!(info.capabilities.len(), 3); assert_eq!(info.defaults_json, r#"{"temperature":0.7}"#); + assert_eq!( + info.pricing_json, + r#"{"input_per_million":15.0,"output_per_million":75.0,"currency":"USD"}"# + ); } #[test] diff --git a/crates/amplifier-core/src/models.rs b/crates/amplifier-core/src/models.rs index 266c62e7..7c4e45f7 100644 --- a/crates/amplifier-core/src/models.rs +++ b/crates/amplifier-core/src/models.rs @@ -304,6 +304,11 @@ pub struct ModelInfo { /// Model-specific default config values (e.g., temperature, max_tokens). #[serde(default)] pub defaults: HashMap, + + /// Per-model pricing information. None when pricing is not available + /// (e.g., local providers like ollama, self-hosted backends like vllm). + #[serde(default)] + pub pricing: Option, } /// A configuration field that a provider needs, with prompt metadata. @@ -352,6 +357,36 @@ pub struct ConfigField { pub requires_model: bool, } +/// Per-model pricing information. +/// +/// Rates are per million tokens, in the specified currency. Surfaced via +/// `/v1/models` so HTTP-bridge applications (e.g., amplifier-app-opencode) +/// can display cost estimates without maintaining their own pricing tables. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct Pricing { + /// Cost per million input tokens. + pub input_per_million: f64, + + /// Cost per million output tokens. + pub output_per_million: f64, + + /// Cost per million cache-read input tokens (None if not supported). + #[serde(default)] + pub cache_read_per_million: Option, + + /// Cost per million cache-write input tokens (None if not supported). + #[serde(default)] + pub cache_write_per_million: Option, + + /// ISO 4217 currency code. + #[serde(default = "default_currency")] + pub currency: String, +} + +fn default_currency() -> String { + "USD".to_string() +} + /// Provider metadata. /// /// Describes capabilities, authentication requirements, and defaults for a provider. @@ -727,6 +762,7 @@ mod tests { max_output_tokens: 4096, capabilities: vec!["streaming".into()], defaults: Default::default(), + pricing: None, }; assert_eq!(info.id, "gpt-4"); } @@ -740,6 +776,13 @@ mod tests { max_output_tokens: 8192, capabilities: vec!["tools".into(), "vision".into(), "streaming".into()], defaults: HashMap::from([("temperature".into(), json!(0.7))]), + pricing: Some(Pricing { + input_per_million: 3.0, + output_per_million: 15.0, + cache_read_per_million: None, + cache_write_per_million: None, + currency: "USD".into(), + }), }; let json_str = serde_json::to_string(&info).unwrap(); let deserialized: ModelInfo = serde_json::from_str(&json_str).unwrap(); diff --git a/crates/amplifier-guest/src/lib.rs b/crates/amplifier-guest/src/lib.rs index c7385061..6be07cc4 100644 --- a/crates/amplifier-guest/src/lib.rs +++ b/crates/amplifier-guest/src/lib.rs @@ -1289,6 +1289,7 @@ mod provider_tests { max_output_tokens: 1024, capabilities: vec!["chat".to_string()], defaults: HashMap::new(), + pricing: None, }]) } diff --git a/crates/amplifier-guest/src/types.rs b/crates/amplifier-guest/src/types.rs index ada17d44..eb86aaaa 100644 --- a/crates/amplifier-guest/src/types.rs +++ b/crates/amplifier-guest/src/types.rs @@ -171,6 +171,27 @@ pub struct ProviderInfo { pub defaults: HashMap, } +/// Per-model pricing information. +/// +/// Rates are per million tokens, in the specified currency. Mirrors +/// `amplifier_core::models::Pricing` on the native side (this crate has no +/// dependency on `amplifier-core`, so the struct is duplicated here). +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct Pricing { + pub input_per_million: f64, + pub output_per_million: f64, + #[serde(default)] + pub cache_read_per_million: Option, + #[serde(default)] + pub cache_write_per_million: Option, + #[serde(default = "default_currency")] + pub currency: String, +} + +fn default_currency() -> String { + "USD".to_string() +} + /// Metadata about a specific model. #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct ModelInfo { @@ -180,6 +201,10 @@ pub struct ModelInfo { pub max_output_tokens: i64, pub capabilities: Vec, pub defaults: HashMap, + /// Per-model pricing information. None when pricing is not available + /// (e.g., local providers like ollama, self-hosted backends like vllm). + #[serde(default)] + pub pricing: Option, } /// Request for an LLM chat completion. @@ -473,9 +498,17 @@ mod tests { max_output_tokens: 4096, capabilities: vec!["chat".to_string(), "tools".to_string()], defaults: HashMap::new(), + pricing: Some(Pricing { + input_per_million: 30.0, + output_per_million: 60.0, + cache_read_per_million: None, + cache_write_per_million: None, + currency: "USD".to_string(), + }), }; assert_eq!(info.context_window, 128000); assert_eq!(info.max_output_tokens, 4096); + assert_eq!(info.pricing.as_ref().unwrap().input_per_million, 30.0); } // --- ChatRequest tests --- @@ -638,6 +671,13 @@ mod tests { max_output_tokens: 4096, capabilities: vec!["chat".to_string(), "tools".to_string()], defaults: HashMap::new(), + pricing: Some(Pricing { + input_per_million: 30.0, + output_per_million: 60.0, + cache_read_per_million: Some(15.0), + cache_write_per_million: Some(37.5), + currency: "USD".to_string(), + }), }; let json_str = serde_json::to_string(&original).unwrap(); let deserialized: ModelInfo = serde_json::from_str(&json_str).unwrap(); diff --git a/proto/amplifier_module.proto b/proto/amplifier_module.proto index e12f988a..35bc2cc3 100644 --- a/proto/amplifier_module.proto +++ b/proto/amplifier_module.proto @@ -415,6 +415,9 @@ message ModelInfo { int32 max_output_tokens = 4; repeated string capabilities = 5; string defaults_json = 6; + // JSON-encoded Pricing, or empty string if pricing is unavailable + // (e.g., local providers, self-hosted backends). + string pricing_json = 7; } message ProviderInfo { diff --git a/python/amplifier_core/__init__.py b/python/amplifier_core/__init__.py index 13303a8a..1de1d227 100644 --- a/python/amplifier_core/__init__.py +++ b/python/amplifier_core/__init__.py @@ -77,6 +77,7 @@ from .models import HookResult from .models import ModelInfo from .models import ModuleInfo +from .models import Pricing from .models import ProviderInfo from .models import SessionStatus from .models import ToolResult @@ -118,6 +119,7 @@ "ConfigField", "ModelInfo", "ModuleInfo", + "Pricing", "ProviderInfo", "SessionStatus", "ApprovalRequest", diff --git a/python/amplifier_core/models.py b/python/amplifier_core/models.py index 218d6499..9d7c2046 100644 --- a/python/amplifier_core/models.py +++ b/python/amplifier_core/models.py @@ -322,6 +322,39 @@ class HookResult(BaseModel): ) +class Pricing(BaseModel): + """Per-model pricing information. + + Rates are per million tokens, in the specified currency. Surfaced via + /v1/models so HTTP-bridge applications (e.g., amplifier-app-opencode) + can display cost estimates without maintaining their own pricing tables. + + Rate fields use `float` (not `Decimal`). These are display-only estimates + surfaced to consumers via `/v1/models` for UI-level cost display. For + per-turn cost accounting, use `Usage.cost_usd`, which is `Decimal` and + uses a field validator that explicitly rejects `float`. + """ + + input_per_million: float = Field(..., description="Cost per million input tokens") + output_per_million: float = Field(..., description="Cost per million output tokens") + cache_read_per_million: float | None = Field( + default=None, + description="Cost per million cache-read input tokens (None if not supported)", + ) + cache_write_per_million: float | None = Field( + default=None, + description="Cost per million cache-write input tokens (None if not supported)", + ) + currency: str = Field(default="USD", description="ISO 4217 currency code") + + @field_validator("currency") + @classmethod + def _validate_currency(cls, v: str) -> str: + if not re.match(r"^[A-Z]{3}$", v): + raise ValueError(f"currency must be a 3-letter ISO 4217 code, got {v!r}") + return v + + class ModelInfo(BaseModel): """Model metadata for provider models. @@ -342,6 +375,13 @@ class ModelInfo(BaseModel): default_factory=dict, description="Model-specific default config values (e.g., temperature, max_tokens)", ) + pricing: Pricing | None = Field( + default=None, + description=( + "Per-model pricing information. None when pricing is not available " + "(e.g., local providers like ollama, self-hosted backends like vllm)." + ), + ) class ConfigField(BaseModel): diff --git a/tests/fixtures/wasm/src/echo-provider/src/lib.rs b/tests/fixtures/wasm/src/echo-provider/src/lib.rs index 4650ecae..f4c46cf3 100644 --- a/tests/fixtures/wasm/src/echo-provider/src/lib.rs +++ b/tests/fixtures/wasm/src/echo-provider/src/lib.rs @@ -31,6 +31,7 @@ impl Provider for EchoProvider { max_output_tokens: 1024, capabilities: vec!["chat".to_string()], defaults: HashMap::new(), + pricing: None, }]) } diff --git a/tests/test_model_info_pricing.py b/tests/test_model_info_pricing.py new file mode 100644 index 00000000..1f93e88a --- /dev/null +++ b/tests/test_model_info_pricing.py @@ -0,0 +1,108 @@ +"""Serialization contract for ModelInfo.pricing and the Pricing model. + +pricing is optional on ModelInfo: None when a provider has no rate data +(e.g., local providers like ollama, self-hosted backends like vllm), and a +populated Pricing object when a provider can supply rates. Both states must +round-trip cleanly through model_dump() / model_validate() so HTTP bridges +(e.g., amplifier-app-opencode) can rely on the field without special-casing +either branch. +""" + +import json + +from pydantic import ValidationError + +from amplifier_core.models import ModelInfo +from amplifier_core.models import Pricing + + +class TestPricingRoundTrip: + """Pricing model dumps and reconstructs without loss.""" + + def test_round_trip_full(self): + pricing = Pricing( + input_per_million=3.0, + output_per_million=15.0, + cache_read_per_million=0.3, + cache_write_per_million=3.75, + currency="USD", + ) + dumped = pricing.model_dump() + rebuilt = Pricing.model_validate(dumped) + + assert rebuilt == pricing + assert rebuilt.input_per_million == 3.0 + assert rebuilt.output_per_million == 15.0 + assert rebuilt.cache_read_per_million == 0.3 + assert rebuilt.cache_write_per_million == 3.75 + assert rebuilt.currency == "USD" + + def test_optional_fields_default_to_none(self): + pricing = Pricing(input_per_million=1.0, output_per_million=5.0) + + assert pricing.cache_read_per_million is None + assert pricing.cache_write_per_million is None + assert pricing.currency == "USD" + + def test_pricing_json_dumps_survives(self): + p = Pricing(input_per_million=3.0, output_per_million=15.0) + # Verify json.dumps(model_dump()) works without needing mode="json" + # (since we no longer have date fields, this should be trivially safe) + json.dumps(p.model_dump()) + # And explicit mode="json" for consistency + json.dumps(p.model_dump(mode="json")) + # And model_dump_json() directly + Pricing.model_validate_json(p.model_dump_json()) + + def test_currency_must_be_iso_4217(self): + # valid 3-letter uppercase + Pricing(input_per_million=1.0, output_per_million=2.0, currency="EUR") + # invalid + for bad in ["usd", "US", "USDD", "us$", "123"]: + try: + Pricing(input_per_million=1.0, output_per_million=2.0, currency=bad) + raise AssertionError(f"expected ValidationError for currency={bad!r}") + except ValidationError: + pass + + +class TestModelInfoPricingField: + """ModelInfo.pricing: optional, backwards-compatible, round-trips both states.""" + + def test_pricing_none_serializes_cleanly(self): + """Local/self-hosted providers (no rate data) omit pricing without error.""" + model = ModelInfo( + id="local-model", + display_name="Local Model", + context_window=8192, + max_output_tokens=4096, + ) + + dumped = model.model_dump() + + assert dumped["pricing"] is None + rebuilt = ModelInfo.model_validate(dumped) + assert rebuilt.pricing is None + + def test_pricing_populated_round_trips(self): + pricing = Pricing( + input_per_million=3.0, + output_per_million=15.0, + cache_read_per_million=0.3, + cache_write_per_million=3.75, + ) + model = ModelInfo( + id="claude-sonnet-4-5", + display_name="Claude Sonnet 4.5", + context_window=200_000, + max_output_tokens=64_000, + pricing=pricing, + ) + + dumped = model.model_dump() + rebuilt = ModelInfo.model_validate(dumped) + + assert rebuilt.pricing is not None + assert rebuilt.pricing == pricing + assert dumped["pricing"]["input_per_million"] == 3.0 + assert dumped["pricing"]["output_per_million"] == 15.0 diff --git a/uv.lock b/uv.lock index 652a8412..c2886455 100644 --- a/uv.lock +++ b/uv.lock @@ -4,7 +4,7 @@ requires-python = ">=3.11" [[package]] name = "amplifier-core" -version = "1.5.1" +version = "1.6.0" source = { editable = "." } dependencies = [ { name = "click" },