Luigit
repositories / smith

smith

There are many coding harnesses - but this one is fast

owned by admin

smith-ai/src/transport.rs

Raw
//! Vendor HTTP transport: request bodies, posting, and SSE decoding
//! (`SMH-SPEC-SPEC0001`, Network transport).
//!
//! Body builders are pure functions over [`smith::provider::ProviderRequest`] so their vendor
//! shapes are unit-tested without network. The transport posts and feeds the
//! matching vendor decoder on a dedicated thread, surfacing normalized
//! events through a channel-backed stream.

use crate::anthropic::AnthropicDecoder;
use crate::conformance::StreamDecoder;
use crate::gemini::GeminiDecoder;
use crate::models::ProviderKind;
use crate::openai::OpenAiDecoder;
use futures::StreamExt;
use smith::error::{ProviderFault, Result, SmithError};
use smith::http::{
    HttpBodyEnd, HttpBodyItem, HttpExchange, HttpExecutorRef, HttpHeader, HttpMethod,
};
use smith::message::{ContentBlock, Role};
use smith::provider::{ProviderRequest, RecordSink, StreamFn};
use smith::tool::CancelHandle;
use std::sync::Arc;

/// Where provider effect records go until the agent drains them into the
/// session; a closed drain drops records.
pub type ProviderRecordSpool = RecordSink;

fn spool_record(spool: &ProviderRecordSpool, record: serde_json::Value) {
    let _ = spool.send(record);
}

/// Vendor endpoint description for one model.
#[derive(Clone)]
pub struct VendorEndpoint {
    /// Adapter family.
    pub kind: ProviderKind,
    /// Base URL without trailing slash.
    pub base_url: String,
    /// Model id for the request body.
    pub model: String,
    /// Bearer or key credential value.
    pub api_key: String,
    /// The host HTTP engine; provider traffic never posts directly.
    pub executor: HttpExecutorRef,
    /// Redacted effect records, drained by the agent into the session.
    pub records: ProviderRecordSpool,
}

/// Build the vendor request body for one provider request.
///
/// # Errors
///
/// Returns a provider protocol fault when the conversation contains content
/// the vendor body cannot represent.
pub fn build_body(
    endpoint: &VendorEndpoint,
    request: &ProviderRequest,
) -> Result<serde_json::Value> {
    match endpoint.kind {
        ProviderKind::OpenAiCompatible | ProviderKind::Plugin => Ok(openai_body(endpoint, request)),
        ProviderKind::Anthropic => Ok(anthropic_body(endpoint, request)),
        ProviderKind::Google => Ok(google_body(endpoint, request)),
    }
}

fn openai_body(endpoint: &VendorEndpoint, request: &ProviderRequest) -> serde_json::Value {
    let mut messages = Vec::new();
    for message in &request.messages {
        match message.role {
            Role::System => messages.push(serde_json::json!({
                "role": "system",
                "content": flatten_text(message),
            })),
            Role::User => messages.push(serde_json::json!({
                "role": "user",
                "content": flatten_text(message),
            })),
            Role::Assistant => messages.push(serde_json::json!({
                "role": "assistant",
                "content": flatten_text(message),
            })),
            Role::Tool => {
                // Each tool result block becomes its own tool message.
                for block in &message.blocks {
                    if let ContentBlock::ToolResult {
                        call_id, output, ..
                    } = block
                    {
                        messages.push(serde_json::json!({
                            "role": "tool",
                            "tool_call_id": call_id.to_string(),
                            "content": output,
                        }));
                    }
                }
            }
        }
    }
    // OpenAI expects assistant tool calls on the assistant message; the
    // normalized conversation records them as blocks, mapped here.
    append_assistant_tool_calls(&request.messages, &mut messages);
    let mut body = serde_json::json!({
        "model": endpoint.model,
        "messages": messages,
        "stream": true,
        "stream_options": {"include_usage": true},
    });
    if !request.tools.is_empty() {
        body["tools"] = serde_json::json!(
            request
                .tools
                .iter()
                .map(|tool| serde_json::json!({
                    "type": "function",
                    "function": {
                        "name": tool.name,
                        "description": tool.description,
                        "parameters": tool.input_schema,
                    }
                }))
                .collect::<Vec<_>>()
        );
    }
    if let Some(max) = request.params.max_output_tokens {
        body["max_tokens"] = serde_json::json!(max);
    }
    body
}

