mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-09-29 01:42:21 +00:00
refactor(llm): extract chat plumbing to openai_chat submodule
Splits the OpenAI Chat Completions wire/translate/request/response/stream layer out of openai_compatible.rs into a new providers/openai_chat/ submodule (wire, translate, request, response, stream, hooks). The shared module exposes pub(crate) complete() and stream() entry points that accept a ChatHooks struct of optional fn pointers, letting future adapters (OpenRouter) layer provider-specific behavior on top without duplicating the wire layer. openai_compatible.rs shrinks from 1649 to 227 lines: the ProviderAdapter impl is now a one-line delegator passing ChatHooks::NONE, which is behaviorally identical to the pre-refactor code path. No behavior change. All 392 lib tests pass. Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
ea9a1467a9
commit
7624481d8f
9 changed files with 1826 additions and 1491 deletions
|
|
@ -4,6 +4,7 @@ pub mod fabro_server;
|
|||
pub mod gemini;
|
||||
pub mod http_api;
|
||||
pub mod openai;
|
||||
pub(crate) mod openai_chat;
|
||||
pub mod openai_compatible;
|
||||
|
||||
pub use anthropic::Adapter as AnthropicAdapter;
|
||||
|
|
|
|||
37
lib/crates/fabro-llm/src/providers/openai_chat/hooks.rs
Normal file
37
lib/crates/fabro-llm/src/providers/openai_chat/hooks.rs
Normal file
|
|
@ -0,0 +1,37 @@
|
|||
//! Extension hooks that let provider-specific adapters (e.g. OpenRouter)
|
||||
//! inject behavior into the shared OpenAI Chat Completions pipeline without
|
||||
//! duplicating the wire layer.
|
||||
|
||||
use crate::types::{Request, Response};
|
||||
|
||||
/// Optional pre-send and post-receive hooks for the chat pipeline.
|
||||
///
|
||||
/// Each field is an opt-in `fn` pointer. The default-constructed value
|
||||
/// ([`Self::NONE`]) is behaviorally identical to no hooks at all, which is
|
||||
/// what `OpenAiCompatibleAdapter` uses.
|
||||
#[derive(Clone, Copy)]
|
||||
pub(crate) struct ChatHooks {
|
||||
/// Mutate the final JSON body just before sending. Used by OpenRouter
|
||||
/// to translate typed `Request.reasoning_effort` into OR's
|
||||
/// `{reasoning: {effort: ...}}` shape (and any other request-shape
|
||||
/// translations).
|
||||
pub(crate) mutate_request: Option<fn(&mut serde_json::Value, &Request)>,
|
||||
/// Read provider-specific fields out of the raw response body and
|
||||
/// attach them to the unified [`Response`]. Used by OpenRouter to
|
||||
/// extract authoritative `usage.cost`.
|
||||
pub(crate) enrich_response: Option<fn(&mut Response, &serde_json::Value)>,
|
||||
}
|
||||
|
||||
impl ChatHooks {
|
||||
/// No-op hooks. Behaviorally identical to no hooks.
|
||||
pub(crate) const NONE: Self = Self {
|
||||
mutate_request: None,
|
||||
enrich_response: None,
|
||||
};
|
||||
}
|
||||
|
||||
impl Default for ChatHooks {
|
||||
fn default() -> Self {
|
||||
Self::NONE
|
||||
}
|
||||
}
|
||||
81
lib/crates/fabro-llm/src/providers/openai_chat/mod.rs
Normal file
81
lib/crates/fabro-llm/src/providers/openai_chat/mod.rs
Normal file
|
|
@ -0,0 +1,81 @@
|
|||
//! Shared implementation of the OpenAI Chat Completions wire protocol.
|
||||
//!
|
||||
//! Used by [`super::openai_compatible::Adapter`] for vanilla
|
||||
//! OpenAI-compatible providers (Together, Groq, vLLM, etc.) and (after
|
||||
//! PR C) by `super::openrouter::Adapter` to layer OR-specific behavior on
|
||||
//! top via [`ChatHooks`].
|
||||
|
||||
pub(crate) mod hooks;
|
||||
pub(crate) mod request;
|
||||
pub(crate) mod response;
|
||||
pub(crate) mod stream;
|
||||
pub(crate) mod translate;
|
||||
pub(crate) mod wire;
|
||||
|
||||
use fabro_model::Catalog;
|
||||
pub(crate) use hooks::ChatHooks;
|
||||
|
||||
use crate::error::Error;
|
||||
use crate::provider::StreamEventStream;
|
||||
use crate::providers::common::send_and_read_response;
|
||||
use crate::types::{Request, Response};
|
||||
|
||||
/// Run a non-streaming Chat Completions request through the shared
|
||||
/// pipeline, applying the supplied [`ChatHooks`] for request mutation and
|
||||
/// response enrichment.
|
||||
///
|
||||
/// `build_request` is a closure that returns an authenticated
|
||||
/// `fabro_http::RequestBuilder` for a given URL — each caller (OpenAI
|
||||
/// compatible, OpenRouter) handles auth header construction itself.
|
||||
pub(crate) async fn complete(
|
||||
http: &super::http_api::HttpApi,
|
||||
build_request: impl Fn(&str) -> fabro_http::RequestBuilder + Send,
|
||||
catalog: Option<&Catalog>,
|
||||
provider_name: &str,
|
||||
request: &Request,
|
||||
hooks: ChatHooks,
|
||||
) -> Result<Response, Error> {
|
||||
let api_body =
|
||||
request::build_chat_request_with_catalog(request, None, provider_name, catalog, hooks);
|
||||
let url = format!("{}/chat/completions", http.base_url);
|
||||
|
||||
let mut req = build_request(&url).json(&api_body);
|
||||
if let Some(t) = http.request_timeout {
|
||||
req = req.timeout(t);
|
||||
}
|
||||
let (body, headers) = send_and_read_response(req, provider_name, "type").await?;
|
||||
|
||||
response::parse_chat_response(&body, &headers, provider_name, request, hooks)
|
||||
}
|
||||
|
||||
/// Run a streaming Chat Completions request through the shared pipeline.
|
||||
pub(crate) async fn stream(
|
||||
http: &super::http_api::HttpApi,
|
||||
build_request: impl Fn(&str) -> fabro_http::RequestBuilder + Send,
|
||||
catalog: Option<&Catalog>,
|
||||
provider_name: &str,
|
||||
request: &Request,
|
||||
hooks: ChatHooks,
|
||||
) -> Result<StreamEventStream, Error> {
|
||||
let api_body = request::build_chat_request_with_catalog(
|
||||
request,
|
||||
Some(true),
|
||||
provider_name,
|
||||
catalog,
|
||||
hooks,
|
||||
);
|
||||
let url = format!("{}/chat/completions", http.base_url);
|
||||
|
||||
let req = build_request(&url).json(&api_body);
|
||||
|
||||
let custom_tool_names = translate::custom_tool_names(request);
|
||||
stream::send_and_stream(
|
||||
req,
|
||||
provider_name.to_string(),
|
||||
request.model.clone(),
|
||||
http.stream_read_timeout,
|
||||
hooks,
|
||||
custom_tool_names,
|
||||
)
|
||||
.await
|
||||
}
|
||||
214
lib/crates/fabro-llm/src/providers/openai_chat/request.rs
Normal file
214
lib/crates/fabro-llm/src/providers/openai_chat/request.rs
Normal file
|
|
@ -0,0 +1,214 @@
|
|||
//! Build a Chat Completions request body from a unified
|
||||
//! [`Request`](crate::types::Request).
|
||||
|
||||
use fabro_model::Catalog;
|
||||
|
||||
use super::hooks::ChatHooks;
|
||||
use super::translate::{
|
||||
translate_messages, translate_response_format, translate_tool_choice, translate_tools,
|
||||
};
|
||||
use super::wire::ApiRequest;
|
||||
use crate::providers::common::api_model_id;
|
||||
use crate::types::Request;
|
||||
|
||||
/// Build the API request body from a unified `Request`.
|
||||
///
|
||||
/// Returns a `serde_json::Value` so that `provider_options.<provider_name>`
|
||||
/// fields can be merged into the request before sending, and so that
|
||||
/// [`ChatHooks::mutate_request`] can apply provider-specific shape
|
||||
/// translations.
|
||||
pub(crate) fn build_chat_request_with_catalog(
|
||||
request: &Request,
|
||||
stream: Option<bool>,
|
||||
provider_name: &str,
|
||||
catalog: Option<&Catalog>,
|
||||
hooks: ChatHooks,
|
||||
) -> serde_json::Value {
|
||||
let chat_messages = translate_messages(&request.messages);
|
||||
let tools = request.tools.as_ref().map(|t| translate_tools(t));
|
||||
let tool_choice = request.tool_choice.as_ref().map(translate_tool_choice);
|
||||
let response_format = request
|
||||
.response_format
|
||||
.as_ref()
|
||||
.map(translate_response_format);
|
||||
|
||||
let api_request = ApiRequest {
|
||||
model: api_model_id(catalog, &request.model),
|
||||
messages: chat_messages,
|
||||
temperature: request.temperature,
|
||||
max_tokens: request.max_tokens,
|
||||
top_p: request.top_p,
|
||||
stop: request.stop_sequences.clone(),
|
||||
tools,
|
||||
tool_choice,
|
||||
response_format,
|
||||
stream,
|
||||
};
|
||||
|
||||
let mut body = serde_json::to_value(&api_request).unwrap_or_default();
|
||||
merge_provider_options(&mut body, request.provider_options.as_ref(), provider_name);
|
||||
if let Some(mutate) = hooks.mutate_request {
|
||||
mutate(&mut body, request);
|
||||
}
|
||||
body
|
||||
}
|
||||
|
||||
/// Merge `provider_options.<provider_name>` fields into the serialized API
|
||||
/// request body.
|
||||
///
|
||||
/// The provider name is configurable (e.g. "groq", "together",
|
||||
/// "openai-compatible"), allowing each instance to have its own namespace in
|
||||
/// `provider_options`.
|
||||
pub(crate) fn merge_provider_options(
|
||||
body: &mut serde_json::Value,
|
||||
provider_options: Option<&serde_json::Value>,
|
||||
provider_name: &str,
|
||||
) {
|
||||
let Some(opts) = provider_options.and_then(|opts| opts.get(provider_name)) else {
|
||||
return;
|
||||
};
|
||||
let Some(body_map) = body.as_object_mut() else {
|
||||
return;
|
||||
};
|
||||
let Some(opts_map) = opts.as_object() else {
|
||||
return;
|
||||
};
|
||||
|
||||
for (key, value) in opts_map {
|
||||
body_map.insert(key.clone(), value.clone());
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn build_chat_request(
|
||||
request: &Request,
|
||||
stream: Option<bool>,
|
||||
provider_name: &str,
|
||||
) -> serde_json::Value {
|
||||
build_chat_request_with_catalog(request, stream, provider_name, None, ChatHooks::NONE)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use fabro_model::catalog::LlmCatalogSettings;
|
||||
|
||||
use super::*;
|
||||
use crate::providers::openai_chat::wire::ApiRequest;
|
||||
use crate::types::Message;
|
||||
|
||||
fn minimal_request() -> Request {
|
||||
Request {
|
||||
model: "llama-3.1-70b".to_string(),
|
||||
messages: vec![Message::user("Hello")],
|
||||
provider: None,
|
||||
tools: None,
|
||||
tool_choice: None,
|
||||
response_format: None,
|
||||
temperature: None,
|
||||
top_p: None,
|
||||
max_tokens: None,
|
||||
stop_sequences: None,
|
||||
reasoning_effort: None,
|
||||
speed: None,
|
||||
metadata: None,
|
||||
provider_options: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn api_request_stream_field_serialization() {
|
||||
let req = ApiRequest {
|
||||
model: "test".into(),
|
||||
messages: vec![],
|
||||
temperature: None,
|
||||
max_tokens: None,
|
||||
top_p: None,
|
||||
stop: None,
|
||||
tools: None,
|
||||
tool_choice: None,
|
||||
response_format: None,
|
||||
stream: Some(true),
|
||||
};
|
||||
let json = serde_json::to_value(&req).unwrap();
|
||||
assert_eq!(json["stream"], true);
|
||||
|
||||
// When stream is None, it should be omitted.
|
||||
let req_no_stream = ApiRequest {
|
||||
model: "test".into(),
|
||||
messages: vec![],
|
||||
temperature: None,
|
||||
max_tokens: None,
|
||||
top_p: None,
|
||||
stop: None,
|
||||
tools: None,
|
||||
tool_choice: None,
|
||||
response_format: None,
|
||||
stream: None,
|
||||
};
|
||||
let json_no_stream = serde_json::to_value(&req_no_stream).unwrap();
|
||||
assert!(json_no_stream.get("stream").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_options_none_produces_standard_body() {
|
||||
let request = minimal_request();
|
||||
let body = build_chat_request(&request, None, "groq");
|
||||
assert_eq!(body["model"], "llama-3.1-70b");
|
||||
assert!(body.get("stream").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn catalog_api_id_is_used_for_provider_request_body() {
|
||||
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"]
|
||||
|
||||
[models."acme-large"]
|
||||
provider = "acme"
|
||||
api_id = "acme/model-large"
|
||||
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
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
let catalog = fabro_model::Catalog::from_builtin_with_overrides(&settings).unwrap();
|
||||
let mut request = minimal_request();
|
||||
request.model = "acme-large".to_string();
|
||||
|
||||
let body = build_chat_request_with_catalog(
|
||||
&request,
|
||||
None,
|
||||
"acme",
|
||||
Some(&catalog),
|
||||
ChatHooks::NONE,
|
||||
);
|
||||
|
||||
assert_eq!(request.model, "acme-large");
|
||||
assert_eq!(body["model"], "acme/model-large");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merge_provider_options_with_non_object_value() {
|
||||
let mut body = serde_json::json!({"model": "test"});
|
||||
let opts = serde_json::json!({"groq": "not-an-object"});
|
||||
merge_provider_options(&mut body, Some(&opts), "groq");
|
||||
// Should not crash and body should be unchanged
|
||||
assert_eq!(body["model"], "test");
|
||||
}
|
||||
}
|
||||
107
lib/crates/fabro-llm/src/providers/openai_chat/response.rs
Normal file
107
lib/crates/fabro-llm/src/providers/openai_chat/response.rs
Normal file
|
|
@ -0,0 +1,107 @@
|
|||
//! Parse a non-streaming Chat Completions HTTP response body into a unified
|
||||
//! [`Response`](crate::types::Response).
|
||||
|
||||
use fabro_http::HeaderMap;
|
||||
|
||||
use super::hooks::ChatHooks;
|
||||
use super::translate::{custom_tool_names, map_finish_reason, parse_tool_arguments};
|
||||
use super::wire::ApiResponse;
|
||||
use crate::error::{Error, ProviderErrorDetail, ProviderErrorKind};
|
||||
use crate::providers::common::parse_rate_limit_headers;
|
||||
use crate::types::{
|
||||
ContentPart, Message, Request, Response, Role, ThinkingData, TokenCounts, ToolCall,
|
||||
};
|
||||
|
||||
/// Parse a non-streaming Chat Completions response body into a unified
|
||||
/// [`Response`]. Applies [`ChatHooks::enrich_response`] after constructing
|
||||
/// the response, before returning it.
|
||||
pub(crate) fn parse_chat_response(
|
||||
body: &str,
|
||||
headers: &HeaderMap,
|
||||
provider_name: &str,
|
||||
request: &Request,
|
||||
hooks: ChatHooks,
|
||||
) -> Result<Response, Error> {
|
||||
let api_resp: ApiResponse = serde_json::from_str(body)
|
||||
.map_err(|e| Error::network(format!("failed to parse response: {e}"), e))?;
|
||||
|
||||
let choice = api_resp.choices.first().ok_or_else(|| Error::Provider {
|
||||
kind: ProviderErrorKind::Server,
|
||||
detail: Box::new(ProviderErrorDetail::new(
|
||||
"no choices in response",
|
||||
provider_name,
|
||||
)),
|
||||
})?;
|
||||
|
||||
let mut content_parts = Vec::new();
|
||||
if let Some(reasoning) = &choice.message.reasoning_content {
|
||||
if !reasoning.is_empty() {
|
||||
content_parts.push(ContentPart::Thinking(ThinkingData {
|
||||
text: reasoning.clone(),
|
||||
signature: None,
|
||||
redacted: false,
|
||||
}));
|
||||
}
|
||||
}
|
||||
if let Some(text) = &choice.message.content {
|
||||
if !text.is_empty() {
|
||||
content_parts.push(ContentPart::text(text));
|
||||
}
|
||||
}
|
||||
if let Some(tool_calls) = &choice.message.tool_calls {
|
||||
let custom_tool_names = custom_tool_names(request);
|
||||
for tc in tool_calls {
|
||||
let arguments = parse_tool_arguments(
|
||||
&tc.function.name,
|
||||
&tc.function.arguments,
|
||||
&custom_tool_names,
|
||||
);
|
||||
let mut tool_call = ToolCall::new(&tc.id, &tc.function.name, arguments);
|
||||
tool_call.raw_arguments = Some(tc.function.arguments.clone());
|
||||
content_parts.push(ContentPart::ToolCall(tool_call));
|
||||
}
|
||||
}
|
||||
|
||||
let finish_reason = map_finish_reason(choice.finish_reason.as_deref());
|
||||
|
||||
let usage = api_resp
|
||||
.usage
|
||||
.as_ref()
|
||||
.map_or_else(TokenCounts::default, |u| TokenCounts {
|
||||
input_tokens: u.prompt_tokens,
|
||||
output_tokens: u.completion_tokens,
|
||||
cache_read_tokens: u
|
||||
.prompt_tokens_details
|
||||
.as_ref()
|
||||
.and_then(|d| d.cached_tokens)
|
||||
.unwrap_or(0),
|
||||
cache_write_tokens: u.cache_write_tokens.unwrap_or(0),
|
||||
..TokenCounts::default()
|
||||
});
|
||||
|
||||
let raw: Option<serde_json::Value> = serde_json::from_str(body).ok();
|
||||
|
||||
let mut response = Response {
|
||||
id: api_resp.id,
|
||||
model: api_resp.model,
|
||||
provider: provider_name.to_string(),
|
||||
message: Message {
|
||||
role: Role::Assistant,
|
||||
content: content_parts,
|
||||
name: None,
|
||||
tool_call_id: None,
|
||||
},
|
||||
finish_reason,
|
||||
usage,
|
||||
raw: raw.clone(),
|
||||
warnings: vec![],
|
||||
rate_limit: parse_rate_limit_headers(headers),
|
||||
};
|
||||
|
||||
if let Some(enrich) = hooks.enrich_response {
|
||||
let raw_for_enrich = raw.unwrap_or(serde_json::Value::Null);
|
||||
enrich(&mut response, &raw_for_enrich);
|
||||
}
|
||||
|
||||
Ok(response)
|
||||
}
|
||||
758
lib/crates/fabro-llm/src/providers/openai_chat/stream.rs
Normal file
758
lib/crates/fabro-llm/src/providers/openai_chat/stream.rs
Normal file
|
|
@ -0,0 +1,758 @@
|
|||
//! Streaming Chat Completions SSE parsing and finish-event assembly.
|
||||
|
||||
use futures::{StreamExt, stream};
|
||||
|
||||
use super::hooks::ChatHooks;
|
||||
use super::translate::{map_finish_reason, parse_tool_arguments};
|
||||
use super::wire::{AccumulatedToolCall, StreamChunk};
|
||||
use crate::error::{Error, error_from_status_code};
|
||||
use crate::provider::StreamEventStream;
|
||||
use crate::providers::common::{
|
||||
LineReader, parse_error_body, parse_rate_limit_headers, parse_retry_after,
|
||||
};
|
||||
use crate::types::{
|
||||
ContentPart, FinishReason, Message, RateLimitInfo, Response, Role, StreamEvent, ThinkingData,
|
||||
TokenCounts, ToolCall,
|
||||
};
|
||||
|
||||
/// State for flattening batched events into individual stream events.
|
||||
pub(crate) struct FlattenState {
|
||||
pub(crate) inner:
|
||||
std::pin::Pin<Box<dyn futures::Stream<Item = Result<Vec<StreamEvent>, Error>> + Send>>,
|
||||
pub(crate) pending: Vec<StreamEvent>,
|
||||
}
|
||||
|
||||
/// Accumulated state while processing the SSE stream.
|
||||
pub(crate) struct StreamState {
|
||||
pub(crate) line_reader: LineReader,
|
||||
pub(crate) provider_name: String,
|
||||
pub(crate) model: String,
|
||||
pub(crate) response_id: String,
|
||||
pub(crate) response_model: String,
|
||||
pub(crate) accumulated_text: String,
|
||||
pub(crate) accumulated_reasoning: String,
|
||||
pub(crate) tool_calls: Vec<AccumulatedToolCall>,
|
||||
pub(crate) usage: TokenCounts,
|
||||
pub(crate) finish_reason: FinishReason,
|
||||
pub(crate) text_started: bool,
|
||||
pub(crate) done: bool,
|
||||
/// True after `finish_events()` has been called (guards against
|
||||
/// duplicates).
|
||||
pub(crate) finished: bool,
|
||||
pub(crate) rate_limit: Option<RateLimitInfo>,
|
||||
/// Hooks (e.g. for enriching the final `Response` with provider-specific
|
||||
/// fields like OpenRouter's `usage.cost`).
|
||||
pub(crate) hooks: ChatHooks,
|
||||
/// Raw JSON of the most recent chunk that carried a `usage` block.
|
||||
/// Captured so that [`ChatHooks::enrich_response`] can pull
|
||||
/// provider-specific fields (e.g. `usage.cost`) out of the stream.
|
||||
pub(crate) last_usage_raw: Option<serde_json::Value>,
|
||||
/// 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>,
|
||||
}
|
||||
|
||||
impl StreamState {
|
||||
pub(crate) fn new(
|
||||
response: fabro_http::Response,
|
||||
provider_name: String,
|
||||
model: String,
|
||||
rate_limit: Option<RateLimitInfo>,
|
||||
stream_read_timeout: Option<std::time::Duration>,
|
||||
hooks: ChatHooks,
|
||||
custom_tool_names: Vec<String>,
|
||||
) -> Self {
|
||||
Self {
|
||||
line_reader: LineReader::new(response, stream_read_timeout),
|
||||
provider_name,
|
||||
model,
|
||||
response_id: String::new(),
|
||||
response_model: String::new(),
|
||||
accumulated_text: String::new(),
|
||||
accumulated_reasoning: String::new(),
|
||||
tool_calls: Vec::new(),
|
||||
usage: TokenCounts::default(),
|
||||
finish_reason: FinishReason::Stop,
|
||||
text_started: false,
|
||||
done: false,
|
||||
finished: false,
|
||||
rate_limit,
|
||||
hooks,
|
||||
last_usage_raw: None,
|
||||
custom_tool_names,
|
||||
}
|
||||
}
|
||||
|
||||
/// Read the next complete line from the SSE byte stream.
|
||||
pub(crate) async fn next_line(&mut self) -> Result<Option<String>, Error> {
|
||||
if self.done {
|
||||
return Ok(None);
|
||||
}
|
||||
if let Some(line) = self.line_reader.read_next_chunk("\n").await? {
|
||||
Ok(Some(line))
|
||||
} else {
|
||||
self.done = true;
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
|
||||
/// Process a parsed SSE chunk and return events to emit, if any.
|
||||
///
|
||||
/// `raw` is the same chunk re-parsed as `serde_json::Value` — when the
|
||||
/// chunk carries `usage`, it gets cached so that
|
||||
/// [`ChatHooks::enrich_response`] can read provider-specific fields out
|
||||
/// of it in [`Self::finish_events`].
|
||||
pub(crate) fn process_chunk(
|
||||
&mut self,
|
||||
chunk: &StreamChunk,
|
||||
raw: &serde_json::Value,
|
||||
) -> Option<Vec<StreamEvent>> {
|
||||
// Capture response metadata from the first chunk.
|
||||
if let Some(id) = &chunk.id {
|
||||
if self.response_id.is_empty() {
|
||||
self.response_id.clone_from(id);
|
||||
}
|
||||
}
|
||||
if let Some(model) = &chunk.model {
|
||||
if self.response_model.is_empty() {
|
||||
self.response_model.clone_from(model);
|
||||
}
|
||||
}
|
||||
|
||||
// Capture usage if present (often in a dedicated chunk).
|
||||
if let Some(usage) = &chunk.usage {
|
||||
self.usage = TokenCounts {
|
||||
input_tokens: usage.prompt_tokens,
|
||||
output_tokens: usage.completion_tokens,
|
||||
cache_read_tokens: usage
|
||||
.prompt_tokens_details
|
||||
.as_ref()
|
||||
.and_then(|d| d.cached_tokens)
|
||||
.unwrap_or(0),
|
||||
cache_write_tokens: usage.cache_write_tokens.unwrap_or(0),
|
||||
..TokenCounts::default()
|
||||
};
|
||||
self.last_usage_raw = Some(raw.clone());
|
||||
}
|
||||
|
||||
let choices = chunk.choices.as_ref()?;
|
||||
let choice = choices.first()?;
|
||||
|
||||
let mut events = Vec::new();
|
||||
|
||||
// Check for finish_reason.
|
||||
if let Some(reason) = &choice.finish_reason {
|
||||
self.finish_reason = map_finish_reason(Some(reason.as_str()));
|
||||
}
|
||||
|
||||
let delta = choice.delta.as_ref()?;
|
||||
|
||||
// Accumulate reasoning/thinking content (Kimi, etc.).
|
||||
if let Some(reasoning) = &delta.reasoning_content {
|
||||
if !reasoning.is_empty() {
|
||||
self.accumulated_reasoning.push_str(reasoning);
|
||||
}
|
||||
}
|
||||
|
||||
// Handle text content delta.
|
||||
if let Some(content) = &delta.content {
|
||||
if !content.is_empty() {
|
||||
if !self.text_started {
|
||||
self.text_started = true;
|
||||
events.push(StreamEvent::TextStart { text_id: None });
|
||||
}
|
||||
self.accumulated_text.push_str(content);
|
||||
events.push(StreamEvent::text_delta(content, None));
|
||||
}
|
||||
}
|
||||
|
||||
// Handle tool call deltas.
|
||||
if let Some(tool_calls) = &delta.tool_calls {
|
||||
for tc in tool_calls {
|
||||
let index = tc.index;
|
||||
|
||||
// Grow the accumulated tool calls vector if needed.
|
||||
while self.tool_calls.len() <= index {
|
||||
self.tool_calls.push(AccumulatedToolCall {
|
||||
id: String::new(),
|
||||
name: String::new(),
|
||||
arguments: String::new(),
|
||||
started: false,
|
||||
});
|
||||
}
|
||||
|
||||
let accumulated = &mut self.tool_calls[index];
|
||||
|
||||
// First chunk for this tool call carries id and name.
|
||||
if let Some(id) = &tc.id {
|
||||
accumulated.id.clone_from(id);
|
||||
}
|
||||
if let Some(func) = &tc.function {
|
||||
if let Some(name) = &func.name {
|
||||
accumulated.name.clone_from(name);
|
||||
}
|
||||
if let Some(args) = &func.arguments {
|
||||
accumulated.arguments.push_str(args);
|
||||
}
|
||||
}
|
||||
|
||||
let partial_tool_call =
|
||||
ToolCall::new(&accumulated.id, &accumulated.name, serde_json::json!(null));
|
||||
|
||||
if accumulated.started {
|
||||
events.push(StreamEvent::ToolCallDelta {
|
||||
tool_call: partial_tool_call,
|
||||
});
|
||||
} else {
|
||||
accumulated.started = true;
|
||||
events.push(StreamEvent::ToolCallStart {
|
||||
tool_call: partial_tool_call,
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if events.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(events)
|
||||
}
|
||||
}
|
||||
|
||||
/// Generate the final events when `[DONE]` is received.
|
||||
pub(crate) fn finish_events(&mut self) -> Vec<StreamEvent> {
|
||||
self.finished = true;
|
||||
let mut events = Vec::new();
|
||||
|
||||
// End text segment if it was started.
|
||||
if self.text_started {
|
||||
events.push(StreamEvent::TextEnd { text_id: None });
|
||||
}
|
||||
|
||||
// End all tool calls with complete data.
|
||||
let mut content_parts = Vec::new();
|
||||
|
||||
// Include reasoning/thinking content if present (Kimi, etc.).
|
||||
if !self.accumulated_reasoning.is_empty() {
|
||||
content_parts.push(ContentPart::Thinking(ThinkingData {
|
||||
text: std::mem::take(&mut self.accumulated_reasoning),
|
||||
signature: None,
|
||||
redacted: false,
|
||||
}));
|
||||
}
|
||||
|
||||
if !self.accumulated_text.is_empty() {
|
||||
content_parts.push(ContentPart::text(&self.accumulated_text));
|
||||
}
|
||||
|
||||
for accumulated in &self.tool_calls {
|
||||
let arguments = parse_tool_arguments(
|
||||
&accumulated.name,
|
||||
&accumulated.arguments,
|
||||
&self.custom_tool_names,
|
||||
);
|
||||
let mut tool_call = ToolCall::new(&accumulated.id, &accumulated.name, arguments);
|
||||
tool_call.raw_arguments = Some(accumulated.arguments.clone());
|
||||
|
||||
events.push(StreamEvent::ToolCallEnd {
|
||||
tool_call: tool_call.clone(),
|
||||
});
|
||||
content_parts.push(ContentPart::ToolCall(tool_call));
|
||||
}
|
||||
|
||||
// Infer finish reason from tool calls if not explicitly set.
|
||||
if !self.tool_calls.is_empty() && self.finish_reason == FinishReason::Stop {
|
||||
self.finish_reason = FinishReason::ToolCalls;
|
||||
}
|
||||
|
||||
let response_model = if self.response_model.is_empty() {
|
||||
self.model.clone()
|
||||
} else {
|
||||
self.response_model.clone()
|
||||
};
|
||||
|
||||
let mut response = Response {
|
||||
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(),
|
||||
};
|
||||
|
||||
if let Some(enrich) = self.hooks.enrich_response {
|
||||
let raw = self
|
||||
.last_usage_raw
|
||||
.clone()
|
||||
.unwrap_or(serde_json::Value::Null);
|
||||
enrich(&mut response, &raw);
|
||||
}
|
||||
|
||||
events.push(StreamEvent::finish(
|
||||
self.finish_reason.clone(),
|
||||
self.usage.clone(),
|
||||
response,
|
||||
));
|
||||
|
||||
events
|
||||
}
|
||||
}
|
||||
|
||||
/// Build the streaming `StreamEventStream` from an HTTP response that has
|
||||
/// already been confirmed as success-status. Handles SSE parsing, chunk
|
||||
/// accumulation, and final-event assembly.
|
||||
pub(crate) fn run_stream(
|
||||
http_resp: fabro_http::Response,
|
||||
provider_name: String,
|
||||
model: String,
|
||||
rate_limit: Option<RateLimitInfo>,
|
||||
stream_read_timeout: Option<std::time::Duration>,
|
||||
hooks: ChatHooks,
|
||||
custom_tool_names: Vec<String>,
|
||||
) -> StreamEventStream {
|
||||
let stream = stream::unfold(
|
||||
StreamState::new(
|
||||
http_resp,
|
||||
provider_name,
|
||||
model,
|
||||
rate_limit,
|
||||
stream_read_timeout,
|
||||
hooks,
|
||||
custom_tool_names,
|
||||
),
|
||||
|mut state| async move {
|
||||
loop {
|
||||
let line = match state.next_line().await {
|
||||
Ok(Some(line)) => line,
|
||||
Ok(None) => {
|
||||
// Stream ended without [DONE]. Some providers
|
||||
// (e.g. Minimax) omit the sentinel. Emit
|
||||
// accumulated finish events if we have content
|
||||
// and haven't already emitted them.
|
||||
if !state.finished && (state.text_started || !state.tool_calls.is_empty()) {
|
||||
let events = state.finish_events();
|
||||
return Some((Ok(events), state));
|
||||
}
|
||||
return None;
|
||||
}
|
||||
Err(e) => return Some((Err(e), state)),
|
||||
};
|
||||
|
||||
let line = line.trim();
|
||||
if line.is_empty() || line.starts_with(':') {
|
||||
continue;
|
||||
}
|
||||
|
||||
let data = match line.strip_prefix("data:") {
|
||||
Some(d) => d.trim(),
|
||||
None => continue,
|
||||
};
|
||||
|
||||
if data == "[DONE]" {
|
||||
let events = state.finish_events();
|
||||
return Some((Ok(events), state));
|
||||
}
|
||||
|
||||
let raw: serde_json::Value = match serde_json::from_str(data) {
|
||||
Ok(v) => v,
|
||||
Err(e) => {
|
||||
return Some((
|
||||
Err(Error::stream_error(
|
||||
format!("failed to parse SSE chunk: {e}"),
|
||||
e,
|
||||
)),
|
||||
state,
|
||||
));
|
||||
}
|
||||
};
|
||||
let chunk: StreamChunk = match serde_json::from_value(raw.clone()) {
|
||||
Ok(c) => c,
|
||||
Err(e) => {
|
||||
return Some((
|
||||
Err(Error::stream_error(
|
||||
format!("failed to parse SSE chunk: {e}"),
|
||||
e,
|
||||
)),
|
||||
state,
|
||||
));
|
||||
}
|
||||
};
|
||||
|
||||
if let Some(events) = state.process_chunk(&chunk, &raw) {
|
||||
return Some((Ok(events), state));
|
||||
}
|
||||
}
|
||||
},
|
||||
);
|
||||
|
||||
// Flatten batched events into individual stream events.
|
||||
let flat_stream = stream::unfold(
|
||||
FlattenState {
|
||||
inner: Box::pin(stream),
|
||||
pending: Vec::new(),
|
||||
},
|
||||
|mut flatten_state| async {
|
||||
loop {
|
||||
if let Some(event) = flatten_state.pending.pop() {
|
||||
return Some((Ok(event), flatten_state));
|
||||
}
|
||||
|
||||
match flatten_state.inner.next().await {
|
||||
Some(Ok(mut events)) => {
|
||||
// Reverse so we can pop from the end in order.
|
||||
events.reverse();
|
||||
flatten_state.pending = events;
|
||||
}
|
||||
Some(Err(e)) => return Some((Err(e), flatten_state)),
|
||||
None => return None,
|
||||
}
|
||||
}
|
||||
},
|
||||
);
|
||||
|
||||
Box::pin(flat_stream)
|
||||
}
|
||||
|
||||
/// Drive an already-built `fabro_http::RequestBuilder` to send the chat
|
||||
/// streaming request and produce the unified [`StreamEventStream`].
|
||||
pub(crate) async fn send_and_stream(
|
||||
req: fabro_http::RequestBuilder,
|
||||
provider_name: String,
|
||||
model: String,
|
||||
stream_read_timeout: Option<std::time::Duration>,
|
||||
hooks: ChatHooks,
|
||||
custom_tool_names: Vec<String>,
|
||||
) -> Result<StreamEventStream, Error> {
|
||||
let http_resp = req
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| Error::network(e.to_string(), e))?;
|
||||
|
||||
let status = http_resp.status();
|
||||
if !status.is_success() {
|
||||
let retry_after = parse_retry_after(http_resp.headers());
|
||||
let body = http_resp
|
||||
.text()
|
||||
.await
|
||||
.map_err(|e| Error::network(e.to_string(), e))?;
|
||||
let (msg, code, raw) = parse_error_body(&body, "type");
|
||||
return Err(error_from_status_code(
|
||||
status.as_u16(),
|
||||
msg,
|
||||
provider_name,
|
||||
code,
|
||||
raw,
|
||||
retry_after,
|
||||
));
|
||||
}
|
||||
|
||||
let rate_limit = parse_rate_limit_headers(http_resp.headers());
|
||||
|
||||
Ok(run_stream(
|
||||
http_resp,
|
||||
provider_name,
|
||||
model,
|
||||
rate_limit,
|
||||
stream_read_timeout,
|
||||
hooks,
|
||||
custom_tool_names,
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::providers::openai_chat::wire::{AccumulatedToolCall, StreamChunk};
|
||||
|
||||
#[test]
|
||||
fn stream_chunk_minimax_format() {
|
||||
let json = r#"{"id":"abc","choices":[{"index":0,"delta":{"content":"hello","role":"assistant","name":"MiniMax AI","audio_content":""}}],"created":1772268546,"model":"MiniMax-M2.5","object":"chat.completion.chunk","usage":null,"input_sensitive":false,"output_sensitive":false}"#;
|
||||
let chunk: StreamChunk = serde_json::from_str(json).unwrap();
|
||||
let choices = chunk.choices.unwrap();
|
||||
let delta = choices[0].delta.as_ref().unwrap();
|
||||
assert_eq!(delta.content.as_deref(), Some("hello"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_chunk_text_delta_parsing() {
|
||||
let json = r#"{"id":"chatcmpl-1","model":"gpt-4","choices":[{"delta":{"content":"Hello"},"finish_reason":null}]}"#;
|
||||
let chunk: StreamChunk = serde_json::from_str(json).unwrap();
|
||||
assert_eq!(chunk.id.as_deref(), Some("chatcmpl-1"));
|
||||
assert_eq!(chunk.model.as_deref(), Some("gpt-4"));
|
||||
let choices = chunk.choices.unwrap();
|
||||
assert_eq!(choices.len(), 1);
|
||||
let delta = choices[0].delta.as_ref().unwrap();
|
||||
assert_eq!(delta.content.as_deref(), Some("Hello"));
|
||||
assert!(choices[0].finish_reason.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_chunk_tool_call_parsing() {
|
||||
let json = r#"{"id":"chatcmpl-1","model":"gpt-4","choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_1","function":{"name":"get_weather","arguments":"{\"ci"}}]},"finish_reason":null}]}"#;
|
||||
let chunk: StreamChunk = serde_json::from_str(json).unwrap();
|
||||
let choices = chunk.choices.unwrap();
|
||||
let delta = choices[0].delta.as_ref().unwrap();
|
||||
let tc = &delta.tool_calls.as_ref().unwrap()[0];
|
||||
assert_eq!(tc.index, 0);
|
||||
assert_eq!(tc.id.as_deref(), Some("call_1"));
|
||||
let func = tc.function.as_ref().unwrap();
|
||||
assert_eq!(func.name.as_deref(), Some("get_weather"));
|
||||
assert_eq!(func.arguments.as_deref(), Some("{\"ci"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_chunk_usage_parsing() {
|
||||
let json = r#"{"id":"chatcmpl-1","model":"gpt-4","choices":[],"usage":{"prompt_tokens":10,"completion_tokens":20,"total_tokens":30}}"#;
|
||||
let chunk: StreamChunk = serde_json::from_str(json).unwrap();
|
||||
let usage = chunk.usage.unwrap();
|
||||
assert_eq!(usage.prompt_tokens, 10);
|
||||
assert_eq!(usage.completion_tokens, 20);
|
||||
assert!(usage.prompt_tokens_details.is_none());
|
||||
assert!(usage.cache_write_tokens.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_chunk_usage_parses_cached_tokens() {
|
||||
let json = r#"{"id":"chatcmpl-1","model":"gpt-4","choices":[],"usage":{"prompt_tokens":100,"completion_tokens":50,"prompt_tokens_details":{"cached_tokens":80},"cache_write_tokens":12}}"#;
|
||||
let chunk: StreamChunk = serde_json::from_str(json).unwrap();
|
||||
let usage = chunk.usage.unwrap();
|
||||
assert_eq!(usage.prompt_tokens, 100);
|
||||
assert_eq!(usage.completion_tokens, 50);
|
||||
assert_eq!(
|
||||
usage
|
||||
.prompt_tokens_details
|
||||
.as_ref()
|
||||
.and_then(|d| d.cached_tokens),
|
||||
Some(80)
|
||||
);
|
||||
assert_eq!(usage.cache_write_tokens, Some(12));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_state_process_usage_populates_cache_counts() {
|
||||
let http_resp =
|
||||
fabro_http::Response::from(http::Response::builder().status(200).body("").unwrap());
|
||||
let mut state = StreamState::new(
|
||||
http_resp,
|
||||
"openrouter".into(),
|
||||
"model".into(),
|
||||
None,
|
||||
Some(std::time::Duration::from_secs(30)),
|
||||
ChatHooks::NONE,
|
||||
Vec::new(),
|
||||
);
|
||||
|
||||
let raw: serde_json::Value = serde_json::from_str(
|
||||
r#"{"id":"c1","model":"m1","choices":[],"usage":{"prompt_tokens":100,"completion_tokens":50,"prompt_tokens_details":{"cached_tokens":80},"cache_write_tokens":12}}"#,
|
||||
)
|
||||
.unwrap();
|
||||
let chunk: StreamChunk = serde_json::from_value(raw.clone()).unwrap();
|
||||
let _ = state.process_chunk(&chunk, &raw);
|
||||
assert_eq!(state.usage.input_tokens, 100);
|
||||
assert_eq!(state.usage.output_tokens, 50);
|
||||
assert_eq!(state.usage.cache_read_tokens, 80);
|
||||
assert_eq!(state.usage.cache_write_tokens, 12);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_chunk_finish_reason_parsing() {
|
||||
let json = r#"{"id":"chatcmpl-1","model":"gpt-4","choices":[{"delta":{},"finish_reason":"stop"}]}"#;
|
||||
let chunk: StreamChunk = serde_json::from_str(json).unwrap();
|
||||
let choices = chunk.choices.unwrap();
|
||||
assert_eq!(choices[0].finish_reason.as_deref(), Some("stop"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_state_process_text_chunks() {
|
||||
let http_resp =
|
||||
fabro_http::Response::from(http::Response::builder().status(200).body("").unwrap());
|
||||
let mut state = StreamState::new(
|
||||
http_resp,
|
||||
"test".into(),
|
||||
"model".into(),
|
||||
None,
|
||||
Some(std::time::Duration::from_secs(30)),
|
||||
ChatHooks::NONE,
|
||||
Vec::new(),
|
||||
);
|
||||
|
||||
// First text chunk should emit TextStart + TextDelta.
|
||||
let raw1: serde_json::Value = serde_json::from_str(
|
||||
r#"{"id":"c1","model":"m1","choices":[{"delta":{"content":"Hello"},"finish_reason":null}]}"#,
|
||||
).unwrap();
|
||||
let chunk1: StreamChunk = serde_json::from_value(raw1.clone()).unwrap();
|
||||
let events1 = state.process_chunk(&chunk1, &raw1).unwrap();
|
||||
assert_eq!(events1.len(), 2);
|
||||
assert!(matches!(events1[0], StreamEvent::TextStart { .. }));
|
||||
assert!(matches!(events1[1], StreamEvent::TextDelta { .. }));
|
||||
|
||||
// Second text chunk should emit only TextDelta (no second TextStart).
|
||||
let raw2: serde_json::Value = serde_json::from_str(
|
||||
r#"{"id":"c1","model":"m1","choices":[{"delta":{"content":" world"},"finish_reason":null}]}"#,
|
||||
).unwrap();
|
||||
let chunk2: StreamChunk = serde_json::from_value(raw2.clone()).unwrap();
|
||||
let events2 = state.process_chunk(&chunk2, &raw2).unwrap();
|
||||
assert_eq!(events2.len(), 1);
|
||||
assert!(matches!(events2[0], StreamEvent::TextDelta { .. }));
|
||||
|
||||
assert_eq!(state.accumulated_text, "Hello world");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_state_process_tool_call_chunks() {
|
||||
let http_resp =
|
||||
fabro_http::Response::from(http::Response::builder().status(200).body("").unwrap());
|
||||
let mut state = StreamState::new(
|
||||
http_resp,
|
||||
"test".into(),
|
||||
"model".into(),
|
||||
None,
|
||||
Some(std::time::Duration::from_secs(30)),
|
||||
ChatHooks::NONE,
|
||||
Vec::new(),
|
||||
);
|
||||
|
||||
// First tool call chunk (has id and name) -> ToolCallStart.
|
||||
let raw1: serde_json::Value = serde_json::from_str(
|
||||
r#"{"id":"c1","model":"m1","choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_1","function":{"name":"fn1","arguments":"{\"k"}}]},"finish_reason":null}]}"#,
|
||||
).unwrap();
|
||||
let chunk1: StreamChunk = serde_json::from_value(raw1.clone()).unwrap();
|
||||
let events1 = state.process_chunk(&chunk1, &raw1).unwrap();
|
||||
assert_eq!(events1.len(), 1);
|
||||
assert!(matches!(events1[0], StreamEvent::ToolCallStart { .. }));
|
||||
|
||||
// Subsequent chunk (more arguments) -> ToolCallDelta.
|
||||
let raw2: serde_json::Value = serde_json::from_str(
|
||||
r#"{"id":"c1","model":"m1","choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"ey\"}"}}]},"finish_reason":null}]}"#,
|
||||
).unwrap();
|
||||
let chunk2: StreamChunk = serde_json::from_value(raw2.clone()).unwrap();
|
||||
let events2 = state.process_chunk(&chunk2, &raw2).unwrap();
|
||||
assert_eq!(events2.len(), 1);
|
||||
assert!(matches!(events2[0], StreamEvent::ToolCallDelta { .. }));
|
||||
|
||||
assert_eq!(state.tool_calls[0].arguments, r#"{"key"}"#);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_state_finish_events_text_only() {
|
||||
let http_resp =
|
||||
fabro_http::Response::from(http::Response::builder().status(200).body("").unwrap());
|
||||
let mut state = StreamState::new(
|
||||
http_resp,
|
||||
"test-provider".into(),
|
||||
"test-model".into(),
|
||||
None,
|
||||
Some(std::time::Duration::from_secs(30)),
|
||||
ChatHooks::NONE,
|
||||
Vec::new(),
|
||||
);
|
||||
state.response_id = "resp-1".into();
|
||||
state.response_model = "gpt-4".into();
|
||||
state.accumulated_text = "Hello world".into();
|
||||
state.text_started = true;
|
||||
state.usage = TokenCounts {
|
||||
input_tokens: 5,
|
||||
output_tokens: 10,
|
||||
..TokenCounts::default()
|
||||
};
|
||||
|
||||
let events = state.finish_events();
|
||||
// TextEnd + Finish
|
||||
assert_eq!(events.len(), 2);
|
||||
assert!(matches!(events[0], StreamEvent::TextEnd { .. }));
|
||||
match &events[1] {
|
||||
StreamEvent::Finish {
|
||||
finish_reason,
|
||||
usage,
|
||||
response,
|
||||
} => {
|
||||
assert_eq!(*finish_reason, FinishReason::Stop);
|
||||
assert_eq!(usage.input_tokens, 5);
|
||||
assert_eq!(usage.output_tokens, 10);
|
||||
assert_eq!(response.text(), "Hello world");
|
||||
assert_eq!(response.id, "resp-1");
|
||||
assert_eq!(response.model, "gpt-4");
|
||||
assert_eq!(response.provider, "test-provider");
|
||||
}
|
||||
other => panic!("Expected Finish, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_state_finish_events_with_tool_calls() {
|
||||
let http_resp =
|
||||
fabro_http::Response::from(http::Response::builder().status(200).body("").unwrap());
|
||||
let mut state = StreamState::new(
|
||||
http_resp,
|
||||
"test".into(),
|
||||
"model".into(),
|
||||
None,
|
||||
Some(std::time::Duration::from_secs(30)),
|
||||
ChatHooks::NONE,
|
||||
Vec::new(),
|
||||
);
|
||||
state.response_id = "resp-1".into();
|
||||
state.tool_calls.push(AccumulatedToolCall {
|
||||
id: "call_1".into(),
|
||||
name: "get_weather".into(),
|
||||
arguments: r#"{"city":"SF"}"#.into(),
|
||||
started: true,
|
||||
});
|
||||
|
||||
let events = state.finish_events();
|
||||
// ToolCallEnd + Finish (no TextEnd since text_started is false)
|
||||
assert_eq!(events.len(), 2);
|
||||
match &events[0] {
|
||||
StreamEvent::ToolCallEnd { tool_call } => {
|
||||
assert_eq!(tool_call.id, "call_1");
|
||||
assert_eq!(tool_call.name, "get_weather");
|
||||
assert_eq!(tool_call.raw_arguments.as_deref(), Some(r#"{"city":"SF"}"#));
|
||||
}
|
||||
other => panic!("Expected ToolCallEnd, got {other:?}"),
|
||||
}
|
||||
match &events[1] {
|
||||
StreamEvent::Finish {
|
||||
finish_reason,
|
||||
response,
|
||||
..
|
||||
} => {
|
||||
assert_eq!(*finish_reason, FinishReason::ToolCalls);
|
||||
let calls = response.tool_calls();
|
||||
assert_eq!(calls.len(), 1);
|
||||
assert_eq!(calls[0].name, "get_weather");
|
||||
}
|
||||
other => panic!("Expected Finish, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_state_uses_request_model_as_fallback() {
|
||||
let http_resp =
|
||||
fabro_http::Response::from(http::Response::builder().status(200).body("").unwrap());
|
||||
let mut state = StreamState::new(
|
||||
http_resp,
|
||||
"test".into(),
|
||||
"fallback-model".into(),
|
||||
None,
|
||||
Some(std::time::Duration::from_secs(30)),
|
||||
ChatHooks::NONE,
|
||||
Vec::new(),
|
||||
);
|
||||
// response_model is empty, so finish_events should use the request model.
|
||||
let events = state.finish_events();
|
||||
match &events[0] {
|
||||
StreamEvent::Finish { response, .. } => {
|
||||
assert_eq!(response.model, "fallback-model");
|
||||
}
|
||||
other => panic!("Expected Finish, got {other:?}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
417
lib/crates/fabro-llm/src/providers/openai_chat/translate.rs
Normal file
417
lib/crates/fabro-llm/src/providers/openai_chat/translate.rs
Normal file
|
|
@ -0,0 +1,417 @@
|
|||
//! Translations between Fabro's unified [`Request`](crate::types::Request)
|
||||
//! shape and the OpenAI Chat Completions wire format.
|
||||
|
||||
use super::wire::{ChatFunction, ChatMessage, ChatToolCall};
|
||||
use crate::types::{
|
||||
ContentPart, FinishReason, Message, Request, ResponseFormat, ResponseFormatType, Role,
|
||||
ToolChoice, ToolDefinition,
|
||||
};
|
||||
|
||||
pub(crate) fn map_finish_reason(reason: Option<&str>) -> FinishReason {
|
||||
match reason {
|
||||
Some("stop") | None => FinishReason::Stop,
|
||||
Some("length") => FinishReason::Length,
|
||||
Some("tool_calls") => FinishReason::ToolCalls,
|
||||
Some("content_filter") => FinishReason::ContentFilter,
|
||||
Some(other) => FinishReason::Other(other.to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Build the content string from a message's parts, including fallback text
|
||||
/// for unsupported content types (Audio, Document).
|
||||
pub(crate) fn content_text_with_fallbacks(parts: &[ContentPart]) -> String {
|
||||
let mut segments: Vec<String> = Vec::new();
|
||||
for part in parts {
|
||||
match part {
|
||||
ContentPart::Text(text) => segments.push(text.clone()),
|
||||
ContentPart::Audio(_) => {
|
||||
segments.push("[Audio content not supported by this provider]".to_string());
|
||||
}
|
||||
ContentPart::Document(doc) => {
|
||||
let desc = doc.file_name.as_ref().map_or_else(
|
||||
|| "[Document content not supported by this provider]".to_string(),
|
||||
|name| {
|
||||
format!("[Document '{name}': content type not supported by this provider]")
|
||||
},
|
||||
);
|
||||
segments.push(desc);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
segments.join("")
|
||||
}
|
||||
|
||||
pub(crate) fn translate_messages(messages: &[Message]) -> Vec<ChatMessage> {
|
||||
messages
|
||||
.iter()
|
||||
.flat_map(|msg| {
|
||||
// Tool messages must be split into one ChatMessage per ToolResult,
|
||||
// each with its own tool_call_id. The Chat Completions API requires
|
||||
// every tool_call_id from the assistant to have a matching tool message.
|
||||
if msg.role == Role::Tool {
|
||||
return msg
|
||||
.content
|
||||
.iter()
|
||||
.filter_map(|part| {
|
||||
if let ContentPart::ToolResult(tr) = part {
|
||||
let output = tr
|
||||
.content
|
||||
.as_str()
|
||||
.map_or_else(|| tr.content.to_string(), str::to_string);
|
||||
Some(ChatMessage {
|
||||
role: "tool".to_string(),
|
||||
content: Some(output),
|
||||
reasoning_content: None,
|
||||
tool_call_id: Some(tr.tool_call_id.clone()),
|
||||
tool_calls: None,
|
||||
})
|
||||
} else {
|
||||
None
|
||||
}
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
}
|
||||
|
||||
let role = match msg.role {
|
||||
Role::System | Role::Developer => "system",
|
||||
Role::User => "user",
|
||||
Role::Assistant => "assistant",
|
||||
Role::Tool => unreachable!(
|
||||
"Role::Tool is handled in the early-return branch above this match"
|
||||
),
|
||||
};
|
||||
|
||||
let mut tool_calls: Vec<ChatToolCall> = Vec::new();
|
||||
if msg.role == Role::Assistant {
|
||||
for part in &msg.content {
|
||||
if let ContentPart::ToolCall(tc) = part {
|
||||
let arguments = tc
|
||||
.raw_arguments
|
||||
.clone()
|
||||
.unwrap_or_else(|| tc.arguments.to_string());
|
||||
tool_calls.push(ChatToolCall {
|
||||
id: tc.id.clone(),
|
||||
kind: "function".to_string(),
|
||||
function: ChatFunction {
|
||||
name: tc.name.clone(),
|
||||
arguments,
|
||||
},
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let text = content_text_with_fallbacks(&msg.content);
|
||||
let content = if text.is_empty() { None } else { Some(text) };
|
||||
let tool_calls = if tool_calls.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(tool_calls)
|
||||
};
|
||||
|
||||
// Extract reasoning/thinking content for assistant messages.
|
||||
let reasoning_content = if msg.role == Role::Assistant {
|
||||
let reasoning: String = msg
|
||||
.content
|
||||
.iter()
|
||||
.filter_map(|part| match part {
|
||||
ContentPart::Thinking(t) if !t.redacted => Some(t.text.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join("");
|
||||
if reasoning.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(reasoning)
|
||||
}
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
vec![ChatMessage {
|
||||
role: role.to_string(),
|
||||
content,
|
||||
reasoning_content,
|
||||
tool_call_id: msg.tool_call_id.clone(),
|
||||
tool_calls,
|
||||
}]
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(crate) fn translate_tools(tools: &[ToolDefinition]) -> Vec<serde_json::Value> {
|
||||
tools
|
||||
.iter()
|
||||
.map(|t| {
|
||||
serde_json::json!({
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": t.name,
|
||||
"description": t.description,
|
||||
"parameters": t.parameters,
|
||||
}
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(crate) fn translate_tool_choice(choice: &ToolChoice) -> serde_json::Value {
|
||||
match choice {
|
||||
ToolChoice::Auto => serde_json::json!("auto"),
|
||||
ToolChoice::None => serde_json::json!("none"),
|
||||
ToolChoice::Required => serde_json::json!("required"),
|
||||
ToolChoice::Named { tool_name } => {
|
||||
serde_json::json!({"type": "function", "function": {"name": tool_name}})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Names of tools on the request marked as "custom" (freeform), used to
|
||||
/// decide how to handle tool-call arguments that don't parse as JSON.
|
||||
pub(crate) fn custom_tool_names(request: &Request) -> Vec<String> {
|
||||
request
|
||||
.tools
|
||||
.as_deref()
|
||||
.unwrap_or_default()
|
||||
.iter()
|
||||
.filter(|tool| tool.is_custom())
|
||||
.map(|tool| tool.name.clone())
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Parse tool-call `arguments` to JSON. For custom (freeform) tools whose
|
||||
/// arguments aren't valid JSON, preserve the raw string so callers like
|
||||
/// `apply_patch` receive the actual patch text instead of an empty object.
|
||||
pub(crate) fn parse_tool_arguments(
|
||||
tool_name: &str,
|
||||
raw_arguments: &str,
|
||||
custom_tool_names: &[String],
|
||||
) -> serde_json::Value {
|
||||
match serde_json::from_str(raw_arguments) {
|
||||
Ok(arguments) => arguments,
|
||||
Err(_) if custom_tool_names.iter().any(|name| name == tool_name) => {
|
||||
serde_json::Value::String(raw_arguments.to_string())
|
||||
}
|
||||
Err(_) => serde_json::json!({}),
|
||||
}
|
||||
}
|
||||
|
||||
/// Translate unified `ResponseFormat` to Chat Completions `response_format`.
|
||||
pub(crate) fn translate_response_format(format: &ResponseFormat) -> serde_json::Value {
|
||||
match format.kind {
|
||||
ResponseFormatType::Text => serde_json::json!({"type": "text"}),
|
||||
ResponseFormatType::JsonObject => serde_json::json!({"type": "json_object"}),
|
||||
ResponseFormatType::JsonSchema => {
|
||||
let mut json_schema = serde_json::json!({
|
||||
"name": "response",
|
||||
"strict": format.strict,
|
||||
});
|
||||
if let Some(schema) = &format.json_schema {
|
||||
json_schema["schema"] = schema.clone();
|
||||
}
|
||||
serde_json::json!({
|
||||
"type": "json_schema",
|
||||
"json_schema": json_schema,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::types::{AudioData, ContentPart, DocumentData, Message, Role, ToolCall};
|
||||
|
||||
#[test]
|
||||
fn translate_assistant_message_with_tool_calls_only() {
|
||||
let msg = Message {
|
||||
role: Role::Assistant,
|
||||
content: vec![ContentPart::ToolCall(ToolCall::new(
|
||||
"call_1",
|
||||
"get_weather",
|
||||
serde_json::json!({"city": "SF"}),
|
||||
))],
|
||||
name: None,
|
||||
tool_call_id: None,
|
||||
};
|
||||
let translated = translate_messages(&[msg]);
|
||||
assert_eq!(translated.len(), 1);
|
||||
assert_eq!(translated[0].role, "assistant");
|
||||
assert!(translated[0].content.is_none());
|
||||
let tool_calls = translated[0].tool_calls.as_ref().unwrap();
|
||||
assert_eq!(tool_calls.len(), 1);
|
||||
assert_eq!(tool_calls[0].id, "call_1");
|
||||
assert_eq!(tool_calls[0].kind, "function");
|
||||
assert_eq!(tool_calls[0].function.name, "get_weather");
|
||||
assert_eq!(tool_calls[0].function.arguments, r#"{"city":"SF"}"#);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn translate_assistant_message_with_text_and_tool_calls() {
|
||||
let msg = Message {
|
||||
role: Role::Assistant,
|
||||
content: vec![
|
||||
ContentPart::text("Let me check the weather"),
|
||||
ContentPart::ToolCall(ToolCall::new(
|
||||
"call_2",
|
||||
"get_weather",
|
||||
serde_json::json!({"city": "NYC"}),
|
||||
)),
|
||||
],
|
||||
name: None,
|
||||
tool_call_id: None,
|
||||
};
|
||||
let translated = translate_messages(&[msg]);
|
||||
assert_eq!(
|
||||
translated[0].content.as_deref(),
|
||||
Some("Let me check the weather")
|
||||
);
|
||||
let tool_calls = translated[0].tool_calls.as_ref().unwrap();
|
||||
assert_eq!(tool_calls.len(), 1);
|
||||
assert_eq!(tool_calls[0].function.name, "get_weather");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn translate_assistant_message_with_raw_arguments() {
|
||||
let mut tc = ToolCall::new("call_3", "search", serde_json::json!({"q": "rust"}));
|
||||
tc.raw_arguments = Some(r#"{"q": "rust"}"#.to_string());
|
||||
let msg = Message {
|
||||
role: Role::Assistant,
|
||||
content: vec![ContentPart::ToolCall(tc)],
|
||||
name: None,
|
||||
tool_call_id: None,
|
||||
};
|
||||
let translated = translate_messages(&[msg]);
|
||||
let tool_calls = translated[0].tool_calls.as_ref().unwrap();
|
||||
// Should prefer raw_arguments over serializing arguments
|
||||
assert_eq!(tool_calls[0].function.arguments, r#"{"q": "rust"}"#);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn translate_tool_message_has_tool_call_id() {
|
||||
let msg = Message::tool_result(
|
||||
"call_1",
|
||||
serde_json::Value::String("72F and sunny".into()),
|
||||
false,
|
||||
);
|
||||
let translated = translate_messages(&[msg]);
|
||||
assert_eq!(translated[0].role, "tool");
|
||||
assert_eq!(translated[0].tool_call_id.as_deref(), Some("call_1"));
|
||||
assert!(translated[0].tool_calls.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn translate_user_message_has_no_tool_calls() {
|
||||
let msg = Message::user("Hello");
|
||||
let translated = translate_messages(&[msg]);
|
||||
assert_eq!(translated[0].role, "user");
|
||||
assert_eq!(translated[0].content.as_deref(), Some("Hello"));
|
||||
assert!(translated[0].tool_calls.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn assistant_tool_calls_serialize_correctly() {
|
||||
let msg = Message {
|
||||
role: Role::Assistant,
|
||||
content: vec![ContentPart::ToolCall(ToolCall::new(
|
||||
"call_1",
|
||||
"get_weather",
|
||||
serde_json::json!({"city": "SF"}),
|
||||
))],
|
||||
name: None,
|
||||
tool_call_id: None,
|
||||
};
|
||||
let translated = translate_messages(&[msg]);
|
||||
let json = serde_json::to_value(&translated[0]).unwrap();
|
||||
assert!(json.get("content").is_none());
|
||||
assert!(json.get("tool_call_id").is_none());
|
||||
let tool_calls = json["tool_calls"].as_array().unwrap();
|
||||
assert_eq!(tool_calls.len(), 1);
|
||||
assert_eq!(tool_calls[0]["type"], "function");
|
||||
assert_eq!(tool_calls[0]["id"], "call_1");
|
||||
assert_eq!(tool_calls[0]["function"]["name"], "get_weather");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn audio_content_produces_text_fallback() {
|
||||
let msg = Message {
|
||||
role: Role::User,
|
||||
content: vec![ContentPart::Audio(AudioData {
|
||||
url: Some("https://example.com/audio.wav".to_string()),
|
||||
data: None,
|
||||
media_type: None,
|
||||
})],
|
||||
name: None,
|
||||
tool_call_id: None,
|
||||
};
|
||||
let translated = translate_messages(&[msg]);
|
||||
assert_eq!(
|
||||
translated[0].content.as_deref(),
|
||||
Some("[Audio content not supported by this provider]")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_content_produces_text_fallback_with_filename() {
|
||||
let msg = Message {
|
||||
role: Role::User,
|
||||
content: vec![ContentPart::Document(DocumentData {
|
||||
url: Some("https://example.com/doc.pdf".to_string()),
|
||||
data: None,
|
||||
media_type: None,
|
||||
file_name: Some("report.pdf".to_string()),
|
||||
})],
|
||||
name: None,
|
||||
tool_call_id: None,
|
||||
};
|
||||
let translated = translate_messages(&[msg]);
|
||||
assert_eq!(
|
||||
translated[0].content.as_deref(),
|
||||
Some("[Document 'report.pdf': content type not supported by this provider]")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_content_produces_text_fallback_without_filename() {
|
||||
let msg = Message {
|
||||
role: Role::User,
|
||||
content: vec![ContentPart::Document(DocumentData {
|
||||
url: None,
|
||||
data: Some(vec![1, 2, 3]),
|
||||
media_type: None,
|
||||
file_name: None,
|
||||
})],
|
||||
name: None,
|
||||
tool_call_id: None,
|
||||
};
|
||||
let translated = translate_messages(&[msg]);
|
||||
assert_eq!(
|
||||
translated[0].content.as_deref(),
|
||||
Some("[Document content not supported by this provider]")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mixed_text_and_audio_content_concatenates() {
|
||||
let msg = Message {
|
||||
role: Role::User,
|
||||
content: vec![
|
||||
ContentPart::text("Check this: "),
|
||||
ContentPart::Audio(AudioData {
|
||||
url: None,
|
||||
data: Some(vec![1, 2]),
|
||||
media_type: None,
|
||||
}),
|
||||
],
|
||||
name: None,
|
||||
tool_call_id: None,
|
||||
};
|
||||
let translated = translate_messages(&[msg]);
|
||||
assert_eq!(
|
||||
translated[0].content.as_deref(),
|
||||
Some("Check this: [Audio content not supported by this provider]")
|
||||
);
|
||||
}
|
||||
}
|
||||
184
lib/crates/fabro-llm/src/providers/openai_chat/wire.rs
Normal file
184
lib/crates/fabro-llm/src/providers/openai_chat/wire.rs
Normal file
|
|
@ -0,0 +1,184 @@
|
|||
//! Serde wire types for the OpenAI Chat Completions request/response
|
||||
//! protocol, plus the streaming-side `AccumulatedToolCall` helper.
|
||||
//!
|
||||
//! These structs are intentionally tiny and faithful to the wire — they
|
||||
//! are shared between [`crate::providers::openai_compatible`] and (in
|
||||
//! follow-up work) a dedicated OpenRouter adapter.
|
||||
|
||||
// --- Request types (Chat Completions format) ---
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
pub(crate) struct ApiRequest {
|
||||
pub(crate) model: String,
|
||||
pub(crate) messages: Vec<ChatMessage>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub(crate) temperature: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub(crate) max_tokens: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub(crate) top_p: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub(crate) stop: Option<Vec<String>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub(crate) tools: Option<Vec<serde_json::Value>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub(crate) tool_choice: Option<serde_json::Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub(crate) response_format: Option<serde_json::Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub(crate) stream: Option<bool>,
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
pub(crate) struct ChatMessage {
|
||||
pub(crate) role: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub(crate) content: Option<String>,
|
||||
/// Reasoning/thinking content echoed back for providers that require it
|
||||
/// (Kimi).
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub(crate) reasoning_content: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub(crate) tool_call_id: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub(crate) tool_calls: Option<Vec<ChatToolCall>>,
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
pub(crate) struct ChatToolCall {
|
||||
pub(crate) id: String,
|
||||
#[serde(rename = "type")]
|
||||
pub(crate) kind: String,
|
||||
pub(crate) function: ChatFunction,
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
pub(crate) struct ChatFunction {
|
||||
pub(crate) name: String,
|
||||
pub(crate) arguments: String,
|
||||
}
|
||||
|
||||
// --- Response types (non-streaming) ---
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
pub(crate) struct ApiResponse {
|
||||
pub(crate) id: String,
|
||||
pub(crate) model: String,
|
||||
pub(crate) choices: Vec<ApiChoice>,
|
||||
pub(crate) usage: Option<ApiUsage>,
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
pub(crate) struct ApiChoice {
|
||||
pub(crate) message: ApiChoiceMessage,
|
||||
pub(crate) finish_reason: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
pub(crate) struct ApiChoiceMessage {
|
||||
pub(crate) content: Option<String>,
|
||||
pub(crate) reasoning_content: Option<String>,
|
||||
pub(crate) tool_calls: Option<Vec<ApiToolCall>>,
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
pub(crate) struct ApiToolCall {
|
||||
pub(crate) id: String,
|
||||
pub(crate) function: ApiFunction,
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
pub(crate) struct ApiFunction {
|
||||
pub(crate) name: String,
|
||||
pub(crate) arguments: String,
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
#[allow(
|
||||
clippy::struct_field_names,
|
||||
reason = "Field names mirror the provider API payload."
|
||||
)]
|
||||
pub(crate) struct ApiUsage {
|
||||
pub(crate) prompt_tokens: i64,
|
||||
pub(crate) completion_tokens: i64,
|
||||
#[serde(default)]
|
||||
pub(crate) prompt_tokens_details: Option<ApiPromptTokensDetails>,
|
||||
#[serde(default)]
|
||||
pub(crate) cache_write_tokens: Option<i64>,
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize, Default)]
|
||||
pub(crate) struct ApiPromptTokensDetails {
|
||||
#[serde(default)]
|
||||
pub(crate) cached_tokens: Option<i64>,
|
||||
}
|
||||
|
||||
// --- Streaming response types ---
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
pub(crate) struct StreamChunk {
|
||||
pub(crate) id: Option<String>,
|
||||
pub(crate) model: Option<String>,
|
||||
pub(crate) choices: Option<Vec<StreamChoice>>,
|
||||
pub(crate) usage: Option<ApiUsage>,
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
pub(crate) struct StreamChoice {
|
||||
pub(crate) delta: Option<StreamDelta>,
|
||||
pub(crate) finish_reason: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
pub(crate) struct StreamDelta {
|
||||
pub(crate) content: Option<String>,
|
||||
/// Reasoning/thinking content (used by Kimi and other reasoning models).
|
||||
pub(crate) reasoning_content: Option<String>,
|
||||
pub(crate) tool_calls: Option<Vec<StreamToolCall>>,
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
pub(crate) struct StreamToolCall {
|
||||
pub(crate) index: usize,
|
||||
pub(crate) id: Option<String>,
|
||||
pub(crate) function: Option<StreamFunction>,
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
pub(crate) struct StreamFunction {
|
||||
pub(crate) name: Option<String>,
|
||||
pub(crate) arguments: Option<String>,
|
||||
}
|
||||
|
||||
// --- Accumulated tool call state for streaming ---
|
||||
|
||||
pub(crate) struct AccumulatedToolCall {
|
||||
pub(crate) id: String,
|
||||
pub(crate) name: String,
|
||||
pub(crate) arguments: String,
|
||||
pub(crate) started: bool,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn api_response_usage_parses_cached_tokens() {
|
||||
// Non-streaming response uses ApiResponse, which is a distinct
|
||||
// code path from the streaming StreamChunk parse. Verify the
|
||||
// same cached-token fields land on TokenCounts via that path.
|
||||
let json = r#"{"id":"chatcmpl-1","model":"anthropic/claude-sonnet-4.6","choices":[{"message":{"role":"assistant","content":"hi"},"finish_reason":"stop"}],"usage":{"prompt_tokens":100,"completion_tokens":50,"prompt_tokens_details":{"cached_tokens":80},"cache_write_tokens":12}}"#;
|
||||
let resp: ApiResponse = serde_json::from_str(json).unwrap();
|
||||
let u = resp.usage.unwrap();
|
||||
assert_eq!(u.prompt_tokens, 100);
|
||||
assert_eq!(u.completion_tokens, 50);
|
||||
assert_eq!(
|
||||
u.prompt_tokens_details
|
||||
.as_ref()
|
||||
.and_then(|d| d.cached_tokens),
|
||||
Some(80)
|
||||
);
|
||||
assert_eq!(u.cache_write_tokens, Some(12));
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load diff
Loading…
Add table
Reference in a new issue