mentra-provider 0.6.0

Shared provider core for Mentra
Documentation
use async_trait::async_trait;
use serde_json::Value;
use std::collections::HashMap;
use std::sync::Arc;

pub(crate) mod model;
pub(crate) mod sse;
pub(crate) mod stream_model;

use crate::AuthScheme;
use crate::BuiltinProvider;
use crate::CompactionRequest;
use crate::CompactionResponse;
use crate::CredentialSource;
use crate::ModelCatalog;
use crate::ModelInfo;
use crate::ProviderCapabilities;
use crate::ProviderDefinition;
use crate::ProviderError;
use crate::ProviderEventStream;
use crate::ProviderSession;
use crate::ProviderSessionFactory;
use crate::RegisteredProvider;
use crate::Request;
use crate::StaticCredentialSource;
use crate::WireApi;

const DEFAULT_BASE_URL: &str = "https://api.anthropic.com";
const ANTHROPIC_VERSION: &str = "2023-06-01";

/// Returns the default Anthropic-compatible provider definition.
pub fn definition() -> ProviderDefinition {
    let mut definition = ProviderDefinition::new(BuiltinProvider::Anthropic);
    definition.descriptor.display_name = Some("Anthropic".to_string());
    definition.descriptor.description = Some("Anthropic Messages API provider".to_string());
    definition.wire_api = WireApi::AnthropicMessages;
    definition.auth_scheme = AuthScheme::Header {
        name: "x-api-key".to_string(),
    };
    definition.capabilities = ProviderCapabilities {
        supports_model_listing: true,
        supports_streaming: true,
        supports_websockets: false,
        supports_tool_calls: true,
        supports_images: true,
        supports_history_compaction: true,
        supports_memory_summarization: true,
        supports_deferred_tools: true,
        supports_hosted_tool_search: true,
        supports_hosted_web_search: false,
        supports_image_generation: false,
        supports_reasoning_effort: true,
        reports_reasoning_tokens: false,
        reports_thoughts_tokens: false,
        supports_structured_tool_results: false,
        supports_embeddings: false,
    };
    definition.base_url = Some(DEFAULT_BASE_URL.to_string());
    definition.headers = Some(HashMap::from([(
        "anthropic-version".to_string(),
        ANTHROPIC_VERSION.to_string(),
    )]));
    definition
}

pub struct AnthropicProvider<C = StaticCredentialSource> {
    client: reqwest::Client,
    credential_source: Arc<C>,
    definition: ProviderDefinition,
}

impl<C> Clone for AnthropicProvider<C> {
    fn clone(&self) -> Self {
        Self {
            client: self.client.clone(),
            credential_source: Arc::clone(&self.credential_source),
            definition: self.definition.clone(),
        }
    }
}

impl AnthropicProvider<StaticCredentialSource> {
    pub fn new(api_key: impl Into<String>) -> Self {
        Self::with_credential_source(StaticCredentialSource::new(api_key))
    }
}

impl<C> AnthropicProvider<C>
where
    C: CredentialSource + 'static,
{
    pub fn with_credential_source(credential_source: C) -> Self {
        Self::with_shared_credential_source(Arc::new(credential_source))
    }

    pub fn with_shared_credential_source(credential_source: Arc<C>) -> Self {
        Self::with_definition_and_shared_credential_source(definition(), credential_source)
    }

    pub fn with_definition_and_credential_source(
        definition: ProviderDefinition,
        credential_source: C,
    ) -> Self {
        Self::with_definition_and_shared_credential_source(definition, Arc::new(credential_source))
    }

    pub fn with_definition_and_shared_credential_source(
        definition: ProviderDefinition,
        credential_source: Arc<C>,
    ) -> Self {
        // The idle timeout, not a total one: a streamed turn can legitimately
        // run for minutes, but a gap between chunks means the provider stopped
        // talking. `read_timeout` resets on every successful read, so it bounds
        // the silence without bounding the turn. The resulting error is a
        // `Transport` error, which the runtime already treats as transient and
        // retries.
        let client = reqwest::Client::builder()
            .read_timeout(definition.stream_idle_timeout)
            .build()
            .expect("Failed to build client");

        Self {
            client,
            credential_source,
            definition,
        }
    }
}

