1use serde::{Deserialize, Serialize};
2use serde_json::Value;
3
4use crate::ProviderError;
5
6#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
7#[serde(rename_all = "snake_case")]
8pub enum Surface {
9 ChatCompletions,
10 Responses,
11}
12
13#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
14pub struct Capabilities {
15 pub chat: bool,
16 pub responses: bool,
17 pub vision: bool,
18 pub reasoning: bool,
19 pub embeddings: bool,
20}
21
22#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
23pub struct ModelUsage {
24 pub input_tokens: u64,
25 pub output_tokens: u64,
26 #[serde(default)]
27 pub reasoning_tokens: u64,
28 pub cache_read_tokens: u64,
29 pub cache_write_tokens: u64,
30}
31
32impl ModelUsage {
33 pub fn total_tokens(self) -> u64 {
34 self.input_tokens
35 .saturating_add(self.output_tokens)
36 .saturating_add(self.cache_read_tokens)
37 .saturating_add(self.cache_write_tokens)
38 }
39}
40
41#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
42pub struct ProviderRequest {
43 pub model: String,
44 pub body: Value,
45}
46
47#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
48pub struct ProviderResponse {
49 pub body: Value,
50 pub usage: ModelUsage,
51}
52
53#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
54pub enum ProviderStreamEvent {
55 Data { event: Option<String>, data: Value },
56 Done(ModelUsage),
57}
58
59pub trait ProviderStreamDecoder: Send {
60 fn decode(&mut self, event: crate::SseEvent)
61 -> Result<Vec<ProviderStreamEvent>, ProviderError>;
62
63 fn finish(&mut self) -> Result<Vec<ProviderStreamEvent>, ProviderError> {
64 Ok(Vec::new())
65 }
66}
67
68pub trait ProviderAdapter: Send + Sync {
69 fn name(&self) -> &'static str;
70 fn capabilities(&self) -> Capabilities;
71 fn encode_request(
72 &self,
73 surface: Surface,
74 request: ProviderRequest,
75 ) -> Result<Value, ProviderError>;
76 fn decode_response(
77 &self,
78 surface: Surface,
79 response: Value,
80 ) -> Result<ProviderResponse, ProviderError>;
81 fn stream_decoder(
82 &self,
83 surface: Surface,
84 ) -> Result<Box<dyn ProviderStreamDecoder>, ProviderError>;
85}
86
87pub(crate) fn chat_usage(value: &Value) -> ModelUsage {
88 let usage = value.get("usage").unwrap_or(value);
89 let input_tokens = usage
90 .get("prompt_tokens")
91 .or_else(|| usage.get("input_tokens"))
92 .and_then(Value::as_u64)
93 .unwrap_or_default();
94 let openai_cache_read_tokens = usage
95 .pointer("/prompt_tokens_details/cached_tokens")
96 .or_else(|| usage.pointer("/input_tokens_details/cached_tokens"))
97 .and_then(Value::as_u64);
98 ModelUsage {
99 input_tokens: input_tokens.saturating_sub(openai_cache_read_tokens.unwrap_or_default()),
100 output_tokens: usage
101 .get("completion_tokens")
102 .or_else(|| usage.get("output_tokens"))
103 .and_then(Value::as_u64)
104 .unwrap_or_default(),
105 reasoning_tokens: usage
106 .pointer("/completion_tokens_details/reasoning_tokens")
107 .or_else(|| usage.pointer("/output_tokens_details/reasoning_tokens"))
108 .and_then(Value::as_u64)
109 .unwrap_or_default(),
110 cache_read_tokens: openai_cache_read_tokens
111 .or_else(|| usage.get("cache_read_input_tokens").and_then(Value::as_u64))
112 .unwrap_or_default(),
113 cache_write_tokens: usage
114 .get("cache_creation_input_tokens")
115 .and_then(Value::as_u64)
116 .unwrap_or_default(),
117 }
118}
119
120#[cfg(test)]
121mod tests {
122 use super::*;
123 use serde_json::json;
124
125 #[test]
126 fn openai_cached_input_tokens_are_normalized_to_disjoint_usage() {
127 assert_eq!(
128 chat_usage(&json!({
129 "prompt_tokens": 19,
130 "completion_tokens": 7,
131 "prompt_tokens_details": { "cached_tokens": 3 }
132 })),
133 ModelUsage {
134 input_tokens: 16,
135 output_tokens: 7,
136 cache_read_tokens: 3,
137 ..ModelUsage::default()
138 }
139 );
140 }
141
142 #[test]
143 fn anthropic_cache_read_tokens_are_already_disjoint() {
144 assert_eq!(
145 chat_usage(&json!({
146 "input_tokens": 19,
147 "output_tokens": 7,
148 "cache_read_input_tokens": 3
149 })),
150 ModelUsage {
151 input_tokens: 19,
152 output_tokens: 7,
153 cache_read_tokens: 3,
154 ..ModelUsage::default()
155 }
156 );
157 }
158}