1mod 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#[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#[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#[derive(Debug, Clone, Default)]
76pub struct CreateImageEditRequest {
77 pub image: Vec<u8>,
79 pub prompt: String,
81 pub mask: Option<Vec<u8>>,
83 pub model: Option<String>,
84 pub n: Option<i64>,
86 pub size: Option<ImageSize>,
87 pub quality: Option<String>,
89 pub response_format: Option<String>,
91}
92
93#[derive(Debug, Clone, Default)]
98pub struct CreateImageVariationRequest {
99 pub image: Vec<u8>,
101 pub model: Option<String>,
102 pub n: Option<i64>,
104 pub size: Option<ImageSize>,
105 pub response_format: Option<String>,
107}
108
109pub 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
127pub trait InferenceGatewayAPI {
129 fn list_models(&self) -> impl Future<Output = Result<ListModelsResponse, GatewayError>> + Send;
131
132 fn list_models_by_provider(
134 &self,
135 provider: Provider,
136 ) -> impl Future<Output = Result<ListModelsResponse, GatewayError>> + Send;
137
138 fn list_models_with_include(
144 &self,
145 provider: Option<Provider>,
146 include: &[&str],
147 ) -> impl Future<Output = Result<ListModelsResponse, GatewayError>> + Send;
148
149 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 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 fn create_message(
170 &self,
171 provider: Option<Provider>,
172 request: CreateMessagesRequest,
173 ) -> impl Future<Output = Result<MessagesResponse, GatewayError>> + Send;
174
175 fn create_message_stream(
178 &self,
179 provider: Option<Provider>,
180 request: CreateMessagesRequest,
181 ) -> impl Stream<Item = Result<SSEvents, GatewayError>> + Send;
182
183 fn list_tools(&self) -> impl Future<Output = Result<ListToolsResponse, GatewayError>> + Send;
185
186 fn generate_image(
191 &self,
192 provider: Provider,
193 request: CreateImageRequest,
194 ) -> impl Future<Output = Result<ImagesResponse, GatewayError>> + Send;
195
196 fn create_image_edit(
201 &self,
202 provider: Option<Provider>,
203 request: CreateImageEditRequest,
204 ) -> impl Future<Output = Result<ImagesResponse, GatewayError>> + Send;
205
206 fn create_image_variation(
211 &self,
212 provider: Option<Provider>,
213 request: CreateImageVariationRequest,
214 ) -> impl Future<Output = Result<ImagesResponse, GatewayError>> + Send;
215
216 fn create_speech(
225 &self,
226 provider: Option<Provider>,
227 request: CreateSpeechRequest,
228 ) -> impl Future<Output = Result<Vec<u8>, GatewayError>> + Send;
229
230 fn health_check(&self) -> impl Future<Output = Result<bool, GatewayError>> + Send;
232}
233
234impl InferenceGatewayClient {
235 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 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 pub fn with_tools(mut self, tools: Option<Vec<ChatCompletionTool>>) -> Self {
266 self.tools = tools;
267 self
268 }
269
270 pub fn with_token(mut self, token: impl Into<String>) -> Self {
272 self.token = Some(token.into());
273 self
274 }
275
276 pub fn with_max_tokens(mut self, max_tokens: Option<i64>) -> Self {
278 self.max_tokens = max_tokens;
279 self
280 }
281
282 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 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 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;