agent-infra-sdk 0.1.1

Gateway-backed Rust SDK for Agent Infra APIs
Documentation
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};

use reqwest::Url;
use serde::Deserialize;

static TRACE_SEQUENCE: AtomicU64 = AtomicU64::new(1);

#[derive(Default, Deserialize)]
#[serde(rename_all = "camelCase")]
struct ErrorEnvelope {
    code: Option<String>,
    message: Option<String>,
    retryable: Option<bool>,
    request_id: Option<String>,
    error: Option<ErrorField>,
}

#[derive(Deserialize)]
#[serde(untagged)]
enum ErrorField {
    Problem(ErrorProblem),
    Message(String),
}

#[derive(Default, Deserialize)]
#[serde(rename_all = "camelCase")]
struct ErrorProblem {
    code: Option<String>,
    message: Option<String>,
    retryable: Option<bool>,
    request_id: Option<String>,
}

pub(crate) struct ParsedError {
    pub(crate) code: Option<String>,
    pub(crate) message: String,
    pub(crate) retryable: Option<bool>,
    pub(crate) request_id: Option<String>,
}

pub(crate) fn error_envelope(body: &[u8], max_message_bytes: usize) -> ParsedError {
    if let Ok(value) = serde_json::from_slice::<ErrorEnvelope>(body) {
        let (nested_code, nested_message, nested_retryable, nested_request_id) = match value.error {
            Some(ErrorField::Problem(problem)) => (
                problem.code,
                problem.message,
                problem.retryable,
                problem.request_id,
            ),
            Some(ErrorField::Message(message)) => (None, Some(message), None, None),
            None => (None, None, None, None),
        };
        let message = value.message.or(nested_message).unwrap_or_default();
        return ParsedError {
            code: value.code.or(nested_code),
            message: truncate(&message, max_message_bytes),
            retryable: value.retryable.or(nested_retryable),
            request_id: value.request_id.or(nested_request_id),
        };
    }
    ParsedError {
        code: None,
        message: truncate(&String::from_utf8_lossy(body), max_message_bytes),
        retryable: None,
        request_id: None,
    }
}

pub(crate) fn parse_retry_after(value: Option<&reqwest::header::HeaderValue>) -> Option<Duration> {
    parse_retry_after_at(value, SystemTime::now())
}

pub(crate) fn parse_retry_after_at(
    value: Option<&reqwest::header::HeaderValue>,
    now: SystemTime,
) -> Option<Duration> {
    let value = value?.to_str().ok()?.trim();
    if let Ok(seconds) = value.parse::<u64>() {
        return Some(Duration::from_secs(seconds));
    }
    httpdate::parse_http_date(value)
        .ok()
        .map(|retry_at| retry_at.duration_since(now).unwrap_or(Duration::ZERO))
}

fn truncate(value: &str, max_bytes: usize) -> String {
    if value.len() <= max_bytes {
        return value.to_string();
    }
    let mut end = max_bytes;
    while !value.is_char_boundary(end) {
        end -= 1;
    }
    format!("{}", &value[..end])
}

pub(crate) fn redacted_endpoint(value: &str) -> String {
    let Ok(mut url) = Url::parse(value) else {
        return "[REDACTED INVALID ENDPOINT]".into();
    };
    if !url.username().is_empty() {
        let _ = url.set_username("");
    }
    if url.password().is_some() {
        let _ = url.set_password(None);
    }
    url.set_query(None);
    url.set_fragment(None);
    url.to_string()
}

pub(crate) fn secure_transport(url: &Url, trusted_mesh_http: bool) -> bool {
    url.scheme() == "https"
        || (url.scheme() == "http"
            && (trusted_mesh_http
                || url.host_str().is_some_and(|host| {
                    host.eq_ignore_ascii_case("localhost")
                        || host
                            .parse::<std::net::IpAddr>()
                            .is_ok_and(|address| address.is_loopback())
                })))
}

pub(crate) fn new_traceparent() -> String {
    let sequence = TRACE_SEQUENCE.fetch_add(1, Ordering::Relaxed);
    let nanos = SystemTime::now()
        .duration_since(UNIX_EPOCH)
        .unwrap_or_default()
        .as_nanos();
    let trace_high = nanos as u64 ^ sequence.rotate_left(17);
    let trace_low = (nanos >> 64) as u64 ^ sequence.wrapping_mul(0x9e37_79b9_7f4a_7c15);
    let span = trace_high.rotate_left(23) ^ trace_low;
    format!("00-{trace_high:016x}{trace_low:016x}-{span:016x}-01")
}

pub(crate) fn elapsed_micros(started: Instant) -> u64 {
    u64::try_from(started.elapsed().as_micros()).unwrap_or(u64::MAX)
}

pub(crate) fn valid_traceparent(value: &str) -> bool {
    value.len() == 55
        && value.as_bytes().get(2) == Some(&b'-')
        && value.as_bytes().get(35) == Some(&b'-')
        && value.as_bytes().get(52) == Some(&b'-')
        && value
            .bytes()
            .enumerate()
            .all(|(index, byte)| matches!(index, 2 | 35 | 52) || byte.is_ascii_hexdigit())
        && &value[3..35] != "00000000000000000000000000000000"
        && &value[36..52] != "0000000000000000"
}