Skip to main content

tokenmiser_providers/
openai.rs

1//! OpenAI client (also covers DeepSeek, Cerebras, DeepInfra — anything that
2//! speaks OpenAI's `/v1/chat/completions` wire shape).
3
4use async_trait::async_trait;
5use bytes::Bytes;
6use futures::stream::{BoxStream, StreamExt};
7use reqwest::header::{HeaderMap, HeaderValue, AUTHORIZATION, CONTENT_TYPE};
8use tokenmiser_config::ProviderConfig;
9
10use crate::{ChatRequest, ChatResponse, Provider, ProviderError, StreamChunk};
11
12pub struct OpenAIProvider {
13    cfg: ProviderConfig,
14    client: reqwest::Client,
15    api_key: Option<String>,
16}
17
18impl OpenAIProvider {
19    pub fn new(cfg: ProviderConfig) -> Self {
20        let api_key = cfg.api_key_env.as_ref().and_then(|k| std::env::var(k).ok());
21
22        let client = reqwest::Client::builder()
23            .pool_max_idle_per_host(64)
24            .build()
25            .expect("reqwest client construction");
26
27        Self {
28            cfg,
29            client,
30            api_key,
31        }
32    }
33
34    fn headers(&self) -> Result<HeaderMap, ProviderError> {
35        let mut h = HeaderMap::new();
36        h.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
37        if let Some(key) = &self.api_key {
38            let val = HeaderValue::from_str(&format!("Bearer {}", key))
39                .map_err(|e| ProviderError::Malformed(format!("invalid api key: {e}")))?;
40            h.insert(AUTHORIZATION, val);
41        }
42        Ok(h)
43    }
44}
45
46#[async_trait]
47impl Provider for OpenAIProvider {
48    fn name(&self) -> &str {
49        &self.cfg.name
50    }
51
52    fn config(&self) -> &ProviderConfig {
53        &self.cfg
54    }
55
56    async fn complete(&self, req: &ChatRequest) -> Result<ChatResponse, ProviderError> {
57        if self.cfg.api_key_env.is_some() && self.api_key.is_none() {
58            return Err(ProviderError::MissingApiKey(
59                self.cfg.api_key_env.clone().unwrap_or_default(),
60            ));
61        }
62
63        let url = format!(
64            "{}/chat/completions",
65            self.cfg.base_url.trim_end_matches('/')
66        );
67        let mut body = req.clone();
68        body.stream = Some(false); // v0.1: non-streaming only
69
70        let res = self
71            .client
72            .post(&url)
73            .headers(self.headers()?)
74            .json(&body)
75            .send()
76            .await?;
77
78        let status = res.status();
79        let text = res.text().await?;
80
81        if !status.is_success() {
82            return Err(ProviderError::Upstream {
83                status: status.as_u16(),
84                body: text,
85            });
86        }
87
88        let parsed: ChatResponse = serde_json::from_str(&text)?;
89        Ok(parsed)
90    }
91
92    async fn stream(
93        &self,
94        req: &ChatRequest,
95    ) -> Result<BoxStream<'static, Result<StreamChunk, ProviderError>>, ProviderError> {
96        if self.cfg.api_key_env.is_some() && self.api_key.is_none() {
97            return Err(ProviderError::MissingApiKey(
98                self.cfg.api_key_env.clone().unwrap_or_default(),
99            ));
100        }
101
102        let url = format!(
103            "{}/chat/completions",
104            self.cfg.base_url.trim_end_matches('/')
105        );
106        let mut body = req.clone();
107        body.stream = Some(true);
108
109        let res = self
110            .client
111            .post(&url)
112            .headers(self.headers()?)
113            .json(&body)
114            .send()
115            .await?;
116
117        let status = res.status();
118        if !status.is_success() {
119            let text = res.text().await?;
120            return Err(ProviderError::Upstream {
121                status: status.as_u16(),
122                body: text,
123            });
124        }
125
126        // OpenAI/Ollama send SSE on this endpoint. Pass through raw bytes;
127        // proxy-side normalization (cross-provider tool-call diffs, usage
128        // synthesis) is a v0.7.1 follow-up.
129        let stream = res
130            .bytes_stream()
131            .map(|chunk_res| match chunk_res {
132                Ok(b) => Ok(StreamChunk::Sse(Bytes::from(b.to_vec()))),
133                Err(e) => Err(ProviderError::Http(e)),
134            })
135            .chain(futures::stream::once(async { Ok(StreamChunk::Done) }));
136
137        Ok(stream.boxed())
138    }
139}