fn append_assistant_tool_calls(
    messages: &[smith::message::Message],
    out: &mut [serde_json::Value],
) {
    for message in messages {
        if message.role != Role::Assistant {
            continue;
        }
        let calls: Vec<serde_json::Value> = message
            .blocks
            .iter()
            .filter_map(|block| {
                if let ContentBlock::ToolUse {
                    name,
                    input,
                    call_id,
                    ..
                } = block
                {
                    Some(serde_json::json!({
                        "id": call_id.to_string(),
                        "type": "function",
                        "function": {"name": name, "arguments": input.to_string()},
                    }))
                } else {
                    None
                }
            })
            .collect();
        if calls.is_empty() {
            continue;
        }
        if let Some(last) = out.iter_mut().rev().find(|entry| {
            entry.get("role").and_then(serde_json::Value::as_str) == Some("assistant")
        }) {
            last["tool_calls"] = serde_json::json!(calls);
        }
    }
}

fn anthropic_body(endpoint: &VendorEndpoint, request: &ProviderRequest) -> serde_json::Value {
    let mut system = String::new();
    let mut messages = Vec::new();
    for message in &request.messages {
        match message.role {
            Role::System => {
                if !system.is_empty() {
                    system.push('\n');
                }
                system.push_str(&flatten_text(message));
            }
            Role::User => messages.push(serde_json::json!({
                "role": "user",
                "content": [{"type": "text", "text": flatten_text(message)}],
            })),
            Role::Assistant => {
                let mut content = Vec::new();
                for block in &message.blocks {
                    match block {
                        ContentBlock::Text(text) => content.push(serde_json::json!({
                            "type": "text", "text": text,
                        })),
                        ContentBlock::ToolUse {
                            name,
                            input,
                            call_id,
                            ..
                        } => content.push(serde_json::json!({
                            "type": "tool_use",
                            "id": call_id.to_string(),
                            "name": name,
                            "input": input,
                        })),
                        _ => {}
                    }
                }
                messages.push(serde_json::json!({"role": "assistant", "content": content}));
            }
            Role::Tool => {
                let mut content = Vec::new();
                for block in &message.blocks {
                    if let ContentBlock::ToolResult {
                        call_id, output, ..
                    } = block
                    {
                        content.push(serde_json::json!({
                            "type": "tool_result",
                            "tool_use_id": call_id.to_string(),
                            "content": output,
                        }));
                    }
                }
                messages.push(serde_json::json!({"role": "user", "content": content}));
            }
        }
    }
    let mut body = serde_json::json!({
        "model": endpoint.model,
        "max_tokens": request.params.max_output_tokens.unwrap_or(4096),
        "messages": messages,
        "stream": true,
    });
    if !system.is_empty() {
        body["system"] = serde_json::json!(system);
    }
    if !request.tools.is_empty() {
        body["tools"] = serde_json::json!(
            request
                .tools
                .iter()
                .map(|tool| serde_json::json!({
                    "name": tool.name,
                    "description": tool.description,
                    "input_schema": tool.input_schema,
                }))
                .collect::<Vec<_>>()
        );
    }
    body
}

fn google_body(_endpoint: &VendorEndpoint, request: &ProviderRequest) -> serde_json::Value {
    let system: Vec<String> = request
        .messages
        .iter()
        .filter(|message| message.role == Role::System)
        .map(flatten_text)
        .collect();
    let mut contents = Vec::new();
    for message in &request.messages {
        let role = match message.role {
            Role::User | Role::Tool => "user",
            Role::Assistant => "model",
            Role::System => continue,
        };
        let mut parts = Vec::new();
        for block in &message.blocks {
            match block {
                ContentBlock::Text(text) => parts.push(serde_json::json!({"text": text})),
                ContentBlock::ToolUse { name, input, .. } => parts.push(serde_json::json!({
                    "functionCall": {"name": name, "args": input},
                })),
                ContentBlock::ToolResult { output, .. } => {
                    parts.push(serde_json::json!({
                        "functionResponse": {"name": "tool", "response": {"output": output}},
                    }));
                }
                ContentBlock::Thinking(_) => {}
            }
        }
        contents.push(serde_json::json!({"role": role, "parts": parts}));
    }
    let mut body = serde_json::json!({
        "contents": contents,
        "generationConfig": {
            "maxOutputTokens": request.params.max_output_tokens.unwrap_or(4096),
        },
    });
    if !system.is_empty() {
        body["systemInstruction"] = serde_json::json!({"parts": [{"text": system.join("\n")}]});
    }
    if !request.tools.is_empty() {
        body["tools"] = serde_json::json!([{
            "functionDeclarations": request
                .tools
                .iter()
                .map(|tool| serde_json::json!({
                    "name": tool.name,
                    "description": tool.description,
                    "parameters": tool.input_schema,
                }))
                .collect::<Vec<_>>(),
        }]);
    }
    body
}

