starweaver-model 0.3.0

Provider-neutral model protocol and wire adapters for Starweaver
Documentation
use std::collections::BTreeMap;

use async_trait::async_trait;
use futures_util::StreamExt;
use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
use serde_json::Value;

use crate::{ModelError, allow_real_model_requests};

use super::sse::{SseJsonParser, StreamSendError, push_sse_utf8_buffer, send_sse_parser_events};
use super::{HttpMethod, HttpRequest, HttpResponse, ModelEventStream, ModelHttpClient};
use crate::transport::{is_retryable_status, websocket};

/// Reqwest-backed HTTP client.
#[derive(Clone, Debug)]
pub struct ReqwestHttpClient {
    client: reqwest::Client,
}

impl ReqwestHttpClient {
    /// Create a reqwest-backed client with rustls TLS.
    ///
    /// # Errors
    ///
    /// Returns an error when reqwest client construction fails.
    pub fn new() -> Result<Self, ModelError> {
        let client = reqwest::Client::builder()
            .build()
            .map_err(|err| ModelError::Transport(err.to_string()))?;
        Ok(Self { client })
    }

    async fn send_request(&self, request: &HttpRequest) -> Result<reqwest::Response, ModelError> {
        if !allow_real_model_requests() {
            return Err(ModelError::RealModelRequestBlocked {
                url: request.url.clone(),
            });
        }

        let mut builder = match request.method {
            HttpMethod::Post => self.client.post(&request.url),
        }
        .headers(Self::header_map(&request.headers)?)
        .json(&request.body);

        if let Some(timeout) = request.timeout {
            builder = builder.timeout(timeout);
        }

        builder
            .send()
            .await
            .map_err(|err| ModelError::Transport(err.to_string()))
    }

    fn header_map(headers: &BTreeMap<String, String>) -> Result<HeaderMap, ModelError> {
        let mut map = HeaderMap::new();
        for (name, value) in headers {
            let name = HeaderName::from_bytes(name.as_bytes()).map_err(|err| {
                ModelError::Transport(format!("invalid header name {name}: {err}"))
            })?;
            let value = HeaderValue::from_str(value).map_err(|err| {
                ModelError::Transport(format!("invalid header value for {name}: {err}"))
            })?;
            map.insert(name, value);
        }
        Ok(map)
    }
}

#[async_trait]
impl ModelHttpClient for ReqwestHttpClient {
    async fn send(&self, request: HttpRequest) -> Result<HttpResponse, ModelError> {
        let cancellation_token = request.cancellation_token.clone();
        if cancellation_token.is_cancelled() {
            return Err(ModelError::Cancelled {
                reason: "model HTTP request cancellation requested".to_string(),
            });
        }
        let response = tokio::select! {
            biased;
            () = cancellation_token.cancelled() => {
                return Err(ModelError::Cancelled {
                    reason: "model HTTP request cancellation requested".to_string(),
                });
            }
            response = self.send_request(&request) => response?,
        };
        let status = response.status().as_u16();
        let headers = response_headers(&response);
        let body = tokio::select! {
            biased;
            () = cancellation_token.cancelled() => {
                return Err(ModelError::Cancelled {
                    reason: "model HTTP request cancellation requested".to_string(),
                });
            }
            body = response.json::<Value>() => {
                body.map_err(|err| ModelError::Transport(err.to_string()))?
            }
        };

        if (200..300).contains(&status) {
            Ok(HttpResponse {
                status,
                headers,
                body,
            })
        } else {
            Err(ModelError::ProviderStatus {
                status,
                body,
                retryable: is_retryable_status(status),
            })
        }
    }

    async fn send_event_stream_incremental(
        &self,
        request: HttpRequest,
    ) -> Result<ModelEventStream, ModelError> {
        let cancellation_token = request.cancellation_token.clone();
        if cancellation_token.is_cancelled() {
            return Err(ModelError::Cancelled {
                reason: "model event stream cancellation requested".to_string(),
            });
        }
        let response = tokio::select! {
            biased;
            () = cancellation_token.cancelled() => {
                return Err(ModelError::Cancelled {
                    reason: "model event stream cancellation requested".to_string(),
                });
            }
            response = self.send_request(&request) => response?,
        };
        let status = response.status().as_u16();
        if !(200..300).contains(&status) {
            let text = response
                .text()
                .await
                .map_err(|err| ModelError::Transport(err.to_string()))?;
            let body = serde_json::from_str(&text).unwrap_or(Value::String(text));
            return Err(ModelError::ProviderStatus {
                status,
                body,
                retryable: is_retryable_status(status),
            });
        }
        let (sender, receiver) = tokio::sync::mpsc::channel(32);
        let worker_cancellation_token = cancellation_token.clone();
        tokio::spawn(async move {
            let mut parser = SseJsonParser::default();
            let mut bytes = response.bytes_stream();
            let mut utf8_buffer = Vec::new();
            loop {
                let chunk = tokio::select! {
                    biased;
                    () = worker_cancellation_token.cancelled() => {
                        let _ = sender
                            .send(Err(ModelError::Cancelled {
                                reason: "model event stream cancellation requested".to_string(),
                            }))
                            .await;
                        return;
                    }
                    chunk = bytes.next() => chunk,
                };
                let Some(chunk) = chunk else {
                    break;
                };
                match chunk {
                    Ok(bytes) => {
                        utf8_buffer.extend_from_slice(&bytes);
                        match push_sse_utf8_buffer(&sender, &mut parser, &mut utf8_buffer).await {
                            Ok(()) => {}
                            Err(StreamSendError::Closed) => return,
                            Err(StreamSendError::InvalidUtf8(error)) => {
                                let _ = sender
                                    .send(Err(ModelError::ResponseParsing(format!(
                                        "invalid server-sent event UTF-8: {error}"
                                    ))))
                                    .await;
                                return;
                            }
                        }
                    }
                    Err(error) => {
                        let _ = sender
                            .send(Err(ModelError::Transport(error.to_string())))
                            .await;
                        return;
                    }
                }
            }
            if !utf8_buffer.is_empty() {
                match std::str::from_utf8(&utf8_buffer) {
                    Ok(text) => {
                        if !send_sse_parser_events(&sender, parser.push_str(text)).await {
                            return;
                        }
                    }
                    Err(error) => {
                        let _ = sender
                            .send(Err(ModelError::ResponseParsing(format!(
                                "invalid server-sent event UTF-8: {error}"
                            ))))
                            .await;
                        return;
                    }
                }
            }
            let _ = send_sse_parser_events(&sender, parser.finish()).await;
        });
        Ok(ModelEventStream::new_with_cancellation(
            receiver,
            cancellation_token,
        ))
    }

    async fn send_websocket_event_stream_incremental(
        &self,
        request: HttpRequest,
    ) -> Result<ModelEventStream, ModelError> {
        Box::pin(websocket::send_websocket_event_stream_incremental(request)).await
    }

    fn websocket_event_session(&self) -> Box<dyn super::ModelWebSocketEventSession + '_> {
        Box::new(websocket::ReusableWebSocketEventSession::default())
    }
}

fn response_headers(response: &reqwest::Response) -> BTreeMap<String, String> {
    response
        .headers()
        .iter()
        .filter_map(|(name, value)| {
            value
                .to_str()
                .ok()
                .map(|value| (name.as_str().to_string(), value.to_string()))
        })
        .collect()
}