Skip to main content

inference_gateway_sdk/
lib.rs

1//! Inference Gateway SDK for Rust
2//!
3//! This crate provides a Rust client for the Inference Gateway API, allowing interaction
4//! with various LLM providers through a unified interface.
5//!
6//! Data types in [`crate::generated::schemas`] are generated from the upstream
7//! `openapi.yaml` and re-exported at the crate root. Run `task generate-types`
8//! to regenerate them after a spec bump.
9
10mod ext;
11mod generated;
12
13pub use generated::schemas::*;
14
15use std::future::Future;
16
17use futures_util::{Stream, StreamExt};
18use reqwest::multipart::{Form, Part};
19use reqwest::{Client, StatusCode};
20use thiserror::Error;
21
22/// Stream of Server-Sent Events (SSE) yielded by [`InferenceGatewayAPI::generate_content_stream`].
23///
24/// This is the SDK's own SSE wrapper used by the streaming function. It is distinct
25/// from the spec's [`SsEvent`] (which constrains `event` to a fixed enum) - the
26/// streaming function may surface arbitrary event names produced by upstream providers.
27#[derive(Debug, serde::Serialize, serde::Deserialize)]
28pub struct SSEvents {
29    pub data: String,
30    pub event: Option<String>,
31    pub retry: Option<u64>,
32}
33
34/// Custom error types for the Inference Gateway SDK
35#[derive(Error, Debug)]
36pub enum GatewayError {
37    #[error("Unauthorized: {0}")]
38    Unauthorized(String),
39
40    #[error("Forbidden: {0}")]
41    Forbidden(String),
42
43    #[error("Not found: {0}")]
44    NotFound(String),
45
46    #[error("Bad request: {0}")]
47    BadRequest(String),
48
49    #[error("Internal server error: {0}")]
50    InternalError(String),
51
52    #[error("Stream error: {0}")]
53    StreamError(reqwest::Error),
54
55    #[error("Decoding error: {0}")]
56    DecodingError(std::string::FromUtf8Error),
57
58    #[error("Request error: {0}")]
59    RequestError(#[from] reqwest::Error),
60
61    #[error("Deserialization error: {0}")]
62    DeserializationError(serde_json::Error),
63
64    #[error("Serialization error: {0}")]
65    SerializationError(#[from] serde_json::Error),
66
67    #[error("Other error: {0}")]
68    Other(#[from] Box<dyn std::error::Error + Send + Sync>),
69}
70
71/// Request for [`InferenceGatewayAPI::create_image_edit`].
72///
73/// The edits endpoint takes `multipart/form-data` with binary image uploads,
74/// which the codegen cannot express, so this type is hand-written.
75#[derive(Debug, Clone, Default)]
76pub struct CreateImageEditRequest {
77    /// The image to edit, as raw file bytes (png/webp/jpg).
78    pub image: Vec<u8>,
79    /// A text description of the desired image.
80    pub prompt: String,
81    /// Optional PNG mask whose transparent areas indicate where to edit.
82    pub mask: Option<Vec<u8>>,
83    pub model: Option<String>,
84    /// Number of images to generate (1-10).
85    pub n: Option<i64>,
86    pub size: Option<ImageSize>,
87    /// `auto`, `standard`, `low`, `medium`, or `high`.
88    pub quality: Option<String>,
89    /// `url` or `b64_json`.
90    pub response_format: Option<String>,
91}
92
93/// Request for [`InferenceGatewayAPI::create_image_variation`].
94///
95/// The variations endpoint takes `multipart/form-data` with a binary image
96/// upload, which the codegen cannot express, so this type is hand-written.
97#[derive(Debug, Clone, Default)]
98pub struct CreateImageVariationRequest {
99    /// The image to base the variation on, as raw file bytes (png/webp/jpg).
100    pub image: Vec<u8>,
101    pub model: Option<String>,
102    /// Number of images to generate (1-10).
103    pub n: Option<i64>,
104    pub size: Option<ImageSize>,
105    /// `url` or `b64_json`.
106    pub response_format: Option<String>,
107}
108
109/// Client for interacting with the Inference Gateway API
110pub struct InferenceGatewayClient {
111    base_url: String,
112    client: Client,
113    token: Option<String>,
114    tools: Option<Vec<ChatCompletionTool>>,
115    max_tokens: Option<i64>,
116}
117
118impl std::fmt::Debug for InferenceGatewayClient {
119    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
120        f.debug_struct("InferenceGatewayClient")
121            .field("base_url", &self.base_url)
122            .field("token", &self.token.as_ref().map(|_| "*****"))
123            .finish()
124    }
125}
126
127/// Core API interface for the Inference Gateway
128pub trait InferenceGatewayAPI {
129    /// Lists available models from all providers
130    fn list_models(&self) -> impl Future<Output = Result<ListModelsResponse, GatewayError>> + Send;
131
132    /// Lists available models filtered by a specific provider
133    fn list_models_by_provider(
134        &self,
135        provider: Provider,
136    ) -> impl Future<Output = Result<ListModelsResponse, GatewayError>> + Send;
137
138    /// Lists available models with additional metadata included, optionally
139    /// filtered by a provider. Supported `include` values: `"pricing"`,
140    /// `"context_window"`, `"modalities"` - they populate
141    /// [`Model::pricing`], [`Model::context_window`], and
142    /// [`Model::modalities`] respectively.
143    fn list_models_with_include(
144        &self,
145        provider: Option<Provider>,
146        include: &[&str],
147    ) -> impl Future<Output = Result<ListModelsResponse, GatewayError>> + Send;
148
149    /// Generates content using a specified model
150    fn generate_content(
151        &self,
152        provider: Provider,
153        model: &str,
154        messages: Vec<Message>,
155    ) -> impl Future<Output = Result<CreateChatCompletionResponse, GatewayError>> + Send;
156
157    /// Streams content generation as SSE events from the gateway.
158    fn generate_content_stream(
159        &self,
160        provider: Provider,
161        model: &str,
162        messages: Vec<Message>,
163    ) -> impl Stream<Item = Result<SSEvents, GatewayError>> + Send;
164
165    /// Creates a message via the Anthropic-compatible Messages API.
166    ///
167    /// Providers without Messages support return [`GatewayError::BadRequest`];
168    /// use [`InferenceGatewayAPI::generate_content`] for those providers.
169    fn create_message(
170        &self,
171        provider: Option<Provider>,
172        request: CreateMessagesRequest,
173    ) -> impl Future<Output = Result<MessagesResponse, GatewayError>> + Send;
174
175    /// Streams a message via the Messages API as SSE events. Each event's
176    /// `data` field holds a JSON-serialized [`MessagesStreamEvent`].
177    fn create_message_stream(
178        &self,
179        provider: Option<Provider>,
180        request: CreateMessagesRequest,
181    ) -> impl Stream<Item = Result<SSEvents, GatewayError>> + Send;
182
183    /// Lists available MCP tools (only when `EXPOSE_MCP=true` server-side)
184    fn list_tools(&self) -> impl Future<Output = Result<ListToolsResponse, GatewayError>> + Send;
185
186    /// Generates images using a specified model via the OpenAI-compatible Images API.
187    ///
188    /// Providers without Images support return [`GatewayError::BadRequest`];
189    /// use [`InferenceGatewayAPI::generate_content`] for those providers.
190    fn generate_image(
191        &self,
192        provider: Provider,
193        request: CreateImageRequest,
194    ) -> impl Future<Output = Result<ImagesResponse, GatewayError>> + Send;
195
196    /// Edits or extends an image via the OpenAI-compatible Images API
197    /// (`POST /images/edits`, multipart/form-data).
198    ///
199    /// Providers without Images support return [`GatewayError::BadRequest`].
200    fn create_image_edit(
201        &self,
202        provider: Option<Provider>,
203        request: CreateImageEditRequest,
204    ) -> impl Future<Output = Result<ImagesResponse, GatewayError>> + Send;
205
206    /// Creates a variation of an image via the OpenAI-compatible Images API
207    /// (`POST /images/variations`, multipart/form-data).
208    ///
209    /// Providers without Images support return [`GatewayError::BadRequest`].
210    fn create_image_variation(
211        &self,
212        provider: Option<Provider>,
213        request: CreateImageVariationRequest,
214    ) -> impl Future<Output = Result<ImagesResponse, GatewayError>> + Send;
215
216    /// Generates speech audio from text via the OpenAI-compatible Audio API
217    /// (`POST /audio/speech`), returning the synthesized audio as raw bytes.
218    ///
219    /// The response format (and thus the bytes' encoding) is chosen via
220    /// `request.response_format`; it defaults to `mp3`.
221    ///
222    /// Providers without Audio support return [`GatewayError::BadRequest`];
223    /// use [`InferenceGatewayAPI::generate_content`] for those providers.
224    fn create_speech(
225        &self,
226        provider: Option<Provider>,
227        request: CreateSpeechRequest,
228    ) -> impl Future<Output = Result<Vec<u8>, GatewayError>> + Send;
229
230    /// Health probe - returns true on HTTP 200, false otherwise.
231    fn health_check(&self) -> impl Future<Output = Result<bool, GatewayError>> + Send;
232}
233
234impl InferenceGatewayClient {
235    /// Creates a new client targeting `base_url`.
236    pub fn new(base_url: &str) -> Self {
237        Self {
238            base_url: base_url.to_string(),
239            client: Client::new(),
240            token: None,
241            tools: None,
242            max_tokens: None,
243        }
244    }
245
246    /// Creates a client using `INFERENCE_GATEWAY_URL` (or `http://localhost:8080/v1`).
247    pub fn new_default() -> Self {
248        let base_url = std::env::var("INFERENCE_GATEWAY_URL")
249            .unwrap_or_else(|_| "http://localhost:8080/v1".to_string());
250
251        Self {
252            base_url,
253            client: Client::new(),
254            token: None,
255            tools: None,
256            max_tokens: None,
257        }
258    }
259
260    pub fn base_url(&self) -> &str {
261        &self.base_url
262    }
263
264    /// Sets the tools used for subsequent generations.
265    pub fn with_tools(mut self, tools: Option<Vec<ChatCompletionTool>>) -> Self {
266        self.tools = tools;
267        self
268    }
269
270    /// Sets the bearer token used for authentication.
271    pub fn with_token(mut self, token: impl Into<String>) -> Self {
272        self.token = Some(token.into());
273        self
274    }
275
276    /// Sets an upper bound for tokens generated per request.
277    pub fn with_max_tokens(mut self, max_tokens: Option<i64>) -> Self {
278        self.max_tokens = max_tokens;
279        self
280    }
281
282    /// The gateway serves `/health` from the root server, not under the
283    /// versioned API prefix, so this strips a trailing `/v<digits>` segment
284    /// from the configured base URL before appending `/health`.
285    fn health_url(&self) -> String {
286        let trimmed = self.base_url.trim_end_matches('/');
287        let root = match trimmed.rsplit_once('/') {
288            Some((prefix, last))
289                if last.len() >= 2
290                    && last.starts_with('v')
291                    && last[1..].chars().all(|c| c.is_ascii_digit()) =>
292            {
293                prefix
294            }
295            _ => trimmed,
296        };
297        format!("{root}/health")
298    }
299
300    fn messages_url(&self, provider: Option<Provider>) -> String {
301        match provider {
302            Some(provider) => format!("{}/messages?provider={provider}", self.base_url),
303            None => format!("{}/messages", self.base_url),
304        }
305    }
306
307    fn build_chat_request(
308        &self,
309        model: &str,
310        messages: Vec<Message>,
311        stream: bool,
312    ) -> CreateChatCompletionRequest {
313        // `tools` and `max_tokens` are deliberately omitted from streaming
314        // requests; every other field falls back to the schema defaults via
315        // `Default`. See CLAUDE.md for the streaming asymmetry.
316        CreateChatCompletionRequest {
317            model: model.to_string(),
318            messages,
319            stream,
320            tools: if stream {
321                Vec::new()
322            } else {
323                self.tools.clone().unwrap_or_default()
324            },
325            max_tokens: if stream { None } else { self.max_tokens },
326            ..Default::default()
327        }
328    }
329}
330
331async fn map_error_status(status: StatusCode, response: reqwest::Response) -> GatewayError {
332    // Gateway errors are `{"error": "..."}`; Messages endpoints use the
333    // Anthropic shape `{"type": "error", "error": {"type": ..., "message": ...}}`.
334    let fallback = || status.canonical_reason().unwrap_or("unknown").to_string();
335    let message = match response.json::<serde_json::Value>().await {
336        Ok(body) => match body.get("error") {
337            Some(serde_json::Value::String(error)) => error.clone(),
338            Some(error) => error
339                .get("message")
340                .and_then(|m| m.as_str())
341                .map(str::to_string)
342                .unwrap_or_else(fallback),
343            None => fallback(),
344        },
345        Err(_) => fallback(),
346    };
347    match status {
348        StatusCode::UNAUTHORIZED => GatewayError::Unauthorized(message),
349        StatusCode::FORBIDDEN => GatewayError::Forbidden(message),
350        StatusCode::NOT_FOUND => GatewayError::NotFound(message),
351        StatusCode::BAD_REQUEST => GatewayError::BadRequest(message),
352        StatusCode::INTERNAL_SERVER_ERROR => GatewayError::InternalError(message),
353        other => GatewayError::Other(Box::new(std::io::Error::other(format!(
354            "Unexpected status code: {other}"
355        )))),
356    }
357}
358
359fn sse_stream<B>(
360    client: Client,
361    token: Option<String>,
362    url: String,
363    body: B,
364) -> impl Stream<Item = Result<SSEvents, GatewayError>> + Send
365where
366    B: serde::Serialize + Send + 'static,
367{
368    async_stream::try_stream! {
369        let mut request = client.post(&url);
370        if let Some(token) = token {
371            request = request.bearer_auth(token);
372        }
373        let response = request.json(&body).send().await?;
374        let mut stream = response.bytes_stream();
375        let mut current_event: Option<String> = None;
376        let mut current_data: Option<String> = None;
377
378        while let Some(chunk) = stream.next().await {
379            let chunk = chunk?;
380            let chunk_str = String::from_utf8_lossy(&chunk);
381
382            for line in chunk_str.lines() {
383                if line.is_empty() && current_data.is_some() {
384                    yield SSEvents {
385                        data: current_data.take().unwrap(),
386                        event: current_event.take(),
387                        retry: None,
388                    };
389                    continue;
390                }
391
392                if let Some(event) = line.strip_prefix("event:") {
393                    current_event = Some(event.trim().to_string());
394                } else if let Some(data) = line.strip_prefix("data:") {
395                    let processed_data = data.strip_suffix('\n').unwrap_or(data);
396                    current_data = Some(processed_data.trim().to_string());
397                }
398            }
399        }
400    }
401}
402
403impl InferenceGatewayClient {
404    async fn fetch_models(&self, query: &str) -> Result<ListModelsResponse, GatewayError> {
405        let url = if query.is_empty() {
406            format!("{}/models", self.base_url)
407        } else {
408            format!("{}/models?{}", self.base_url, query)
409        };
410        let mut request = self.client.get(&url);
411        if let Some(token) = &self.token {
412            request = request.bearer_auth(token);
413        }
414
415        let response = request.send().await?;
416        match response.status() {
417            StatusCode::OK => Ok(response.json().await?),
418            status => Err(map_error_status(status, response).await),
419        }
420    }
421}
422
423impl InferenceGatewayAPI for InferenceGatewayClient {
424    async fn list_models(&self) -> Result<ListModelsResponse, GatewayError> {
425        self.fetch_models("").await
426    }
427
428    async fn list_models_by_provider(
429        &self,
430        provider: Provider,
431    ) -> Result<ListModelsResponse, GatewayError> {
432        self.fetch_models(&format!("provider={provider}")).await
433    }
434
435    async fn list_models_with_include(
436        &self,
437        provider: Option<Provider>,
438        include: &[&str],
439    ) -> Result<ListModelsResponse, GatewayError> {
440        let mut query = Vec::new();
441        if let Some(provider) = provider {
442            query.push(format!("provider={provider}"));
443        }
444        if !include.is_empty() {
445            query.push(format!("include={}", include.join(",")));
446        }
447        self.fetch_models(&query.join("&")).await
448    }
449
450    async fn generate_content(
451        &self,
452        provider: Provider,
453        model: &str,
454        messages: Vec<Message>,
455    ) -> Result<CreateChatCompletionResponse, GatewayError> {
456        let url = format!("{}/chat/completions?provider={}", self.base_url, provider);
457        let mut request = self.client.post(&url);
458        if let Some(token) = &self.token {
459            request = request.bearer_auth(token);
460        }
461
462        let payload = self.build_chat_request(model, messages, false);
463        let response = request.json(&payload).send().await?;
464
465        match response.status() {
466            StatusCode::OK => Ok(response.json().await?),
467            status => Err(map_error_status(status, response).await),
468        }
469    }
470
471    fn generate_content_stream(
472        &self,
473        provider: Provider,
474        model: &str,
475        messages: Vec<Message>,
476    ) -> impl Stream<Item = Result<SSEvents, GatewayError>> + Send {
477        let url = format!("{}/chat/completions?provider={}", self.base_url, provider);
478        let request_body = self.build_chat_request(model, messages, true);
479        sse_stream(self.client.clone(), self.token.clone(), url, request_body)
480    }
481
482    async fn create_message(
483        &self,
484        provider: Option<Provider>,
485        mut request: CreateMessagesRequest,
486    ) -> Result<MessagesResponse, GatewayError> {
487        request.stream = false;
488        let mut req = self.client.post(self.messages_url(provider));
489        if let Some(token) = &self.token {
490            req = req.bearer_auth(token);
491        }
492
493        let response = req.json(&request).send().await?;
494        match response.status() {
495            StatusCode::OK => Ok(response.json().await?),
496            status => Err(map_error_status(status, response).await),
497        }
498    }
499
500    fn create_message_stream(
501        &self,
502        provider: Option<Provider>,
503        mut request: CreateMessagesRequest,
504    ) -> impl Stream<Item = Result<SSEvents, GatewayError>> + Send {
505        request.stream = true;
506        sse_stream(
507            self.client.clone(),
508            self.token.clone(),
509            self.messages_url(provider),
510            request,
511        )
512    }
513
514    async fn list_tools(&self) -> Result<ListToolsResponse, GatewayError> {
515        let url = format!("{}/mcp/tools", self.base_url);
516        let mut request = self.client.get(&url);
517        if let Some(token) = &self.token {
518            request = request.bearer_auth(token);
519        }
520
521        let response = request.send().await?;
522        match response.status() {
523            StatusCode::OK => Ok(response.json().await?),
524            status => Err(map_error_status(status, response).await),
525        }
526    }
527
528    async fn generate_image(
529        &self,
530        provider: Provider,
531        request: CreateImageRequest,
532    ) -> Result<ImagesResponse, GatewayError> {
533        let url = format!("{}/images/generations?provider={}", self.base_url, provider);
534        let mut req = self.client.post(&url);
535        if let Some(token) = &self.token {
536            req = req.bearer_auth(token);
537        }
538        let response = req.json(&request).send().await?;
539        match response.status() {
540            StatusCode::OK => Ok(response.json().await?),
541            status => Err(map_error_status(status, response).await),
542        }
543    }
544
545    async fn create_image_edit(
546        &self,
547        provider: Option<Provider>,
548        request: CreateImageEditRequest,
549    ) -> Result<ImagesResponse, GatewayError> {
550        let mut url = format!("{}/images/edits", self.base_url);
551        if let Some(provider) = provider {
552            url = format!("{url}?provider={provider}");
553        }
554        let mut form = Form::new()
555            .part("image", Part::bytes(request.image).file_name("image"))
556            .text("prompt", request.prompt);
557        if let Some(mask) = request.mask {
558            form = form.part("mask", Part::bytes(mask).file_name("mask"));
559        }
560        if let Some(model) = request.model {
561            form = form.text("model", model);
562        }
563        if let Some(n) = request.n {
564            form = form.text("n", n.to_string());
565        }
566        if let Some(size) = request.size {
567            form = form.text("size", size.to_string());
568        }
569        if let Some(quality) = request.quality {
570            form = form.text("quality", quality);
571        }
572        if let Some(response_format) = request.response_format {
573            form = form.text("response_format", response_format);
574        }
575        let mut req = self.client.post(&url);
576        if let Some(token) = &self.token {
577            req = req.bearer_auth(token);
578        }
579        let response = req.multipart(form).send().await?;
580        match response.status() {
581            StatusCode::OK => Ok(response.json().await?),
582            status => Err(map_error_status(status, response).await),
583        }
584    }
585
586    async fn create_image_variation(
587        &self,
588        provider: Option<Provider>,
589        request: CreateImageVariationRequest,
590    ) -> Result<ImagesResponse, GatewayError> {
591        let mut url = format!("{}/images/variations", self.base_url);
592        if let Some(provider) = provider {
593            url = format!("{url}?provider={provider}");
594        }
595        let mut form = Form::new().part("image", Part::bytes(request.image).file_name("image"));
596        if let Some(model) = request.model {
597            form = form.text("model", model);
598        }
599        if let Some(n) = request.n {
600            form = form.text("n", n.to_string());
601        }
602        if let Some(size) = request.size {
603            form = form.text("size", size.to_string());
604        }
605        if let Some(response_format) = request.response_format {
606            form = form.text("response_format", response_format);
607        }
608        let mut req = self.client.post(&url);
609        if let Some(token) = &self.token {
610            req = req.bearer_auth(token);
611        }
612        let response = req.multipart(form).send().await?;
613        match response.status() {
614            StatusCode::OK => Ok(response.json().await?),
615            status => Err(map_error_status(status, response).await),
616        }
617    }
618
619    async fn create_speech(
620        &self,
621        provider: Option<Provider>,
622        request: CreateSpeechRequest,
623    ) -> Result<Vec<u8>, GatewayError> {
624        let mut url = format!("{}/audio/speech", self.base_url);
625        if let Some(provider) = provider {
626            url = format!("{url}?provider={provider}");
627        }
628        let mut req = self.client.post(&url);
629        if let Some(token) = &self.token {
630            req = req.bearer_auth(token);
631        }
632        let response = req.json(&request).send().await?;
633        match response.status() {
634            StatusCode::OK => Ok(response.bytes().await?.to_vec()),
635            status => Err(map_error_status(status, response).await),
636        }
637    }
638
639    async fn health_check(&self) -> Result<bool, GatewayError> {
640        let response = self.client.get(self.health_url()).send().await?;
641        Ok(response.status() == StatusCode::OK)
642    }
643}
644
645#[cfg(test)]
646mod tests;