#[async_trait]
impl<C> ModelCatalog for AnthropicProvider<C>
where
    C: CredentialSource + 'static,
{
    async fn list_models(&self) -> Result<Vec<ModelInfo>, ProviderError> {
        let mut models = Vec::new();
        let mut after_id = None;

        loop {
            let credentials = self.credential_source.credentials().await?;
            let request = self
                .client
                .get(
                    self.definition
                        .request_url_with_auth_for_path("v1/models", &credentials)?,
                )
                .headers(self.definition.build_headers(&credentials)?)
                .query(&[
                    ("limit", "1000"),
                    ("after_id", after_id.as_deref().unwrap_or("")),
                ]);

            let response = request.send().await.map_err(ProviderError::Transport)?;

            if !response.status().is_success() {
                return Err(ProviderError::from_http_response(response).await);
            }

            let page = response
                .json::<model::AnthropicModelsPage>()
                .await
                .map_err(ProviderError::Decode)?;

            after_id = page.last_id.clone();
            models.extend(page.data.into_iter().map(|model| model.into()));

            if !page.has_more {
                break;
            }
        }

        Ok(models)
    }
}

#[async_trait]
impl<C> ProviderSessionFactory for AnthropicProvider<C>
where
    C: CredentialSource + 'static,
{
    async fn create_session(&self) -> Result<Box<dyn ProviderSession>, ProviderError> {
        Ok(Box::new((*self).clone()))
    }
}

#[async_trait]
impl<C> ProviderSession for AnthropicProvider<C>
where
    C: CredentialSource + 'static,
{
    async fn stream(&self, request: Request<'_>) -> Result<ProviderEventStream, ProviderError> {
        let requested_model = request.model.to_string();
        let provider = self.definition.provider_id().clone();
        let response = self.send_message(request, true).await?;
        Ok(sse::spawn_event_stream(response, provider, requested_model))
    }

    async fn compact(
        &self,
        request: CompactionRequest<'_>,
    ) -> Result<CompactionResponse, ProviderError> {
        let request = request.into_model_request()?;
        let response = ProviderSession::send(self, request).await?;
        Ok(response.into_compaction_response())
    }

    async fn summarize_memories(
        &self,
        request: crate::MemorySummarizeRequest<'_>,
    ) -> Result<crate::MemorySummarizeResponse, ProviderError> {
        let request = request.into_model_request()?;
        let response = ProviderSession::send(self, request).await?;
        response.into_memory_summarize_response()
    }
}

#[async_trait]
impl<C> RegisteredProvider for AnthropicProvider<C>
where
    C: CredentialSource + 'static,
{
    fn definition(&self) -> ProviderDefinition {
        self.definition.clone()
    }
}

impl<C> AnthropicProvider<C>
where
    C: CredentialSource + 'static,
{
    async fn send_message(
        &self,
        request: Request<'_>,
        stream: bool,
    ) -> Result<reqwest::Response, ProviderError> {
        let session = request.provider_request_options.session.clone();
        let request = model::AnthropicRequest::try_from_with_provider(
            request,
            self.definition.provider_id(),
        )?;
        let mut body = serde_json::to_value(request).map_err(ProviderError::Serialize)?;
        if stream {
            body["stream"] = Value::Bool(true);
        }
        let credentials = self.credential_source.credentials().await?;
        let response = self
            .client
            .post(
                self.definition
                    .request_url_with_auth_for_path("v1/messages", &credentials)?,
            )
            .headers(self.definition.build_headers_for_session(
                &credentials,
                Some(&session),
                None,
            )?)
            .json(&body)
            .send()
            .await
            .map_err(ProviderError::Transport)?;

        if !response.status().is_success() {
            return Err(ProviderError::from_http_response(response).await);
        }

        Ok(response)
    }
}

