Skip to main content

gateway_core/
provider.rs

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}