fn flatten_text(message: &smith::message::Message) -> String {
    let mut out = String::new();
    for block in &message.blocks {
        if let ContentBlock::Text(text) = block {
            if !out.is_empty() {
                out.push('\n');
            }
            out.push_str(text);
        }
    }
    out
}

/// The request URL for one endpoint.
#[must_use]
pub fn endpoint_url(endpoint: &VendorEndpoint) -> String {
    match endpoint.kind {
        ProviderKind::OpenAiCompatible | ProviderKind::Plugin => {
            format!("{}/chat/completions", endpoint.base_url)
        }
        ProviderKind::Anthropic => format!("{}/v1/messages", endpoint.base_url),
        ProviderKind::Google => format!(
            "{}/v1beta/models/{}:streamGenerateContent?alt=sse",
            endpoint.base_url, endpoint.model
        ),
    }
}

/// Headers the vendor requires.
#[must_use]
pub fn endpoint_headers(endpoint: &VendorEndpoint) -> Vec<(String, String)> {
    match endpoint.kind {
        ProviderKind::OpenAiCompatible | ProviderKind::Plugin => {
            vec![(
                "Authorization".to_string(),
                format!("Bearer {}", endpoint.api_key),
            )]
        }
        ProviderKind::Anthropic => vec![
            ("x-api-key".to_string(), endpoint.api_key.clone()),
            ("anthropic-version".to_string(), "2023-06-01".to_string()),
        ],
        ProviderKind::Google => vec![("x-goog-api-key".to_string(), endpoint.api_key.clone())],
    }
}

fn make_decoder(kind: ProviderKind) -> Box<dyn StreamDecoder + Send> {
    match kind {
        ProviderKind::OpenAiCompatible | ProviderKind::Plugin => Box::new(OpenAiDecoder::new()),
        ProviderKind::Anthropic => Box::new(AnthropicDecoder::new()),
        ProviderKind::Google => Box::new(GeminiDecoder::new()),
    }
}

/// Build a [`StreamFn`] that posts to one vendor endpoint.
///
/// The transport is deliberately synchronous-threaded: the agent loop polls
/// streams synchronously, so a reader thread feeds decoded events through a
/// channel.
///
/// # Errors
///
/// The returned stream's first item is `Err` when the request cannot be
/// built, credentials are missing, or the HTTP call fails.
#[must_use]
pub fn transport(endpoint: VendorEndpoint) -> StreamFn {
    Arc::new(move |request: ProviderRequest, cancel: CancelHandle| {
        let endpoint = endpoint.clone();
        let (sender, receiver) = std::sync::mpsc::channel::<smith::provider::StreamItem>();
        std::thread::spawn(move || {
            let result = post_and_decode(&endpoint, &request, &cancel, &sender);
            if let Err(err) = result
                && sender.send(Err(err)).is_err()
            {
                // Receiver dropped: the turn was abandoned.
            }
        });
        futures::stream::unfold(receiver, |receiver| async move {
            receiver.recv().map_or(None, |item| Some((item, receiver)))
        })
        .boxed()
    })
}

