mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-09-29 01:42:21 +00:00
feat(llm): add Response.cost_usd and CostSource across all adapters
Adds two new fields to the unified completion Response: - cost_usd: Option<f64> — USD cost for the completion - cost_source: Option<CostSource> — Authoritative | Estimated A new fabro_llm::cost::estimate_cost_usd helper centralizes the catalog math (tokens × catalog price → USD micros → f64). Every adapter response path (Anthropic, OpenAI, Gemini, openai_chat) calls the helper when a catalog is available, producing Estimated costs. Adapters with authoritative provider-side pricing (e.g. OpenRouter, in a later commit) will bypass the helper and set Authoritative. The openai_chat streaming path now plumbs catalog + speed through StreamState so finish_events can call the helper before assembling the final Response. Public API surface: CompletionResponse in fabro-api.yaml is extended with cost_usd and cost_source. Progenitor + the TS client are regenerated; apps/fabro-web typechecks clean since both fields are optional. 42 Response constructor sites updated: - 5 live adapter sites call estimate_cost_usd - ~32 synthesized/fallback/test sites set (None, None) explicitly - 5 test fixtures in types.rs set (None, None) 397 lib tests pass (was 392, +5 new helper unit tests). Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
7624481d8f
commit
480ff56aba
23 changed files with 612 additions and 50 deletions
|
|
@ -7410,6 +7410,22 @@ components:
|
|||
$ref: "#/components/schemas/CompletionUsage"
|
||||
output:
|
||||
description: Parsed structured output when schema was provided.
|
||||
cost_usd:
|
||||
type: number
|
||||
format: double
|
||||
nullable: true
|
||||
description: |
|
||||
Total USD cost for the completion. Populated when the provider
|
||||
returns an authoritative billing figure (e.g. OpenRouter) or
|
||||
when the model has catalog pricing.
|
||||
cost_source:
|
||||
type: string
|
||||
enum: [authoritative, estimated]
|
||||
nullable: true
|
||||
description: |
|
||||
Whether `cost_usd` came from the provider's billing data
|
||||
(`authoritative`) or was computed from catalog prices
|
||||
(`estimated`).
|
||||
|
||||
PaginatedSavedQueryList:
|
||||
description: Paginated list of saved queries.
|
||||
|
|
|
|||
|
|
@ -1747,6 +1747,8 @@ def farewell(name):
|
|||
},
|
||||
finish_reason: FinishReason::ToolCalls,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -1823,6 +1825,8 @@ def farewell(name):
|
|||
},
|
||||
finish_reason: FinishReason::ToolCalls,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
|
|||
|
|
@ -184,6 +184,8 @@ pub fn text_response(text: &str) -> Response {
|
|||
output_tokens: 5,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -260,6 +262,8 @@ pub fn tool_call_response(
|
|||
output_tokens: 5,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -387,6 +391,8 @@ pub fn multi_tool_call_response(calls: Vec<(&str, &str, serde_json::Value)>) ->
|
|||
output_tokens: 5,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
|
|||
|
|
@ -224,6 +224,8 @@ impl ProviderAdapter for AuthenticatedFabroServerAdapter {
|
|||
output_tokens: server_response.usage.output_tokens,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
|
|||
|
|
@ -584,6 +584,8 @@ mod tests {
|
|||
output_tokens: 20,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -605,6 +607,8 @@ mod tests {
|
|||
message: Message::assistant(&text),
|
||||
finish_reason: FinishReason::Stop,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
|
|||
211
lib/crates/fabro-llm/src/cost.rs
Normal file
211
lib/crates/fabro-llm/src/cost.rs
Normal file
|
|
@ -0,0 +1,211 @@
|
|||
//! Catalog-derived cost estimation for completion responses.
|
||||
|
||||
use fabro_model::billing::{ModelRef, Speed, TokenCounts};
|
||||
use fabro_model::{Catalog, ProviderId};
|
||||
|
||||
use crate::types::CostSource;
|
||||
|
||||
/// Estimate the USD cost of a completion from the catalog's per-token
|
||||
/// pricing for the model. Returns `(None, None)` if the catalog is absent,
|
||||
/// the model is not in the catalog, or the model has no pricing.
|
||||
///
|
||||
/// Used by adapters that don't receive an authoritative cost from the
|
||||
/// provider. Adapters with authoritative pricing (e.g. OpenRouter
|
||||
/// `usage.cost`) bypass this and set `Authoritative` directly.
|
||||
#[must_use]
|
||||
pub fn estimate_cost_usd(
|
||||
catalog: Option<&Catalog>,
|
||||
provider: &str,
|
||||
model: &str,
|
||||
tokens: &TokenCounts,
|
||||
speed: Option<Speed>,
|
||||
) -> (Option<f64>, Option<CostSource>) {
|
||||
let Some(catalog) = catalog else {
|
||||
return (None, None);
|
||||
};
|
||||
let model_ref = ModelRef {
|
||||
provider: ProviderId::from(provider),
|
||||
model_id: model.to_string(),
|
||||
speed,
|
||||
};
|
||||
let Some(micros) = catalog.price_tokens(&model_ref, tokens) else {
|
||||
return (None, None);
|
||||
};
|
||||
#[allow(
|
||||
clippy::cast_precision_loss,
|
||||
reason = "micros fit comfortably in f64 for any realistic completion cost"
|
||||
)]
|
||||
let usd = micros as f64 / 1_000_000.0;
|
||||
(Some(usd), Some(CostSource::Estimated))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use fabro_model::catalog::LlmCatalogSettings;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn catalog_with_openai_model() -> Catalog {
|
||||
let settings: LlmCatalogSettings = toml::from_str(
|
||||
r#"
|
||||
[providers.openai]
|
||||
display_name = "OpenAI"
|
||||
adapter = "openai"
|
||||
agent_profile = "openai"
|
||||
|
||||
[models."gpt-test"]
|
||||
provider = "openai"
|
||||
display_name = "GPT Test"
|
||||
family = "gpt"
|
||||
default = true
|
||||
|
||||
[models."gpt-test".limits]
|
||||
context_window = 200000
|
||||
max_output = 4096
|
||||
|
||||
[models."gpt-test".features]
|
||||
tools = true
|
||||
vision = false
|
||||
reasoning = false
|
||||
|
||||
[models."gpt-test".costs]
|
||||
input_cost_per_mtok = 1.0
|
||||
output_cost_per_mtok = 2.0
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
Catalog::from_settings(&settings).unwrap()
|
||||
}
|
||||
|
||||
fn catalog_without_costs() -> Catalog {
|
||||
let settings: LlmCatalogSettings = toml::from_str(
|
||||
r#"
|
||||
[providers.openai]
|
||||
display_name = "OpenAI"
|
||||
adapter = "openai"
|
||||
agent_profile = "openai"
|
||||
|
||||
[models."gpt-no-cost"]
|
||||
provider = "openai"
|
||||
display_name = "GPT No Cost"
|
||||
family = "gpt"
|
||||
default = true
|
||||
|
||||
[models."gpt-no-cost".limits]
|
||||
context_window = 200000
|
||||
max_output = 4096
|
||||
|
||||
[models."gpt-no-cost".features]
|
||||
tools = true
|
||||
vision = false
|
||||
reasoning = false
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
Catalog::from_settings(&settings).unwrap()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn returns_none_when_catalog_is_none() {
|
||||
let tokens = TokenCounts {
|
||||
input_tokens: 1000,
|
||||
output_tokens: 500,
|
||||
..TokenCounts::default()
|
||||
};
|
||||
let (cost, source) = estimate_cost_usd(None, "openai", "gpt-test", &tokens, None);
|
||||
assert_eq!(cost, None);
|
||||
assert_eq!(source, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn returns_estimated_when_model_priced() {
|
||||
let catalog = catalog_with_openai_model();
|
||||
let tokens = TokenCounts {
|
||||
input_tokens: 1_000_000, // 1M tokens at $1/Mtok = $1.00
|
||||
output_tokens: 500_000, // 500k tokens at $2/Mtok = $1.00
|
||||
..TokenCounts::default()
|
||||
};
|
||||
let (cost, source) = estimate_cost_usd(Some(&catalog), "openai", "gpt-test", &tokens, None);
|
||||
assert_eq!(source, Some(CostSource::Estimated));
|
||||
let cost = cost.expect("cost should be Some");
|
||||
// 1.00 + 1.00 = 2.00 USD
|
||||
assert!((cost - 2.0).abs() < 1e-9, "expected ~$2.00, got {cost}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn returns_none_when_model_missing_from_catalog() {
|
||||
let catalog = catalog_with_openai_model();
|
||||
let tokens = TokenCounts {
|
||||
input_tokens: 1000,
|
||||
output_tokens: 500,
|
||||
..TokenCounts::default()
|
||||
};
|
||||
let (cost, source) =
|
||||
estimate_cost_usd(Some(&catalog), "openai", "nonexistent-model", &tokens, None);
|
||||
assert_eq!(cost, None);
|
||||
assert_eq!(source, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn returns_none_when_model_has_no_pricing() {
|
||||
let catalog = catalog_without_costs();
|
||||
let tokens = TokenCounts {
|
||||
input_tokens: 1000,
|
||||
output_tokens: 500,
|
||||
..TokenCounts::default()
|
||||
};
|
||||
let (cost, source) =
|
||||
estimate_cost_usd(Some(&catalog), "openai", "gpt-no-cost", &tokens, None);
|
||||
assert_eq!(cost, None);
|
||||
assert_eq!(source, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn micros_to_usd_conversion_is_exact_for_integer_amounts() {
|
||||
// 1_500_000 micros = $1.50 exactly (representable as f64).
|
||||
// Configure pricing so that price_tokens returns exactly 1_500_000.
|
||||
// input_cost_per_mtok = 1.5 USD; 1M input tokens with no output yields
|
||||
// 1_500_000 micros.
|
||||
let settings: LlmCatalogSettings = toml::from_str(
|
||||
r#"
|
||||
[providers.openai]
|
||||
display_name = "OpenAI"
|
||||
adapter = "openai"
|
||||
agent_profile = "openai"
|
||||
|
||||
[models."gpt-exact"]
|
||||
provider = "openai"
|
||||
display_name = "GPT Exact"
|
||||
family = "gpt"
|
||||
default = true
|
||||
|
||||
[models."gpt-exact".limits]
|
||||
context_window = 200000
|
||||
max_output = 4096
|
||||
|
||||
[models."gpt-exact".features]
|
||||
tools = true
|
||||
vision = false
|
||||
reasoning = false
|
||||
|
||||
[models."gpt-exact".costs]
|
||||
input_cost_per_mtok = 1.5
|
||||
output_cost_per_mtok = 0.0
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
let catalog = Catalog::from_settings(&settings).unwrap();
|
||||
let tokens = TokenCounts {
|
||||
input_tokens: 1_000_000,
|
||||
output_tokens: 0,
|
||||
..TokenCounts::default()
|
||||
};
|
||||
let (cost, _source) =
|
||||
estimate_cost_usd(Some(&catalog), "openai", "gpt-exact", &tokens, None);
|
||||
let cost = cost.expect("cost should be Some");
|
||||
assert!(
|
||||
(cost - 1.5).abs() < f64::EPSILON,
|
||||
"expected $1.50 exact, got {cost}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -1135,6 +1135,8 @@ mod tests {
|
|||
output_tokens: 20,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -1163,6 +1165,8 @@ mod tests {
|
|||
output_tokens: 20,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -1271,6 +1275,8 @@ mod tests {
|
|||
output_tokens: 5,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -1288,6 +1294,8 @@ mod tests {
|
|||
output_tokens: 10,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -1357,6 +1365,8 @@ mod tests {
|
|||
output_tokens: 2,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -1562,6 +1572,8 @@ mod tests {
|
|||
message: Message::assistant(&self.full_text),
|
||||
finish_reason: FinishReason::Stop,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -1593,6 +1605,8 @@ mod tests {
|
|||
output_tokens: 20,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -1772,6 +1786,8 @@ mod tests {
|
|||
},
|
||||
finish_reason: FinishReason::ToolCalls,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -1938,6 +1954,8 @@ mod tests {
|
|||
message: Message::assistant("fallback"),
|
||||
finish_reason: FinishReason::Stop,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -1967,6 +1985,8 @@ mod tests {
|
|||
output_tokens: 5,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -1994,6 +2014,8 @@ mod tests {
|
|||
output_tokens: 10,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -2110,6 +2132,8 @@ mod tests {
|
|||
output_tokens: 5,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -2273,6 +2297,8 @@ mod tests {
|
|||
message: Message::assistant("fallback"),
|
||||
finish_reason: FinishReason::Stop,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -2304,6 +2330,8 @@ mod tests {
|
|||
output_tokens: 20,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -2386,6 +2414,8 @@ mod tests {
|
|||
message: Message::assistant("fallback"),
|
||||
finish_reason: FinishReason::Stop,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -2402,6 +2432,8 @@ mod tests {
|
|||
message: Message::assistant(text),
|
||||
finish_reason: FinishReason::Stop,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -2484,6 +2516,8 @@ mod tests {
|
|||
message: Message::assistant("fallback"),
|
||||
finish_reason: FinishReason::Stop,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -2509,6 +2543,8 @@ mod tests {
|
|||
},
|
||||
finish_reason: FinishReason::ToolCalls,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -2533,6 +2569,8 @@ mod tests {
|
|||
message: Message::assistant(text),
|
||||
finish_reason: FinishReason::Stop,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
pub mod adapter_registry;
|
||||
pub mod client;
|
||||
pub mod cost;
|
||||
pub mod error;
|
||||
pub mod generate;
|
||||
pub mod middleware;
|
||||
|
|
|
|||
|
|
@ -208,6 +208,8 @@ mod tests {
|
|||
message: Message::assistant(text),
|
||||
finish_reason: FinishReason::Stop,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
|
|||
use fabro_model::{Catalog, ReasoningEffortFeature};
|
||||
use futures::stream;
|
||||
|
||||
use crate::cost::estimate_cost_usd;
|
||||
use crate::error::{Error, ProviderErrorDetail, ProviderErrorKind, error_from_status_code};
|
||||
use crate::provider::{ProviderAdapter, StreamEventStream};
|
||||
use crate::providers::common::{
|
||||
|
|
@ -747,10 +748,27 @@ struct StreamAccumulator {
|
|||
current_tool_args: String,
|
||||
/// Rate limit info parsed from the initial HTTP response headers.
|
||||
rate_limit: Option<RateLimitInfo>,
|
||||
/// Provider name to attribute to the final Response and to use for
|
||||
/// catalog cost estimation.
|
||||
provider_name: String,
|
||||
/// Optional catalog used to estimate `cost_usd` on the final Response.
|
||||
catalog: Option<Arc<Catalog>>,
|
||||
/// Speed setting forwarded from the request, used by catalog pricing.
|
||||
speed: Option<fabro_model::Speed>,
|
||||
}
|
||||
|
||||
impl StreamAccumulator {
|
||||
#[cfg(test)]
|
||||
fn new(rate_limit: Option<RateLimitInfo>) -> Self {
|
||||
Self::new_with_pricing(rate_limit, "anthropic".to_string(), None, None)
|
||||
}
|
||||
|
||||
fn new_with_pricing(
|
||||
rate_limit: Option<RateLimitInfo>,
|
||||
provider_name: String,
|
||||
catalog: Option<Arc<Catalog>>,
|
||||
speed: Option<fabro_model::Speed>,
|
||||
) -> Self {
|
||||
Self {
|
||||
id: String::new(),
|
||||
model: String::new(),
|
||||
|
|
@ -762,6 +780,9 @@ impl StreamAccumulator {
|
|||
current_thinking: String::new(),
|
||||
current_tool_args: String::new(),
|
||||
rate_limit,
|
||||
provider_name,
|
||||
catalog,
|
||||
speed,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -769,21 +790,30 @@ impl StreamAccumulator {
|
|||
/// parts.
|
||||
fn take_response(&mut self) -> Response {
|
||||
let content_parts = std::mem::take(&mut self.content_parts);
|
||||
let (cost_usd, cost_source) = estimate_cost_usd(
|
||||
self.catalog.as_deref(),
|
||||
&self.provider_name,
|
||||
&self.model,
|
||||
&self.usage,
|
||||
self.speed,
|
||||
);
|
||||
Response {
|
||||
id: self.id.clone(),
|
||||
model: self.model.clone(),
|
||||
provider: "anthropic".to_string(),
|
||||
message: Message {
|
||||
id: self.id.clone(),
|
||||
model: self.model.clone(),
|
||||
provider: self.provider_name.clone(),
|
||||
message: Message {
|
||||
role: Role::Assistant,
|
||||
content: content_parts,
|
||||
name: None,
|
||||
tool_call_id: None,
|
||||
},
|
||||
finish_reason: self.finish_reason.clone(),
|
||||
usage: self.usage.clone(),
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: self.rate_limit.clone(),
|
||||
usage: self.usage.clone(),
|
||||
cost_usd,
|
||||
cost_source,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: self.rate_limit.clone(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -1111,10 +1141,17 @@ impl SseReaderState {
|
|||
json_schema_mode: bool,
|
||||
stream_read_timeout: Option<std::time::Duration>,
|
||||
provider_name: String,
|
||||
catalog: Option<Arc<Catalog>>,
|
||||
speed: Option<fabro_model::Speed>,
|
||||
) -> Self {
|
||||
Self {
|
||||
line_reader: super::common::LineReader::new(http_resp, stream_read_timeout),
|
||||
accumulator: StreamAccumulator::new(rate_limit),
|
||||
accumulator: StreamAccumulator::new_with_pricing(
|
||||
rate_limit,
|
||||
provider_name.clone(),
|
||||
catalog,
|
||||
speed,
|
||||
),
|
||||
pending_events: std::collections::VecDeque::new(),
|
||||
json_schema_mode,
|
||||
provider_name,
|
||||
|
|
@ -1471,6 +1508,14 @@ impl ProviderAdapter for Adapter {
|
|||
} else {
|
||||
map_finish_reason(api_resp.stop_reason.as_deref())
|
||||
};
|
||||
let usage = token_counts_from_api_usage(&api_resp.usage);
|
||||
let (cost_usd, cost_source) = estimate_cost_usd(
|
||||
self.catalog.as_deref(),
|
||||
&self.provider_name,
|
||||
&api_resp.model,
|
||||
&usage,
|
||||
request.speed,
|
||||
);
|
||||
Ok(Response {
|
||||
id: api_resp.id,
|
||||
model: api_resp.model,
|
||||
|
|
@ -1482,7 +1527,9 @@ impl ProviderAdapter for Adapter {
|
|||
tool_call_id: None,
|
||||
},
|
||||
finish_reason,
|
||||
usage: token_counts_from_api_usage(&api_resp.usage),
|
||||
usage,
|
||||
cost_usd,
|
||||
cost_source,
|
||||
raw: serde_json::from_str(&body).ok(),
|
||||
warnings: vec![],
|
||||
rate_limit: parse_rate_limit_headers(&headers),
|
||||
|
|
@ -1527,6 +1574,8 @@ impl ProviderAdapter for Adapter {
|
|||
json_schema_mode,
|
||||
stream_read_timeout,
|
||||
self.provider_name.clone(),
|
||||
self.catalog.clone(),
|
||||
request.speed,
|
||||
),
|
||||
|mut state| async move {
|
||||
loop {
|
||||
|
|
@ -2430,6 +2479,8 @@ reasoning = true
|
|||
},
|
||||
finish_reason: FinishReason::ToolCalls,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
|
|||
|
|
@ -148,6 +148,12 @@ impl ProviderAdapter for Adapter {
|
|||
output_tokens: server_resp.usage.output_tokens,
|
||||
..Default::default()
|
||||
},
|
||||
// The remote fabro-server has its own catalog and may report
|
||||
// cost on its wire response. Until that wire format carries
|
||||
// `cost_usd`/`cost_source`, leave these `None` rather than
|
||||
// pretending to estimate locally.
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ use fabro_http::HeaderMap;
|
|||
use fabro_model::Catalog;
|
||||
use futures::stream;
|
||||
|
||||
use crate::cost::estimate_cost_usd;
|
||||
use crate::error::{
|
||||
Error, ProviderErrorDetail, ProviderErrorKind, error_from_grpc_status, error_from_status_code,
|
||||
};
|
||||
|
|
@ -642,9 +643,20 @@ fn process_sse_stream(
|
|||
model: String,
|
||||
rate_limit: Option<RateLimitInfo>,
|
||||
stream_read_timeout: Option<std::time::Duration>,
|
||||
provider_name: String,
|
||||
catalog: Option<Arc<Catalog>>,
|
||||
speed: Option<fabro_model::Speed>,
|
||||
) -> StreamEventStream {
|
||||
Box::pin(stream::unfold(
|
||||
SseStreamState::new(http_resp, model, rate_limit, stream_read_timeout),
|
||||
SseStreamState::new(
|
||||
http_resp,
|
||||
model,
|
||||
rate_limit,
|
||||
stream_read_timeout,
|
||||
provider_name,
|
||||
catalog,
|
||||
speed,
|
||||
),
|
||||
|mut state| async move {
|
||||
// If we have buffered events, yield them first.
|
||||
if let Some(event) = state.pending_events.pop_front() {
|
||||
|
|
@ -750,6 +762,13 @@ struct SseStreamState {
|
|||
finished: bool,
|
||||
/// Rate limit info parsed from HTTP response headers.
|
||||
rate_limit: Option<RateLimitInfo>,
|
||||
/// Provider name to attribute to the final Response and use for catalog
|
||||
/// cost estimation.
|
||||
provider_name: String,
|
||||
/// Optional catalog used to estimate `cost_usd` on the final Response.
|
||||
catalog: Option<Arc<Catalog>>,
|
||||
/// Speed setting forwarded from the request, used by catalog pricing.
|
||||
speed: Option<fabro_model::Speed>,
|
||||
}
|
||||
|
||||
impl SseStreamState {
|
||||
|
|
@ -758,6 +777,9 @@ impl SseStreamState {
|
|||
model: String,
|
||||
rate_limit: Option<RateLimitInfo>,
|
||||
stream_read_timeout: Option<std::time::Duration>,
|
||||
provider_name: String,
|
||||
catalog: Option<Arc<Catalog>>,
|
||||
speed: Option<fabro_model::Speed>,
|
||||
) -> Self {
|
||||
Self {
|
||||
line_reader: super::common::LineReader::new(http_resp, stream_read_timeout),
|
||||
|
|
@ -774,6 +796,9 @@ impl SseStreamState {
|
|||
finish_reason_str: None,
|
||||
finished: false,
|
||||
rate_limit,
|
||||
provider_name,
|
||||
catalog,
|
||||
speed,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -906,21 +931,30 @@ impl SseStreamState {
|
|||
content_parts.push(ContentPart::ToolCall(tc.clone()));
|
||||
}
|
||||
|
||||
let (cost_usd, cost_source) = estimate_cost_usd(
|
||||
self.catalog.as_deref(),
|
||||
&self.provider_name,
|
||||
&self.model,
|
||||
&self.usage,
|
||||
self.speed,
|
||||
);
|
||||
let response = Response {
|
||||
id: uuid::Uuid::new_v4().to_string(),
|
||||
model: self.model.clone(),
|
||||
provider: "gemini".to_string(),
|
||||
message: Message {
|
||||
id: uuid::Uuid::new_v4().to_string(),
|
||||
model: self.model.clone(),
|
||||
provider: self.provider_name.clone(),
|
||||
message: Message {
|
||||
role: Role::Assistant,
|
||||
content: content_parts,
|
||||
name: None,
|
||||
tool_call_id: None,
|
||||
},
|
||||
finish_reason: finish_reason.clone(),
|
||||
usage: self.usage.clone(),
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: self.rate_limit.clone(),
|
||||
usage: self.usage.clone(),
|
||||
cost_usd,
|
||||
cost_source,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: self.rate_limit.clone(),
|
||||
};
|
||||
|
||||
StreamEvent::finish(finish_reason, self.usage.clone(), response)
|
||||
|
|
@ -1027,6 +1061,13 @@ impl ProviderAdapter for Adapter {
|
|||
|
||||
let usage = parse_usage(api_resp.usage_metadata.as_ref());
|
||||
|
||||
let (cost_usd, cost_source) = estimate_cost_usd(
|
||||
self.catalog.as_deref(),
|
||||
&self.provider_name,
|
||||
&request.model,
|
||||
&usage,
|
||||
request.speed,
|
||||
);
|
||||
Ok(Response {
|
||||
id: uuid::Uuid::new_v4().to_string(),
|
||||
model: request.model.clone(),
|
||||
|
|
@ -1039,6 +1080,8 @@ impl ProviderAdapter for Adapter {
|
|||
},
|
||||
finish_reason,
|
||||
usage,
|
||||
cost_usd,
|
||||
cost_source,
|
||||
raw: serde_json::from_str(&body).ok(),
|
||||
warnings: vec![],
|
||||
rate_limit: parse_rate_limit_headers(&headers),
|
||||
|
|
@ -1070,6 +1113,9 @@ impl ProviderAdapter for Adapter {
|
|||
request.model.clone(),
|
||||
rate_limit,
|
||||
self.http.stream_read_timeout,
|
||||
self.provider_name.clone(),
|
||||
self.catalog.clone(),
|
||||
request.speed,
|
||||
))
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
|
|||
use fabro_model::Catalog;
|
||||
use futures::{StreamExt, stream};
|
||||
|
||||
use crate::cost::estimate_cost_usd;
|
||||
use crate::error::{Error, ProviderErrorDetail, ProviderErrorKind, error_from_status_code};
|
||||
use crate::provider::{
|
||||
ProviderAdapter, StreamEventStream, validate_standard_speed, validate_tool_choice,
|
||||
|
|
@ -794,6 +795,13 @@ struct SseStreamState {
|
|||
emitted_reasoning_start: bool,
|
||||
raw_response: Option<serde_json::Value>,
|
||||
rate_limit: Option<RateLimitInfo>,
|
||||
/// Provider name to attribute to the final Response and to use for
|
||||
/// catalog cost estimation.
|
||||
provider_name: String,
|
||||
/// Optional catalog used to estimate `cost_usd` on the final Response.
|
||||
catalog: Option<Arc<Catalog>>,
|
||||
/// Speed setting forwarded from the request, used by catalog pricing.
|
||||
speed: Option<fabro_model::Speed>,
|
||||
}
|
||||
|
||||
/// Parse a single SSE message block into an (`event_type`, `data`) pair.
|
||||
|
|
@ -1208,6 +1216,13 @@ fn handle_response_completed(
|
|||
state.response_model.clone()
|
||||
};
|
||||
|
||||
let (cost_usd, cost_source) = estimate_cost_usd(
|
||||
state.catalog.as_deref(),
|
||||
&state.provider_name,
|
||||
&model,
|
||||
&state.usage,
|
||||
state.speed,
|
||||
);
|
||||
let response = Response {
|
||||
id: state.response_id.clone(),
|
||||
model,
|
||||
|
|
@ -1220,6 +1235,8 @@ fn handle_response_completed(
|
|||
},
|
||||
finish_reason: state.finish_reason.clone(),
|
||||
usage: state.usage.clone(),
|
||||
cost_usd,
|
||||
cost_source,
|
||||
raw: state.raw_response.clone(),
|
||||
warnings: vec![],
|
||||
rate_limit: state.rate_limit.clone(),
|
||||
|
|
@ -1323,9 +1340,17 @@ impl ProviderAdapter for Adapter {
|
|||
|
||||
let usage = token_counts_from_api_usage(api_resp.usage.as_ref());
|
||||
|
||||
let resp_model = api_resp.model.unwrap_or_else(|| request.model.clone());
|
||||
let (cost_usd, cost_source) = estimate_cost_usd(
|
||||
self.catalog.as_deref(),
|
||||
&self.provider_name,
|
||||
&resp_model,
|
||||
&usage,
|
||||
request.speed,
|
||||
);
|
||||
Ok(Response {
|
||||
id: api_resp.id,
|
||||
model: api_resp.model.unwrap_or_else(|| request.model.clone()),
|
||||
model: resp_model,
|
||||
provider: "openai".to_string(),
|
||||
message: Message {
|
||||
role: Role::Assistant,
|
||||
|
|
@ -1335,6 +1360,8 @@ impl ProviderAdapter for Adapter {
|
|||
},
|
||||
finish_reason,
|
||||
usage,
|
||||
cost_usd,
|
||||
cost_source,
|
||||
raw: serde_json::from_str(&body).ok(),
|
||||
warnings: vec![],
|
||||
rate_limit: parse_rate_limit_headers(&headers),
|
||||
|
|
@ -1397,6 +1424,9 @@ impl ProviderAdapter for Adapter {
|
|||
emitted_reasoning_start: false,
|
||||
raw_response: None,
|
||||
rate_limit,
|
||||
provider_name: self.provider_name.clone(),
|
||||
catalog: self.catalog.clone(),
|
||||
speed: request.speed,
|
||||
};
|
||||
|
||||
let stream = stream::unfold(state, |mut state| async move {
|
||||
|
|
@ -2342,6 +2372,9 @@ mod tests {
|
|||
emitted_reasoning_start: false,
|
||||
raw_response: None,
|
||||
rate_limit: None,
|
||||
provider_name: "openai".to_string(),
|
||||
catalog: None,
|
||||
speed: None,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -12,6 +12,8 @@ pub(crate) mod stream;
|
|||
pub(crate) mod translate;
|
||||
pub(crate) mod wire;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use fabro_model::Catalog;
|
||||
pub(crate) use hooks::ChatHooks;
|
||||
|
||||
|
|
@ -45,14 +47,25 @@ pub(crate) async fn complete(
|
|||
}
|
||||
let (body, headers) = send_and_read_response(req, provider_name, "type").await?;
|
||||
|
||||
response::parse_chat_response(&body, &headers, provider_name, request, hooks)
|
||||
response::parse_chat_response(
|
||||
&body,
|
||||
&headers,
|
||||
provider_name,
|
||||
request,
|
||||
hooks,
|
||||
catalog,
|
||||
request.speed,
|
||||
)
|
||||
}
|
||||
|
||||
/// Run a streaming Chat Completions request through the shared pipeline.
|
||||
///
|
||||
/// `catalog` is taken by `Arc` (not borrow) because the streaming
|
||||
/// `StreamState` must own it for the lifetime of the spawned stream.
|
||||
pub(crate) async fn stream(
|
||||
http: &super::http_api::HttpApi,
|
||||
build_request: impl Fn(&str) -> fabro_http::RequestBuilder + Send,
|
||||
catalog: Option<&Catalog>,
|
||||
catalog: Option<Arc<Catalog>>,
|
||||
provider_name: &str,
|
||||
request: &Request,
|
||||
hooks: ChatHooks,
|
||||
|
|
@ -61,7 +74,7 @@ pub(crate) async fn stream(
|
|||
request,
|
||||
Some(true),
|
||||
provider_name,
|
||||
catalog,
|
||||
catalog.as_deref(),
|
||||
hooks,
|
||||
);
|
||||
let url = format!("{}/chat/completions", http.base_url);
|
||||
|
|
@ -76,6 +89,8 @@ pub(crate) async fn stream(
|
|||
http.stream_read_timeout,
|
||||
hooks,
|
||||
custom_tool_names,
|
||||
catalog,
|
||||
request.speed,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,10 +2,12 @@
|
|||
//! [`Response`](crate::types::Response).
|
||||
|
||||
use fabro_http::HeaderMap;
|
||||
use fabro_model::{Catalog, Speed};
|
||||
|
||||
use super::hooks::ChatHooks;
|
||||
use super::translate::{custom_tool_names, map_finish_reason, parse_tool_arguments};
|
||||
use super::wire::ApiResponse;
|
||||
use crate::cost::estimate_cost_usd;
|
||||
use crate::error::{Error, ProviderErrorDetail, ProviderErrorKind};
|
||||
use crate::providers::common::parse_rate_limit_headers;
|
||||
use crate::types::{
|
||||
|
|
@ -21,6 +23,8 @@ pub(crate) fn parse_chat_response(
|
|||
provider_name: &str,
|
||||
request: &Request,
|
||||
hooks: ChatHooks,
|
||||
catalog: Option<&Catalog>,
|
||||
speed: Option<Speed>,
|
||||
) -> Result<Response, Error> {
|
||||
let api_resp: ApiResponse = serde_json::from_str(body)
|
||||
.map_err(|e| Error::network(format!("failed to parse response: {e}"), e))?;
|
||||
|
|
@ -81,6 +85,9 @@ pub(crate) fn parse_chat_response(
|
|||
|
||||
let raw: Option<serde_json::Value> = serde_json::from_str(body).ok();
|
||||
|
||||
let (cost_usd, cost_source) =
|
||||
estimate_cost_usd(catalog, provider_name, &api_resp.model, &usage, speed);
|
||||
|
||||
let mut response = Response {
|
||||
id: api_resp.id,
|
||||
model: api_resp.model,
|
||||
|
|
@ -93,6 +100,8 @@ pub(crate) fn parse_chat_response(
|
|||
},
|
||||
finish_reason,
|
||||
usage,
|
||||
cost_usd,
|
||||
cost_source,
|
||||
raw: raw.clone(),
|
||||
warnings: vec![],
|
||||
rate_limit: parse_rate_limit_headers(headers),
|
||||
|
|
|
|||
|
|
@ -1,10 +1,14 @@
|
|||
//! Streaming Chat Completions SSE parsing and finish-event assembly.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use fabro_model::{Catalog, Speed};
|
||||
use futures::{StreamExt, stream};
|
||||
|
||||
use super::hooks::ChatHooks;
|
||||
use super::translate::{map_finish_reason, parse_tool_arguments};
|
||||
use super::wire::{AccumulatedToolCall, StreamChunk};
|
||||
use crate::cost::estimate_cost_usd;
|
||||
use crate::error::{Error, error_from_status_code};
|
||||
use crate::provider::StreamEventStream;
|
||||
use crate::providers::common::{
|
||||
|
|
@ -50,6 +54,10 @@ pub(crate) struct StreamState {
|
|||
/// Names of tools on the request marked custom (freeform), used by
|
||||
/// [`parse_tool_arguments`] to preserve raw non-JSON arguments.
|
||||
pub(crate) custom_tool_names: Vec<String>,
|
||||
/// Optional catalog used to estimate `cost_usd` on the final Response.
|
||||
pub(crate) catalog: Option<Arc<Catalog>>,
|
||||
/// Speed setting forwarded from the request, used by catalog pricing.
|
||||
pub(crate) speed: Option<Speed>,
|
||||
}
|
||||
|
||||
impl StreamState {
|
||||
|
|
@ -61,6 +69,8 @@ impl StreamState {
|
|||
stream_read_timeout: Option<std::time::Duration>,
|
||||
hooks: ChatHooks,
|
||||
custom_tool_names: Vec<String>,
|
||||
catalog: Option<Arc<Catalog>>,
|
||||
speed: Option<Speed>,
|
||||
) -> Self {
|
||||
Self {
|
||||
line_reader: LineReader::new(response, stream_read_timeout),
|
||||
|
|
@ -80,6 +90,8 @@ impl StreamState {
|
|||
hooks,
|
||||
last_usage_raw: None,
|
||||
custom_tool_names,
|
||||
catalog,
|
||||
speed,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -271,21 +283,31 @@ impl StreamState {
|
|||
self.response_model.clone()
|
||||
};
|
||||
|
||||
let (cost_usd, cost_source) = estimate_cost_usd(
|
||||
self.catalog.as_deref(),
|
||||
&self.provider_name,
|
||||
&response_model,
|
||||
&self.usage,
|
||||
self.speed,
|
||||
);
|
||||
|
||||
let mut response = Response {
|
||||
id: self.response_id.clone(),
|
||||
model: response_model,
|
||||
provider: self.provider_name.clone(),
|
||||
message: Message {
|
||||
id: self.response_id.clone(),
|
||||
model: response_model,
|
||||
provider: self.provider_name.clone(),
|
||||
message: Message {
|
||||
role: Role::Assistant,
|
||||
content: content_parts,
|
||||
name: None,
|
||||
tool_call_id: None,
|
||||
},
|
||||
finish_reason: self.finish_reason.clone(),
|
||||
usage: self.usage.clone(),
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: self.rate_limit.clone(),
|
||||
usage: self.usage.clone(),
|
||||
cost_usd,
|
||||
cost_source,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: self.rate_limit.clone(),
|
||||
};
|
||||
|
||||
if let Some(enrich) = self.hooks.enrich_response {
|
||||
|
|
@ -317,6 +339,8 @@ pub(crate) fn run_stream(
|
|||
stream_read_timeout: Option<std::time::Duration>,
|
||||
hooks: ChatHooks,
|
||||
custom_tool_names: Vec<String>,
|
||||
catalog: Option<Arc<Catalog>>,
|
||||
speed: Option<Speed>,
|
||||
) -> StreamEventStream {
|
||||
let stream = stream::unfold(
|
||||
StreamState::new(
|
||||
|
|
@ -327,6 +351,8 @@ pub(crate) fn run_stream(
|
|||
stream_read_timeout,
|
||||
hooks,
|
||||
custom_tool_names,
|
||||
catalog,
|
||||
speed,
|
||||
),
|
||||
|mut state| async move {
|
||||
loop {
|
||||
|
|
@ -430,6 +456,8 @@ pub(crate) async fn send_and_stream(
|
|||
stream_read_timeout: Option<std::time::Duration>,
|
||||
hooks: ChatHooks,
|
||||
custom_tool_names: Vec<String>,
|
||||
catalog: Option<Arc<Catalog>>,
|
||||
speed: Option<Speed>,
|
||||
) -> Result<StreamEventStream, Error> {
|
||||
let http_resp = req
|
||||
.send()
|
||||
|
|
@ -464,6 +492,8 @@ pub(crate) async fn send_and_stream(
|
|||
stream_read_timeout,
|
||||
hooks,
|
||||
custom_tool_names,
|
||||
catalog,
|
||||
speed,
|
||||
))
|
||||
}
|
||||
|
||||
|
|
@ -548,6 +578,8 @@ mod tests {
|
|||
Some(std::time::Duration::from_secs(30)),
|
||||
ChatHooks::NONE,
|
||||
Vec::new(),
|
||||
None,
|
||||
None,
|
||||
);
|
||||
|
||||
let raw: serde_json::Value = serde_json::from_str(
|
||||
|
|
@ -582,6 +614,8 @@ mod tests {
|
|||
Some(std::time::Duration::from_secs(30)),
|
||||
ChatHooks::NONE,
|
||||
Vec::new(),
|
||||
None,
|
||||
None,
|
||||
);
|
||||
|
||||
// First text chunk should emit TextStart + TextDelta.
|
||||
|
|
@ -618,6 +652,8 @@ mod tests {
|
|||
Some(std::time::Duration::from_secs(30)),
|
||||
ChatHooks::NONE,
|
||||
Vec::new(),
|
||||
None,
|
||||
None,
|
||||
);
|
||||
|
||||
// First tool call chunk (has id and name) -> ToolCallStart.
|
||||
|
|
@ -653,6 +689,8 @@ mod tests {
|
|||
Some(std::time::Duration::from_secs(30)),
|
||||
ChatHooks::NONE,
|
||||
Vec::new(),
|
||||
None,
|
||||
None,
|
||||
);
|
||||
state.response_id = "resp-1".into();
|
||||
state.response_model = "gpt-4".into();
|
||||
|
|
@ -698,6 +736,8 @@ mod tests {
|
|||
Some(std::time::Duration::from_secs(30)),
|
||||
ChatHooks::NONE,
|
||||
Vec::new(),
|
||||
None,
|
||||
None,
|
||||
);
|
||||
state.response_id = "resp-1".into();
|
||||
state.tool_calls.push(AccumulatedToolCall {
|
||||
|
|
@ -745,6 +785,8 @@ mod tests {
|
|||
Some(std::time::Duration::from_secs(30)),
|
||||
ChatHooks::NONE,
|
||||
Vec::new(),
|
||||
None,
|
||||
None,
|
||||
);
|
||||
// response_model is empty, so finish_events should use the request model.
|
||||
let events = state.finish_events();
|
||||
|
|
|
|||
|
|
@ -111,7 +111,7 @@ impl ProviderAdapter for Adapter {
|
|||
openai_chat::stream(
|
||||
&self.http,
|
||||
|url| self.build_request(url),
|
||||
self.catalog.as_deref(),
|
||||
self.catalog.clone(),
|
||||
&self.provider_name,
|
||||
request,
|
||||
ChatHooks::NONE,
|
||||
|
|
|
|||
|
|
@ -100,6 +100,33 @@ impl Message {
|
|||
}
|
||||
}
|
||||
|
||||
// --- 3.7.1 CostSource ---
|
||||
|
||||
/// Source of the `cost_usd` value on a [`Response`].
|
||||
#[derive(
|
||||
Debug,
|
||||
Clone,
|
||||
Copy,
|
||||
PartialEq,
|
||||
Eq,
|
||||
Serialize,
|
||||
Deserialize,
|
||||
strum::Display,
|
||||
strum::EnumString,
|
||||
strum::IntoStaticStr,
|
||||
strum::VariantArray,
|
||||
)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
#[strum(serialize_all = "snake_case")]
|
||||
pub enum CostSource {
|
||||
/// Provider returned an inline authoritative cost (e.g. OpenRouter
|
||||
/// `usage.cost`). Use this when reconciling billing.
|
||||
Authoritative,
|
||||
/// Computed from `tokens × catalog_price`. Approximation; can drift if
|
||||
/// catalog prices are stale or the provider applies surcharges/credits.
|
||||
Estimated,
|
||||
}
|
||||
|
||||
// --- 3.8 FinishReason ---
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
|
|
@ -311,6 +338,10 @@ pub struct Response {
|
|||
pub message: Message,
|
||||
pub finish_reason: FinishReason,
|
||||
pub usage: TokenCounts,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cost_usd: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cost_source: Option<CostSource>,
|
||||
pub raw: Option<serde_json::Value>,
|
||||
pub warnings: Vec<Warning>,
|
||||
pub rate_limit: Option<RateLimitInfo>,
|
||||
|
|
@ -782,6 +813,8 @@ mod tests {
|
|||
message: Message::assistant("Hello world"),
|
||||
finish_reason: FinishReason::Stop,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -810,6 +843,8 @@ mod tests {
|
|||
},
|
||||
finish_reason: FinishReason::ToolCalls,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -841,6 +876,8 @@ mod tests {
|
|||
},
|
||||
finish_reason: FinishReason::Stop,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -858,6 +895,8 @@ mod tests {
|
|||
message: Message::assistant("Hello"),
|
||||
finish_reason: FinishReason::Stop,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -1028,6 +1067,8 @@ mod tests {
|
|||
output_tokens: 5,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
|
|||
|
|
@ -352,6 +352,8 @@ mod tests {
|
|||
message: Message::assistant(self.response_text.clone()),
|
||||
finish_reason: FinishReason::Stop,
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: Vec::new(),
|
||||
rate_limit: None,
|
||||
|
|
|
|||
|
|
@ -29,20 +29,20 @@ pub use fabro_api::types::{
|
|||
BatchRunLifecycleResponse, BatchRunLifecycleResult, BatchRunLifecycleResultOutcome,
|
||||
BatchRunLifecycleSummary, BillingByModel, BillingStageRef, CloseRunPullRequestResponse,
|
||||
CompletionContentPart, CompletionMessage, CompletionMessageRole, CompletionResponse,
|
||||
CompletionToolChoiceMode, CompletionUsage, CreateCompletionRequest,
|
||||
CreateRunPullRequestRequest, CreateSecretRequest, CreateVariableRequest, DeleteRunResponse,
|
||||
DeleteRunSandbox, DeleteSecretRequest, DenyRunRequest, DiskUsageResponse, DiskUsageRunRow,
|
||||
DiskUsageSummaryRow, ErrorResponseEntry, ForkRequest, ForkResponse, IntegrationConnectionKind,
|
||||
IntegrationConnectionState, IntegrationConnectionStatus, IntegrationProvider,
|
||||
IntegrationStatus, LinkRunPullRequestRequest, MergeRunPullRequestRequest,
|
||||
MergeRunPullRequestResponse, ModelReference, PaginatedEventList, PaginatedRunList,
|
||||
PaginationMeta, PreflightResponse, PreviewUrlRequest, PreviewUrlResponse, Provider,
|
||||
ProviderList, PruneRunEntry, PruneRunsRequest, PruneRunsResponse, RenderWorkflowGraphDirection,
|
||||
RenderWorkflowGraphRequest, RewindRequest, RewindResponse, Run, RunArtifactEntry,
|
||||
RunArtifactListResponse, RunBilling, RunBillingStage, RunBillingTotals, RunError, RunManifest,
|
||||
RunStage, SandboxDetails, SandboxFileEntry, SandboxFileListResponse, SandboxService,
|
||||
SandboxServiceListResponse, SshAccessRequest, SshAccessResponse, StageHandler, StageState,
|
||||
StartRunRequest, SubmitAnswerRequest, SystemCpuResourceScope, SystemCpuResources,
|
||||
CompletionResponseCostSource, CompletionToolChoiceMode, CompletionUsage,
|
||||
CreateCompletionRequest, CreateRunPullRequestRequest, CreateSecretRequest,
|
||||
CreateVariableRequest, DeleteRunResponse, DeleteRunSandbox, DeleteSecretRequest,
|
||||
DenyRunRequest, DiskUsageResponse, DiskUsageRunRow, DiskUsageSummaryRow, ErrorResponseEntry,
|
||||
ForkRequest, ForkResponse, IntegrationConnectionKind, IntegrationConnectionState,
|
||||
IntegrationConnectionStatus, IntegrationProvider, IntegrationStatus, LinkRunPullRequestRequest,
|
||||
MergeRunPullRequestRequest, MergeRunPullRequestResponse, ModelReference, PaginatedEventList,
|
||||
PaginatedRunList, PaginationMeta, PreflightResponse, PreviewUrlRequest, PreviewUrlResponse,
|
||||
Provider, ProviderList, PruneRunEntry, PruneRunsRequest, PruneRunsResponse,
|
||||
RenderWorkflowGraphDirection, RenderWorkflowGraphRequest, RewindRequest, RewindResponse, Run,
|
||||
RunArtifactEntry, RunArtifactListResponse, RunBilling, RunBillingStage, RunBillingTotals,
|
||||
RunError, RunManifest, RunStage, SandboxDetails, SandboxFileEntry, SandboxFileListResponse,
|
||||
SandboxService, SandboxServiceListResponse, SshAccessRequest, SshAccessResponse, StageHandler,
|
||||
StageState, StartRunRequest, SubmitAnswerRequest, SystemCpuResourceScope, SystemCpuResources,
|
||||
SystemDiskResourceScope, SystemDiskResources, SystemInfoResponse, SystemIntegrationStatus,
|
||||
SystemIntegrationsResponse, SystemMemoryResourceScope, SystemMemoryResources,
|
||||
SystemRepairRunIssue, SystemRepairRunsResponse, SystemResourcesResponse, SystemRunCounts,
|
||||
|
|
|
|||
|
|
@ -1,11 +1,14 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use fabro_llm::types::CostSource;
|
||||
|
||||
use super::super::{
|
||||
ApiError, AppState, CompletionContentPart, CompletionMessage, CompletionMessageRole,
|
||||
CompletionResponse, CompletionToolChoiceMode, CompletionUsage, ContentPart,
|
||||
CreateCompletionRequest, Duration, Event, FinishReason, GenerateParams, IntoResponse, Json,
|
||||
KeepAlive, LlmMessage, LlmRequest, RequiredUser, Response, Role, Router, Sse, State,
|
||||
StatusCode, ToolChoice, ToolDefinition, Ulid, error, generate_object, info, post, warn,
|
||||
CompletionResponse, CompletionResponseCostSource, CompletionToolChoiceMode, CompletionUsage,
|
||||
ContentPart, CreateCompletionRequest, Duration, Event, FinishReason, GenerateParams,
|
||||
IntoResponse, Json, KeepAlive, LlmMessage, LlmRequest, RequiredUser, Response, Role, Router,
|
||||
Sse, State, StatusCode, ToolChoice, ToolDefinition, Ulid, error, generate_object, info, post,
|
||||
warn,
|
||||
};
|
||||
|
||||
pub(super) fn routes() -> Router<Arc<AppState>> {
|
||||
|
|
@ -23,6 +26,13 @@ fn finish_reason_to_api_stop_reason(reason: &FinishReason) -> String {
|
|||
}
|
||||
}
|
||||
|
||||
fn cost_source_to_api(source: CostSource) -> CompletionResponseCostSource {
|
||||
match source {
|
||||
CostSource::Authoritative => CompletionResponseCostSource::Authoritative,
|
||||
CostSource::Estimated => CompletionResponseCostSource::Estimated,
|
||||
}
|
||||
}
|
||||
|
||||
fn convert_api_message(msg: &CompletionMessage) -> LlmMessage {
|
||||
let role = match msg.role {
|
||||
CompletionMessageRole::System => Role::System,
|
||||
|
|
@ -254,6 +264,8 @@ async fn create_completion(
|
|||
output_tokens: result.usage.output_tokens,
|
||||
},
|
||||
output: result.output,
|
||||
cost_usd: result.response.cost_usd,
|
||||
cost_source: result.response.cost_source.map(cost_source_to_api),
|
||||
})
|
||||
.into_response(),
|
||||
Err(e) => ApiError::new(StatusCode::BAD_GATEWAY, format!("LLM error: {e}"))
|
||||
|
|
@ -271,6 +283,8 @@ async fn create_completion(
|
|||
output_tokens: response.usage.output_tokens,
|
||||
},
|
||||
output: None,
|
||||
cost_usd: response.cost_usd,
|
||||
cost_source: response.cost_source.map(cost_source_to_api),
|
||||
})
|
||||
.into_response(),
|
||||
Err(e) => ApiError::new(StatusCode::BAD_GATEWAY, format!("LLM error: {e}"))
|
||||
|
|
|
|||
|
|
@ -729,6 +729,8 @@ mod tests {
|
|||
output_tokens: 20,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
@ -757,6 +759,8 @@ mod tests {
|
|||
output_tokens: 20,
|
||||
..Default::default()
|
||||
},
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
raw: None,
|
||||
warnings: vec![],
|
||||
rate_limit: None,
|
||||
|
|
|
|||
|
|
@ -30,4 +30,19 @@ export interface CompletionResponse {
|
|||
'stop_reason': string;
|
||||
'usage': CompletionUsage;
|
||||
'output'?: any;
|
||||
/**
|
||||
* Total USD cost for the completion. Populated when the provider returns an authoritative billing figure (e.g. OpenRouter) or when the model has catalog pricing.
|
||||
*/
|
||||
'cost_usd'?: number;
|
||||
/**
|
||||
* Whether `cost_usd` came from the provider\'s billing data (`authoritative`) or was computed from catalog prices (`estimated`).
|
||||
*/
|
||||
'cost_source'?: CompletionResponseCostSourceEnum;
|
||||
}
|
||||
|
||||
export const CompletionResponseCostSourceEnum = {
|
||||
AUTHORITATIVE: 'authoritative',
|
||||
ESTIMATED: 'estimated'
|
||||
} as const;
|
||||
|
||||
export type CompletionResponseCostSourceEnum = typeof CompletionResponseCostSourceEnum[keyof typeof CompletionResponseCostSourceEnum];
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue