tokenmiser_providers/
openai.rs1use 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); 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 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}