fn post_and_decode(
    endpoint: &VendorEndpoint,
    request: &ProviderRequest,
    cancel: &CancelHandle,
    sender: &std::sync::mpsc::Sender<smith::provider::StreamItem>,
) -> Result<()> {
    let body = build_body(endpoint, request)?;
    let headers = endpoint_headers(endpoint)
        .into_iter()
        .map(|(name, value)| HttpHeader { name, value })
        .collect::<Vec<_>>();
    let exchange = HttpExchange {
        method: HttpMethod::Post,
        url: endpoint_url(endpoint),
        headers: headers.clone(),
        body: serde_json::to_vec(&body).unwrap_or_default(),
    };
    spool_record(
        &endpoint.records,
        serde_json::json!({
            "phase": "intent",
            "method": exchange.method.as_str(),
            "url": exchange.url,
            "headers": smith::http::redact_headers(&headers),
            "request_bytes": exchange.body.len(),
        }),
    );
    let response = endpoint.executor.execute(exchange, cancel.clone())?;
    if !(200..300).contains(&response.head.status) {
        let mut body = Vec::new();
        for item in response.body {
            match item {
                HttpBodyItem::Chunk(chunk) => body.extend_from_slice(&chunk.bytes),
                HttpBodyItem::End(_) => break,
            }
        }
        let message = String::from_utf8_lossy(&body).into_owned();
        spool_record(
            &endpoint.records,
            serde_json::json!({
                "phase": "outcome",
                "status": response.head.status,
                "response_bytes": body.len(),
                "outcome": "failed",
            }),
        );
        return Err(SmithError::Provider {
            fault: status_fault(response.head.status, message),
        });
    }
    let mut decoder = make_decoder(endpoint.kind);
    let mut received: u64 = 0;
    let outcome = |received: u64, outcome: &str| {
        serde_json::json!({
            "phase": "outcome",
            "status": response.head.status,
            "response_bytes": received,
            "outcome": outcome,
        })
    };
    loop {
        if cancel.is_cancelled() {
            spool_record(&endpoint.records, outcome(received, "cancelled"));
            return Err(SmithError::Cancelled);
        }
        let error = match response.body.recv() {
            Ok(HttpBodyItem::Chunk(chunk)) => {
                received += u64::try_from(chunk.bytes.len()).unwrap_or(u64::MAX);
                for event in decoder.push(&chunk.bytes)? {
                    if sender.send(Ok(event)).is_err() {
                        return Ok(());
                    }
                }
                continue;
            }
            Ok(HttpBodyItem::End(HttpBodyEnd::Complete)) => break,
            Ok(HttpBodyItem::End(HttpBodyEnd::Cancelled)) => SmithError::Cancelled,
            Ok(HttpBodyItem::End(HttpBodyEnd::Truncated { limit })) => SmithError::Provider {
                fault: ProviderFault::Transient {
                    message: format!("response body exceeded {limit} bytes"),
                },
            },
            Ok(HttpBodyItem::End(HttpBodyEnd::Failed { message })) => SmithError::Provider {
                fault: ProviderFault::Transient { message },
            },
            Err(_) => SmithError::Provider {
                fault: ProviderFault::Transient {
                    message: "body channel closed without an end".to_string(),
                },
            },
        };
        spool_record(&endpoint.records, outcome(received, "failed"));
        return Err(error);
    }
    for event in decoder.finish()? {
        if sender.send(Ok(event)).is_err() {
            return Ok(());
        }
    }
    spool_record(&endpoint.records, outcome(received, "completed"));
    Ok(())
}

