Skip to main content

trusty_common/chat/
mod.rs

1//! Provider-agnostic streaming chat abstraction with tool-use support.
2//!
3//! Why: trusty-memory and trusty-search both want to support more than one
4//! upstream LLM (OpenRouter for cloud, Ollama / LM Studio for local). Rather
5//! than each crate re-implementing the dispatch, we expose a small
6//! [`ChatProvider`] trait plus concrete implementations and an auto-detector
7//! for a running local model server. The trait also surfaces OpenAI-style
8//! tool/function calling so downstream agents can let the model invoke tools
9//! (search, memory recall, shell, etc.).
10//!
11//! What: defines the [`ChatProvider`] trait, [`ToolDef`] / [`ToolCall`] /
12//! [`ChatEvent`] tool-use types, an [`OpenRouterProvider`] and an
13//! [`OllamaProvider`] that both speak OpenAI-compatible
14//! `/v1/chat/completions` with SSE streaming (including the streamed
15//! `tool_calls` shape), a [`BedrockProvider`] that uses the AWS Bedrock
16//! `Converse` API (behind the `bedrock` feature flag), and
17//! [`auto_detect_local_provider`] which probes `{base_url}/v1/models` with a
18//! 1-second timeout.
19//!
20//! Test: `cargo test -p trusty-common --features unconditional-only` covers
21//! default config values, the unreachable-server path of
22//! `auto_detect_local_provider`, SSE delta streaming, and accumulation of
23//! streamed tool-call fragments.
24
25mod openai_compat;
26
27#[cfg(feature = "bedrock")]
28mod bedrock_impl;
29#[cfg(not(feature = "bedrock"))]
30mod bedrock_stub;
31
32pub use openai_compat::{OllamaProvider, OpenRouterProvider, auto_detect_local_provider};
33
34#[cfg(feature = "bedrock")]
35pub use bedrock_impl::{
36    BedrockProvider, DEFAULT_BEDROCK_MODEL, DEFAULT_BEDROCK_REGION, ENV_REGION_AWS,
37    ENV_REGION_TRUSTY,
38};
39
40// Re-expose the bedrock_impl module as `bedrock_provider` so downstream
41// crates can access constants (e.g. `DEFAULT_BEDROCK_MODEL`) without needing
42// to depend on the bedrock feature themselves.
43#[cfg(feature = "bedrock")]
44pub mod bedrock_provider {
45    pub use super::bedrock_impl::*;
46}
47
48#[cfg(not(feature = "bedrock"))]
49pub use bedrock_stub::BedrockProvider;
50
51// Stub constant so code that references DEFAULT_BEDROCK_MODEL compiles without
52// the bedrock feature. Must stay in sync with bedrock_impl::DEFAULT_BEDROCK_MODEL.
53// Claude Sonnet 4.6 drops the date stamp and -v1:0 suffix (verified vs AWS docs).
54#[cfg(not(feature = "bedrock"))]
55pub const DEFAULT_BEDROCK_MODEL: &str = "us.anthropic.claude-sonnet-4-6";
56
57use crate::ChatMessage;
58use anyhow::Result;
59use async_trait::async_trait;
60use serde::{Deserialize, Serialize};
61use tokio::sync::mpsc::Sender;
62
63// ── Public re-exports so callers get the full surface from `chat::*` ──────────
64
65/// Configuration for a local OpenAI-compatible model server (Ollama, LM
66/// Studio, llama.cpp's server, etc.).
67///
68/// Why: callers want a single struct they can deserialize from config files
69/// and pass to [`auto_detect_local_provider`] without juggling defaults.
70/// What: holds an enable flag, the server's base URL (no trailing slash),
71/// and the default model to request. Defaults target Ollama's standard
72/// localhost binding.
73/// Test: `local_model_config_defaults` asserts the default values.
74#[derive(Debug, Clone, Serialize, Deserialize)]
75pub struct LocalModelConfig {
76    pub enabled: bool,
77    pub base_url: String,
78    pub model: String,
79}
80
81impl Default for LocalModelConfig {
82    fn default() -> Self {
83        Self {
84            enabled: true,
85            base_url: "http://localhost:11434".to_string(),
86            model: "qwen3:30b".to_string(),
87        }
88    }
89}
90
91// ─── Tool-use types ───────────────────────────────────────────────────────────
92
93/// JSON-Schema description of a callable tool, in OpenAI function-calling
94/// shape.
95///
96/// Why: downstream agents (trusty-memory, trusty-search) expose tools like
97/// `memory_recall` or `web_search` to the LLM. The OpenAI tool format is the
98/// de-facto common denominator across OpenRouter, Ollama, LM Studio, and
99/// most cloud providers.
100/// What: `name` and `description` are passed verbatim; `parameters` is a
101/// JSON Schema object (typically `{"type":"object","properties":{...}}`).
102/// Test: `tool_def_serializes_as_function` checks the wire shape.
103#[derive(Debug, Clone, Serialize, Deserialize)]
104pub struct ToolDef {
105    pub name: String,
106    pub description: String,
107    pub parameters: serde_json::Value,
108}
109
110/// A tool invocation the model wants the host to perform.
111///
112/// Why: the streaming chat API emits `tool_calls` in fragments — first an
113/// `id` + `function.name`, then a string of `function.arguments` deltas.
114/// We accumulate fragments and surface one fully-formed [`ToolCall`] per
115/// invocation to the caller.
116/// What: `id` is the upstream's call id (echoed back in subsequent
117/// `role:"tool"` messages); `name` is the function name; `arguments` is a
118/// JSON string (NOT a parsed value — many models emit malformed JSON and
119/// callers want the raw text for error reporting / repair).
120/// Test: `accumulates_streamed_tool_call_fragments`.
121#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
122pub struct ToolCall {
123    pub id: String,
124    pub name: String,
125    pub arguments: String,
126}
127
128/// Sampling knobs forwarded to an OpenAI-compatible provider.
129///
130/// Why (issue #3758): the streaming request wire previously sent only
131/// `model`/`messages`/`tools`, so a streamed reply's style, verbosity, and
132/// stopping behaviour were NOT equivalent to the blocking path's for the same
133/// turn — the caller silently got provider defaults instead of its configured
134/// temperature, token ceiling, and stop sequences. Carrying them in one struct
135/// (rather than three provider constructor arguments) keeps `new()` stable for
136/// existing callers and lets the set grow without churning call sites.
137/// What: every field is optional — `None`/empty means "omit from the request
138/// body and let the provider default apply", which is exactly the pre-#3758
139/// behaviour, so a caller that does not opt in is unaffected. `stop` maps to
140/// the OpenAI `stop` array (the direct-Anthropic dialect calls the same thing
141/// `stop_sequences`).
142/// Test: `sampling_params_serialize_into_request_body`,
143/// `default_sampling_omits_fields`.
144#[derive(Debug, Clone, Default, PartialEq)]
145pub struct SamplingParams {
146    /// Sampling temperature; `None` omits the field.
147    pub temperature: Option<f32>,
148    /// Maximum tokens to generate; `None` omits the field.
149    pub max_tokens: Option<u32>,
150    /// Stop sequences; empty omits the field.
151    pub stop: Vec<String>,
152}
153
154impl SamplingParams {
155    /// The stop sequences as a request-body-ready `Option<&[String]>`.
156    ///
157    /// Why: an EMPTY `stop` array is not the same as an absent one — some
158    /// OpenAI-compatible servers reject `"stop": []` outright, which is the
159    /// same trap [`ChatProvider`] already documents for an empty `tools`
160    /// array. Collapsing empty to `None` at the boundary keeps that decision
161    /// in one place instead of at each provider.
162    /// What: `None` when empty, `Some(slice)` otherwise.
163    /// Test: `default_sampling_omits_fields`.
164    pub fn stop_slice(&self) -> Option<&[String]> {
165        if self.stop.is_empty() {
166            None
167        } else {
168            Some(&self.stop)
169        }
170    }
171}
172
173/// Token-usage tally surfaced at the end of a streamed response.
174///
175/// Why (issue #3767): Bedrock's `ConverseStream` reports usage only once, in
176/// a terminal `metadata` event — unlike the per-token deltas, it has no
177/// natural home in [`ChatEvent::Delta`]. Without a dedicated variant a
178/// streaming provider has no way to hand token counts back to the caller, so
179/// a downstream cost-reporting consumer would silently see zero usage for
180/// every streamed call. Kept as a small standalone struct (rather than
181/// reusing `trusty_common::inference::types::Usage`) so the `chat` module —
182/// which compiles unconditionally, unlike the `inference-client`-gated
183/// `types` module — never depends on a feature-gated type.
184/// What: four token buckets mirroring the shape every provider in this
185/// workspace already reports (prompt/completion split, plus the two
186/// prompt-cache buckets). The four fields are SEPARATE and ADDITIVE — not
187/// nested — matching both the Bedrock `TokenUsage` wire shape
188/// (`input_tokens`/`output_tokens`/`cache_read_input_tokens`/
189/// `cache_write_input_tokens` are independent fields, verified against
190/// `aws-sdk-bedrockruntime` 1.132.0) and this workspace's existing
191/// convention (`crate::perf::TokenUsage`,
192/// `trusty-agents/src/llm/anthropic_native/mod.rs`'s usage parsing, and
193/// `chat::bedrock_impl::parse_usage`'s non-streaming counterpart). A
194/// provider that has none of the cache fields simply leaves them zero.
195/// Test: `bedrock_stream_reports_usage_from_metadata_event` (bedrock_impl).
196#[derive(Debug, Clone, Copy, Default, PartialEq)]
197pub struct ChatUsage {
198    /// Prompt (input) tokens. Separate from `cache_read_tokens` — a
199    /// cost-accounting consumer that assumes this already includes cache
200    /// reads will double-count; one that assumes it excludes the full
201    /// prompt will undercount. Sum with `cache_read_tokens` for the total
202    /// prompt-side token count when that is what's needed.
203    pub prompt_tokens: u32,
204    /// Completion (output) tokens produced.
205    pub completion_tokens: u32,
206    /// Prompt tokens served from the provider's prompt cache — additive
207    /// with `prompt_tokens`, not a subset of it.
208    pub cache_read_tokens: u32,
209    /// Prompt tokens newly written into the prompt cache this turn.
210    pub cache_creation_tokens: u32,
211}
212
213/// Streaming chat event.
214///
215/// Why: replaces the previous "string-only" channel so callers can
216/// distinguish text deltas from tool invocations and from terminal
217/// success/error without parsing magic markers out of the text stream.
218/// What: `Delta` is a content chunk; `ToolCall` is a fully-accumulated tool
219/// invocation; `Usage` carries the call's token tally (emitted zero or one
220/// times, whenever the provider's wire format reports it — see
221/// [`ChatUsage`]); `Done` signals the upstream stream terminated normally;
222/// `Error` carries a human-readable message for stream-mid failures (the
223/// provider also returns `Err` from `chat_stream`, but `Error` lets the
224/// caller display partial-stream failures inline).
225/// Test: `ollama_provider_streams_sse_deltas`.
226///
227/// # Stability
228///
229/// `#[non_exhaustive]`: adding a variant here is otherwise a SemVer-breaking
230/// change for every downstream crate, because an exhaustive `match` over the
231/// old variant set stops compiling with E0004. That is not hypothetical — the
232/// `Usage` variant forced arm additions in five consumers at once
233/// (`trusty-agents`, `trusty-analyze`, `trusty-memory`, `trusty-mpm`,
234/// `trusty-search`) and is the direct cause of the 0.27.0 MINOR bump: a patch
235/// release would have re-resolved the already-published, arm-less consumer
236/// sources against the new variant and hard-failed `cargo install` (the same
237/// failure that forced `trusty-analyze` 0.7.3 to be yanked).
238///
239/// Downstream matches must therefore carry a wildcard arm. This attribute does
240/// NOT retroactively fix consumers published before it landed — it only stops
241/// the next variant addition from repeating the break.
242#[derive(Debug, Clone)]
243#[non_exhaustive]
244pub enum ChatEvent {
245    Delta(String),
246    ToolCall(ToolCall),
247    Usage(ChatUsage),
248    Done,
249    Error(String),
250}
251
252/// Streaming chat provider abstraction.
253///
254/// Why: downstream crates (trusty-memory, trusty-search) want to support
255/// multiple LLM backends without hard-coding which one to call. Providers
256/// expose a uniform streaming interface so the caller can swap them at
257/// runtime based on configuration / availability.
258/// What: implementors stream [`ChatEvent`]s into `tx`. Pass an empty
259/// `tools` vec to disable tool use entirely (the provider MUST then omit
260/// the `tools` field from the upstream request — some models error on an
261/// empty array). Returning `Ok(())` means the stream completed normally;
262/// the caller should also expect a final [`ChatEvent::Done`].
263/// Test: implementations are covered by their own unit tests in this
264/// module plus integration tests in downstream crates.
265#[async_trait]
266pub trait ChatProvider: Send + Sync {
267    /// Human-readable provider name (e.g. `"openrouter"`, `"ollama"`).
268    fn name(&self) -> &str;
269    /// Model identifier sent on every request.
270    fn model(&self) -> &str;
271    /// Stream chat events into `tx`. `tools` empty disables tool use.
272    async fn chat_stream(
273        &self,
274        messages: Vec<ChatMessage>,
275        tools: Vec<ToolDef>,
276        tx: Sender<ChatEvent>,
277    ) -> Result<()>;
278}
279
280#[cfg(test)]
281mod tests {
282    use super::*;
283
284    #[test]
285    fn local_model_config_defaults() {
286        let cfg = LocalModelConfig::default();
287        assert!(cfg.enabled);
288        assert_eq!(cfg.base_url, "http://localhost:11434");
289        assert_eq!(cfg.model, "qwen3:30b");
290    }
291
292    #[test]
293    fn local_model_config_deserializes_from_toml() {
294        let toml_src = r#"
295            enabled = true
296            base_url = "http://localhost:1234"
297            model = "qwen2.5-coder"
298        "#;
299        let cfg: LocalModelConfig = toml::from_str(toml_src).expect("parse TOML");
300        assert!(cfg.enabled);
301        assert_eq!(cfg.base_url, "http://localhost:1234");
302        assert_eq!(cfg.model, "qwen2.5-coder");
303    }
304}