1use std::fmt;
21use std::time::Duration;
22
23use serde::{Deserialize, Serialize};
24
25use crate::ids::{CallId, ModelKey, ModelRef, ProviderKey, RequestId};
26use crate::request::{ContentPart, ToolCall};
27use crate::structured::StructuredOutputError;
28
29#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
31#[serde(rename_all = "snake_case")]
32#[non_exhaustive]
33pub enum FinishReason {
34 Stop,
36 MaxTokens,
39 ToolCalls,
41 ContentFilter,
43 Refusal,
46 Other,
49}
50
51impl FinishReason {
52 #[must_use]
54 pub const fn as_str(self) -> &'static str {
55 match self {
56 Self::Stop => "stop",
57 Self::MaxTokens => "max_tokens",
58 Self::ToolCalls => "tool_calls",
59 Self::ContentFilter => "content_filter",
60 Self::Refusal => "refusal",
61 Self::Other => "other",
62 }
63 }
64
65 #[must_use]
70 pub const fn is_complete(self) -> bool {
71 matches!(self, Self::Stop | Self::ToolCalls)
72 }
73}
74
75impl fmt::Display for FinishReason {
76 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
77 f.write_str(self.as_str())
78 }
79}
80
81#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default, Serialize, Deserialize)]
87#[serde(deny_unknown_fields)]
88pub struct TokenUsage {
89 pub input: u64,
91 pub output: u64,
93 #[serde(default)]
96 pub cached_input: u64,
97 #[serde(default)]
100 pub reasoning: u64,
101}
102
103impl TokenUsage {
104 #[must_use]
106 pub const fn new(input: u64, output: u64) -> Self {
107 Self {
108 input,
109 output,
110 cached_input: 0,
111 reasoning: 0,
112 }
113 }
114
115 #[must_use]
117 pub const fn none() -> Self {
118 Self::new(0, 0)
119 }
120
121 #[must_use]
123 pub const fn with_cached_input(mut self, cached: u64) -> Self {
124 self.cached_input = cached;
125 self
126 }
127
128 #[must_use]
130 pub const fn with_reasoning(mut self, reasoning: u64) -> Self {
131 self.reasoning = reasoning;
132 self
133 }
134
135 #[must_use]
137 pub const fn total(self) -> u64 {
138 self.input.saturating_add(self.output)
139 }
140
141 #[must_use]
143 pub const fn is_unreported(self) -> bool {
144 self.input == 0 && self.output == 0
145 }
146}
147
148#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
154#[serde(tag = "kind", rename_all = "snake_case")]
155#[non_exhaustive]
156pub enum ResponseWarning {
157 SynthesizedCallIds,
161 UnknownFinishReason {
163 reported: String,
165 },
166 UsageUnreported,
168 FeatureDropped {
173 feature: String,
175 },
176 Reconstructed,
178}
179
180#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
182#[serde(deny_unknown_fields)]
183pub struct ModelResponse {
184 pub request_id: RequestId,
188 pub provider: ProviderKey,
190 pub model: ModelKey,
193 pub content: Vec<ContentPart>,
195 pub finish: FinishReason,
197 #[serde(default)]
199 pub usage: TokenUsage,
200 #[serde(default, skip_serializing_if = "Option::is_none")]
203 pub raw_id: Option<String>,
204 pub latency: Duration,
206 #[serde(default, skip_serializing_if = "Vec::is_empty")]
208 pub warnings: Vec<ResponseWarning>,
209}
210
211impl ModelResponse {
212 #[must_use]
214 pub fn new(
215 request_id: RequestId,
216 provider: impl Into<ProviderKey>,
217 model: impl Into<ModelKey>,
218 ) -> Self {
219 Self {
220 request_id,
221 provider: provider.into(),
222 model: model.into(),
223 content: Vec::new(),
224 finish: FinishReason::Stop,
225 usage: TokenUsage::none(),
226 raw_id: None,
227 latency: Duration::ZERO,
228 warnings: Vec::new(),
229 }
230 }
231
232 #[must_use]
234 pub fn with_part(mut self, part: ContentPart) -> Self {
235 self.content.push(part);
236 self
237 }
238
239 #[must_use]
241 pub fn with_text(self, text: impl Into<String>) -> Self {
242 self.with_part(ContentPart::text(text))
243 }
244
245 #[must_use]
247 pub fn with_tool_call(self, call: ToolCall) -> Self {
248 self.with_part(ContentPart::ToolCall(call))
249 }
250
251 #[must_use]
253 pub const fn with_finish(mut self, finish: FinishReason) -> Self {
254 self.finish = finish;
255 self
256 }
257
258 #[must_use]
260 pub const fn with_usage(mut self, usage: TokenUsage) -> Self {
261 self.usage = usage;
262 self
263 }
264
265 #[must_use]
267 pub fn with_raw_id(mut self, raw_id: impl Into<String>) -> Self {
268 self.raw_id = Some(raw_id.into());
269 self
270 }
271
272 #[must_use]
274 pub const fn with_latency(mut self, latency: Duration) -> Self {
275 self.latency = latency;
276 self
277 }
278
279 #[must_use]
281 pub fn with_warning(mut self, warning: ResponseWarning) -> Self {
282 self.warnings.push(warning);
283 self
284 }
285
286 #[must_use]
288 pub fn reference(&self) -> ModelRef {
289 ModelRef {
290 provider: self.provider.clone(),
291 model: self.model.clone(),
292 }
293 }
294
295 #[must_use]
306 pub fn text(&self) -> String {
307 let mut out = String::new();
308 for part in &self.content {
309 if let Some(text) = part.as_text() {
310 out.push_str(text);
311 }
312 }
313 out
314 }
315
316 #[must_use]
318 pub fn tool_calls(&self) -> Vec<&ToolCall> {
319 self.content
320 .iter()
321 .filter_map(ContentPart::as_tool_call)
322 .collect()
323 }
324
325 #[must_use]
327 pub fn tool_call(&self, id: &CallId) -> Option<&ToolCall> {
328 self.tool_calls().into_iter().find(|call| &call.id == id)
329 }
330
331 #[must_use]
333 pub fn is_empty(&self) -> bool {
334 self.tool_calls().is_empty() && self.text().trim().is_empty()
335 }
336
337 pub fn single_json(&self) -> Result<serde_json::Value, StructuredOutputError> {
361 if self.finish == FinishReason::Refusal {
362 return Err(StructuredOutputError::Refusal);
363 }
364 let calls = self.tool_calls();
365 if calls.len() > 1 {
366 return Err(StructuredOutputError::MultipleCandidates {
367 candidates: calls.len(),
368 });
369 }
370 if !self.finish.is_complete() {
371 return Err(StructuredOutputError::NoOutput);
372 }
373 if let Some(call) = calls.first() {
374 return Ok(call.arguments.clone());
375 }
376 let text = self.text();
377 let payload = strip_code_fence(text.trim());
378 if payload.is_empty() {
379 return Err(StructuredOutputError::NoOutput);
380 }
381 serde_json::from_str(payload).map_err(StructuredOutputError::not_json)
382 }
383}
384
385fn strip_code_fence(text: &str) -> &str {
390 let Some(rest) = text.strip_prefix("```") else {
391 return text;
392 };
393 let Some(body_start) = rest.find('\n') else {
394 return text;
395 };
396 if rest[..body_start]
398 .chars()
399 .any(|ch| !ch.is_ascii_alphanumeric())
400 {
401 return text;
402 }
403 let body = &rest[body_start + 1..];
404 match body.trim_end().strip_suffix("```") {
405 Some(inner) => inner.trim(),
406 None => text,
407 }
408}
409
410#[cfg(test)]
411mod tests {
412 use super::*;
413 use serde_json::json;
414
415 fn response() -> ModelResponse {
416 ModelResponse::new(RequestId::nil(), "openai", "gpt-4o")
417 }
418
419 #[test]
420 fn helpers_read_text_and_calls_in_order() {
421 let built = response()
422 .with_text("first ")
423 .with_tool_call(ToolCall::new("call_a", "plan", json!({"n": 1})))
424 .with_text("second")
425 .with_tool_call(ToolCall::new("call_b", "plan", json!({"n": 2})));
426 assert_eq!(built.text(), "first second");
427 let calls = built.tool_calls();
428 assert_eq!(calls.len(), 2);
429 assert_eq!(calls[0].id.as_str(), "call_a");
430 assert_eq!(calls[1].id.as_str(), "call_b");
431 assert!(built.tool_call(&CallId::from("call_b")).is_some());
432 assert!(built.tool_call(&CallId::from("call_z")).is_none());
433 assert!(!built.is_empty());
434 assert_eq!(built.reference().to_string(), "openai/gpt-4o");
435 }
436
437 #[test]
438 fn single_json_takes_the_only_tool_call_and_ignores_the_preamble() {
439 let built = response()
440 .with_text("Certo, ecco il piano.")
441 .with_tool_call(ToolCall::new("call_a", "plan", json!({"acts": []})))
442 .with_finish(FinishReason::ToolCalls);
443 assert_eq!(built.single_json().unwrap(), json!({"acts": []}));
444 }
445
446 #[test]
447 fn single_json_refuses_to_choose_between_two_calls() {
448 let built = response()
449 .with_tool_call(ToolCall::new("a", "plan", json!({"n": 1})))
450 .with_tool_call(ToolCall::new("b", "plan", json!({"n": 2})))
451 .with_finish(FinishReason::ToolCalls);
452 assert!(matches!(
453 built.single_json(),
454 Err(StructuredOutputError::MultipleCandidates { candidates: 2 })
455 ));
456 }
457
458 #[test]
459 fn single_json_parses_text_and_strips_one_fence() {
460 let plain = response().with_text(" {\"a\": 1} ");
461 assert_eq!(plain.single_json().unwrap(), json!({"a": 1}));
462
463 let fenced = response().with_text("```json\n{\"a\": 1}\n```");
464 assert_eq!(fenced.single_json().unwrap(), json!({"a": 1}));
465
466 let bare_fence = response().with_text("```\n{\"a\": 1}\n```");
467 assert_eq!(bare_fence.single_json().unwrap(), json!({"a": 1}));
468
469 let not_a_fence = response().with_text("```json {\"a\": 1}");
471 assert!(matches!(
472 not_a_fence.single_json(),
473 Err(StructuredOutputError::NotJson { .. })
474 ));
475 }
476
477 #[test]
478 fn single_json_never_parses_an_incomplete_answer() {
479 let truncated = response()
480 .with_text("{\"a\": 1")
481 .with_finish(FinishReason::MaxTokens);
482 assert!(matches!(
483 truncated.single_json(),
484 Err(StructuredOutputError::NoOutput)
485 ));
486
487 let filtered = response()
488 .with_text("{\"a\": 1}")
489 .with_finish(FinishReason::ContentFilter);
490 assert!(matches!(
491 filtered.single_json(),
492 Err(StructuredOutputError::NoOutput)
493 ));
494
495 let refused = response()
496 .with_text("I cannot help with that.")
497 .with_finish(FinishReason::Refusal);
498 assert!(matches!(
499 refused.single_json(),
500 Err(StructuredOutputError::Refusal)
501 ));
502
503 let empty = response();
504 assert!(empty.is_empty());
505 assert!(matches!(
506 empty.single_json(),
507 Err(StructuredOutputError::NoOutput)
508 ));
509 }
510
511 #[test]
512 fn finish_reasons_say_whether_output_is_complete() {
513 assert!(FinishReason::Stop.is_complete());
514 assert!(FinishReason::ToolCalls.is_complete());
515 for incomplete in [
516 FinishReason::MaxTokens,
517 FinishReason::ContentFilter,
518 FinishReason::Refusal,
519 FinishReason::Other,
520 ] {
521 assert!(!incomplete.is_complete(), "{incomplete}");
522 }
523 }
524
525 #[test]
526 fn usage_totals_and_reports_absence() {
527 assert!(TokenUsage::none().is_unreported());
528 let usage = TokenUsage::new(100, 20)
529 .with_cached_input(80)
530 .with_reasoning(5);
531 assert_eq!(usage.total(), 120);
532 assert_eq!(usage.cached_input, 80);
533 assert_eq!(usage.reasoning, 5);
534 assert!(!usage.is_unreported());
535 }
536
537 #[test]
538 fn responses_round_trip_through_serde() {
539 let built = response()
540 .with_text("hi")
541 .with_tool_call(ToolCall::new("c1", "plan", json!({})))
542 .with_finish(FinishReason::ToolCalls)
543 .with_usage(TokenUsage::new(10, 3))
544 .with_raw_id("resp_123")
545 .with_latency(Duration::from_millis(420))
546 .with_warning(ResponseWarning::Reconstructed);
547 let json = serde_json::to_string(&built).unwrap();
548 let back: ModelResponse = serde_json::from_str(&json).unwrap();
549 assert_eq!(back, built);
550 }
551}