/// The provider fault a non-2xx status maps to.
fn status_fault(status: u16, message: String) -> ProviderFault {
    match status {
        429 => ProviderFault::RateLimit {
            retry_after_ms: None,
        },
        401 | 403 => ProviderFault::Authentication { message },
        400 | 404 => ProviderFault::Invalid {
            field: None,
            message,
        },
        503 => ProviderFault::Overloaded { message },
        _ => ProviderFault::Transient { message },
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use smith::id::ToolCallId;
    use smith::message::{Message, Role};
    use smith::stream::StreamEvent;
    use smith::tool::ToolMetadata;

    fn endpoint(kind: ProviderKind, model: &str) -> VendorEndpoint {
        VendorEndpoint {
            kind,
            base_url: "https://example.test".to_string(),
            model: model.to_string(),
            api_key: "k".to_string(),
            executor: std::sync::Arc::new(RefusesExecutor),
            records: smith::provider::record_channel().0,
        }
    }

    /// Body-builder tests never execute exchanges.
    struct RefusesExecutor;

    impl smith::http::HttpExecutor for RefusesExecutor {
        fn execute(
            &self,
            _exchange: smith::http::HttpExchange,
            _cancel: smith::tool::CancelHandle,
        ) -> smith::error::Result<smith::http::HttpResponse> {
            Err(SmithError::Provider {
                fault: ProviderFault::Invalid {
                    field: None,
                    message: "executor must not run in body tests".to_string(),
                },
            })
        }
    }

    #[test]
    fn openai_body_carries_messages_tools_and_usage_flag() {
        let call_id = ToolCallId::new();
        let mut assistant = Message::new(Role::Assistant);
        assistant.add_block(ContentBlock::text("Writing."));
        assistant.add_block(ContentBlock::tool_use(
            "write",
            serde_json::json!({"path": "f"}),
            call_id,
        ));
        let mut result = Message::new(Role::Tool);
        result.add_block(ContentBlock::tool_result(call_id, true, "written"));
        let request = ProviderRequest::new(
            "gpt",
            vec![
                Message::with_text(Role::System, "be brief"),
                Message::with_text(Role::User, "hi"),
                assistant,
                result,
            ],
        )
        .with_tools(vec![ToolMetadata::new("read", "Read")]);
        let body = build_body(&endpoint(ProviderKind::OpenAiCompatible, "gpt"), &request).unwrap();
        assert_eq!(body["model"], "gpt");
        assert_eq!(body["stream"], true);
        assert_eq!(body["stream_options"]["include_usage"], true);
        assert_eq!(body["messages"][0]["role"], "system");
        assert_eq!(body["messages"][3]["role"], "tool");
        assert!(body["messages"][2]["tool_calls"].is_array());
        assert_eq!(body["tools"][0]["function"]["name"], "read");
    }

    #[test]
    fn anthropic_body_hoists_system_and_pairs_tool_blocks() {
        let call_id = ToolCallId::new();
        let mut assistant = Message::new(Role::Assistant);
        assistant.add_block(ContentBlock::text("Reading."));
        assistant.add_block(ContentBlock::tool_use(
            "read",
            serde_json::json!({"path": "f"}),
            call_id,
        ));
        let mut result = Message::new(Role::Tool);
        result.add_block(ContentBlock::tool_result(call_id, true, "ok"));
        let request = ProviderRequest::new(
            "claude",
            vec![
                Message::with_text(Role::System, "sys"),
                Message::with_text(Role::User, "hi"),
                assistant,
                result,
            ],
        )
        .with_tools(vec![ToolMetadata::new("read", "Read")]);
        let body = build_body(&endpoint(ProviderKind::Anthropic, "claude"), &request).unwrap();
        assert_eq!(body["system"], "sys");
        assert_eq!(body["messages"].as_array().map(Vec::len), Some(3));
        assert_eq!(
            body["messages"][1]["content"][0],
            serde_json::json!({"type": "text", "text": "Reading."})
        );
        assert_eq!(body["messages"][1]["content"][1]["type"], "tool_use");
        assert_eq!(body["messages"][2]["content"][0]["type"], "tool_result");
        assert_eq!(body["tools"][0]["name"], "read");
    }

    /// Answers every exchange with `status` and `chunks`, then a complete end.
    struct Canned {
        status: u16,
        chunks: Vec<&'static str>,
    }

    impl smith::http::HttpExecutor for Canned {
        fn execute(
            &self,
            _exchange: HttpExchange,
            _cancel: CancelHandle,
        ) -> Result<smith::http::HttpResponse> {
            let (sender, body) = std::sync::mpsc::channel();
            for chunk in &self.chunks {
                let bytes = chunk.as_bytes().to_vec();
                let _ = sender.send(HttpBodyItem::Chunk(smith::http::HttpChunk { bytes }));
            }
            let _ = sender.send(HttpBodyItem::End(HttpBodyEnd::Complete));
            Ok(smith::http::HttpResponse {
                head: smith::http::HttpHead {
                    status: self.status,
                },
                body,
            })
        }
    }

    /// Post one request against `canned`; the result, streamed items, and
    /// spooled records.
    fn post(
        canned: Canned,
    ) -> (
        Result<()>,
        Vec<smith::provider::StreamItem>,
        Vec<serde_json::Value>,
    ) {
        let (records, drain) = smith::provider::record_channel();
        let endpoint = VendorEndpoint {
            executor: Arc::new(canned),
            records,
            ..endpoint(ProviderKind::OpenAiCompatible, "gpt")
        };
        let request = ProviderRequest::new("gpt", vec![Message::with_text(Role::User, "hi")]);
        let (sender, items) = std::sync::mpsc::channel();
        let result = post_and_decode(&endpoint, &request, &CancelHandle::new(), &sender);
        drop(sender);
        (
            result,
            items.try_iter().collect(),
            drain.try_iter().collect(),
        )
    }

    #[test]
    fn non_success_statuses_map_to_provider_faults_carrying_the_body() {
        use crate::conformance::faults::*;
        let cases = [
            (429, rate_limit()),
            (401, authentication()),
            (403, authentication()),
            (400, invalid()),
            (404, invalid()),
            (503, overloaded()),
            (500, transient()),
        ];
        for (status, fault) in cases {
            let (result, items, _) = post(Canned {
                status,
                chunks: vec!["m"],
            });
            assert_eq!(
                result.err(),
                Some(SmithError::Provider { fault }),
                "{status}"
            );
            assert!(items.is_empty(), "{status}");
        }
    }

    #[test]
    fn completed_exchange_streams_events_and_records_received_bytes() {
        let chunks = vec![
            "data: {\"choices\":[{\"delta\":{\"content\":\"Hi\"}}]}\n\n",
            "data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}]}\n\ndata: [DONE]\n\n",
        ];
        let received = chunks.iter().map(|chunk| chunk.len()).sum::<usize>();
        let (result, items, records) = post(Canned {
            status: 200,
            chunks,
        });
        assert_eq!(result.ok(), Some(()));
        assert!(matches!(
            items.as_slice(),
            [
                Ok(StreamEvent::TextDelta { .. }),
                Ok(StreamEvent::Stop { .. })
            ]
        ));
        assert_eq!(
            records.last(),
            Some(&serde_json::json!({
                "phase": "outcome",
                "status": 200,
                "response_bytes": received,
                "outcome": "completed",
            }))
        );
    }

    #[test]
    fn google_body_maps_roles_and_declares_functions() {
        let request = ProviderRequest::new(
            "gemini",
            vec![
                Message::with_text(Role::System, "sys"),
                Message::with_text(Role::User, "hi"),
            ],
        )
        .with_tools(vec![ToolMetadata::new("read", "Read")]);
        let body = build_body(&endpoint(ProviderKind::Google, "gemini"), &request).unwrap();
        assert_eq!(body["systemInstruction"]["parts"][0]["text"], "sys");
        assert_eq!(body["contents"][0]["role"], "user");
        assert_eq!(body["tools"][0]["functionDeclarations"][0]["name"], "read");
    }

    #[test]
    fn urls_and_headers_match_vendor_conventions() {
        let openai = endpoint(ProviderKind::OpenAiCompatible, "gpt");
        assert!(endpoint_url(&openai).ends_with("/chat/completions"));
        let anthropic = endpoint(ProviderKind::Anthropic, "claude");
        assert!(endpoint_url(&anthropic).ends_with("/v1/messages"));
        let google = endpoint(ProviderKind::Google, "gemini");
        assert!(endpoint_url(&google).contains("alt=sse"));
        let headers = endpoint_headers(&anthropic);
        assert!(headers.iter().any(|(name, _)| name == "x-api-key"));
    }

    #[test]
    fn stream_events_flow_through_the_shape_they_came_in() {
        // Body builders stay pure; the unfold transport is covered by the
        // decoders, so this pins the vendor kind to decoder mapping.
        let mut decoder = make_decoder(ProviderKind::OpenAiCompatible);
        let mut events = decoder
            .push(b"data: {\"choices\":[{\"delta\":{\"content\":\"x\"}}]}\n\n")
            .unwrap();
        assert!(matches!(events.pop(), Some(StreamEvent::TextDelta { .. })));
    }
}