Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions bindings/python/tests/test_schema_sync.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
19 changes: 19 additions & 0 deletions crates/amplifier-core/src/bridges/grpc_provider.rs
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,23 @@ fn parse_defaults_json(json_str: &str, id: &str) -> HashMap<String, Value> {
})
}

/// 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<crate::models::Pricing> {
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
Expand Down Expand Up @@ -125,13 +142,15 @@ 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,
context_window: m.context_window as i64,
max_output_tokens: m.max_output_tokens as i64,
capabilities: m.capabilities,
defaults,
pricing,
}
})
.collect();
Expand Down
4 changes: 4 additions & 0 deletions crates/amplifier-core/src/generated/amplifier.module.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
64 changes: 64 additions & 0 deletions crates/amplifier-core/src/generated/conversions.rs
Original file line number Diff line number Diff line change
Expand Up @@ -95,6 +95,11 @@ impl From<crate::models::ModelInfo> 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(),
}
}
}
Expand All @@ -112,6 +117,19 @@ impl From<super::amplifier_module::ModelInfo> 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::<crate::models::Pricing>(&proto.pricing_json) {
Ok(p) => Some(p),
Err(e) => {
log::warn!(
"Failed to parse ModelInfo pricing_json: {e} β€” pricing unavailable"
);
None
}
}
},
}
}
}
Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -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);
Expand All @@ -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);
Expand Down
6 changes: 6 additions & 0 deletions crates/amplifier-core/src/generated/equivalence_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -194,13 +194,19 @@ 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");
assert_eq!(info.context_window, 200_000);
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]
Expand Down
43 changes: 43 additions & 0 deletions crates/amplifier-core/src/models.rs
Original file line number Diff line number Diff line change
Expand Up @@ -304,6 +304,11 @@ pub struct ModelInfo {
/// Model-specific default config values (e.g., temperature, max_tokens).
#[serde(default)]
pub defaults: HashMap<String, Value>,

/// 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<Pricing>,
}

/// A configuration field that a provider needs, with prompt metadata.
Expand Down Expand Up @@ -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<f64>,

/// Cost per million cache-write input tokens (None if not supported).
#[serde(default)]
pub cache_write_per_million: Option<f64>,

/// 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.
Expand Down Expand Up @@ -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");
}
Expand All @@ -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();
Expand Down
1 change: 1 addition & 0 deletions crates/amplifier-guest/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1289,6 +1289,7 @@ mod provider_tests {
max_output_tokens: 1024,
capabilities: vec!["chat".to_string()],
defaults: HashMap::new(),
pricing: None,
}])
}

Expand Down
40 changes: 40 additions & 0 deletions crates/amplifier-guest/src/types.rs
Original file line number Diff line number Diff line change
Expand Up @@ -171,6 +171,27 @@ pub struct ProviderInfo {
pub defaults: HashMap<String, Value>,
}

/// 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<f64>,
#[serde(default)]
pub cache_write_per_million: Option<f64>,
#[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 {
Expand All @@ -180,6 +201,10 @@ pub struct ModelInfo {
pub max_output_tokens: i64,
pub capabilities: Vec<String>,
pub defaults: HashMap<String, Value>,
/// 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<Pricing>,
}

/// Request for an LLM chat completion.
Expand Down Expand Up @@ -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 ---
Expand Down Expand Up @@ -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();
Expand Down
3 changes: 3 additions & 0 deletions proto/amplifier_module.proto
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
2 changes: 2 additions & 0 deletions python/amplifier_core/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -118,6 +119,7 @@
"ConfigField",
"ModelInfo",
"ModuleInfo",
"Pricing",
"ProviderInfo",
"SessionStatus",
"ApprovalRequest",
Expand Down
Loading
Loading