#[cfg(test)]
mod tests {
    use std::borrow::Cow;
    use std::collections::BTreeMap;
    use std::io::{Read, Write};
    use std::net::TcpListener;
    use std::sync::mpsc;
    use std::thread;
    use std::time::{Duration, Instant};

    use super::*;
    use crate::{
        ContentBlock, Message, ProviderRequestOptions, StaticCredentialSource, ToolChoice,
    };

    #[test]
    fn definition_advertises_history_compaction_support() {
        assert!(definition().capabilities.supports_history_compaction);
    }

    /// Answers one request with SSE headers and a first event, then holds the
    /// socket open and sends nothing further — a provider that accepted the
    /// turn and stopped talking. `_shutdown` keeps the connection alive until
    /// the test drops it.
    fn spawn_stalling_sse_server() -> (String, mpsc::Sender<()>) {
        let listener = TcpListener::bind("127.0.0.1:0").expect("bind test server");
        let addr = listener.local_addr().expect("read listener addr");
        let (shutdown, wait_for_shutdown) = mpsc::channel::<()>();

        thread::spawn(move || {
            let (mut stream, _) = listener.accept().expect("accept request");
            let mut temp = [0_u8; 1024];
            let _ = stream.read(&mut temp).expect("read request");
            stream
                .write_all(
                    b"HTTP/1.1 200 OK\r\n\
                      content-type: text/event-stream\r\n\
                      transfer-encoding: chunked\r\n\r\n\
                      2b\r\n\
                      event: ping\ndata: {\"type\":\"ping\"}\n\n\r\n",
                )
                .expect("write stream head");
            stream.flush().expect("flush stream head");
            // Then say nothing at all.
            let _ = wait_for_shutdown.recv();
        });

        (format!("http://{addr}"), shutdown)
    }

    #[tokio::test]
    async fn a_stalled_sse_stream_fails_instead_of_hanging() {
        // A provider that accepts a turn and then stops sending held the stream
        // open until the caller's own deadline, if it had one. The definition's
        // `stream_idle_timeout` now bounds the gap between chunks.
        let (base_url, _shutdown) = spawn_stalling_sse_server();
        let mut definition = definition();
        definition.base_url = Some(base_url);
        definition.stream_idle_timeout = Duration::from_millis(250);

        let provider = AnthropicProvider::with_definition_and_credential_source(
            definition,
            StaticCredentialSource::new("test-key"),
        );

        let started = Instant::now();
        let mut stream = ProviderSession::stream(
            &provider,
            Request {
                model: Cow::Borrowed("claude-test"),
                system: None,
                messages: Cow::Owned(vec![Message::user(ContentBlock::text("hi"))]),
                tools: Cow::Owned(vec![]),
                tool_choice: Some(ToolChoice::Auto),
                temperature: None,
                max_output_tokens: Some(16),
                metadata: Cow::Owned(BTreeMap::new()),
                provider_request_options: ProviderRequestOptions::default(),
            },
        )
        .await
        .expect("headers arrive before the stall");

        let drain = async {
            while let Some(event) = stream.recv().await {
                if let Err(error) = event {
                    return Some(error);
                }
            }
            None
        };
        // The bound is the assertion: without an idle timeout of its own the
        // stream waits on a provider that will never speak again, and the only
        // thing that ends the wait is a deadline someone else imposed.
        let failure = tokio::time::timeout(Duration::from_secs(10), drain)
            .await
            .expect("the stream gave up on its own rather than hanging");

        let error = failure.expect("the stalled stream ends in an error, not silence");
        assert!(
            matches!(error, ProviderError::Transport(_)),
            "a stall is a transport failure so the run retries it: {error}"
        );
        assert!(
            started.elapsed() < Duration::from_secs(10),
            "the stream gave up on its own, not on a test deadline"
        );
    }
}