agent-infra-sdk 0.1.0

Gateway-backed Rust SDK for Agent Infra APIs
Documentation
use super::*;
use std::sync::Mutex;

#[test]
fn retry_after_supports_imf_fixdate_with_a_fixed_clock() {
    let retry_at = reqwest::header::HeaderValue::from_static("Sun, 06 Nov 1994 08:49:37 GMT");
    let before = httpdate::parse_http_date("Sun, 06 Nov 1994 08:49:30 GMT").unwrap();
    let after = httpdate::parse_http_date("Sun, 06 Nov 1994 08:50:00 GMT").unwrap();
    assert_eq!(
        crate::transport_support::parse_retry_after_at(Some(&retry_at), before),
        Some(Duration::from_secs(7))
    );
    assert_eq!(
        crate::transport_support::parse_retry_after_at(Some(&retry_at), after),
        Some(Duration::ZERO),
        "an expired HTTP-date is immediately retryable"
    );
}

#[tokio::test]
async fn retry_after_http_date_cannot_extend_the_request_deadline() {
    let server = wiremock::MockServer::start().await;
    wiremock::Mock::given(wiremock::matchers::method("GET"))
        .and(wiremock::matchers::path("/busy"))
        .respond_with(
            wiremock::ResponseTemplate::new(503)
                .insert_header("retry-after", "Sun, 06 Nov 2094 08:49:37 GMT")
                .set_body_json(serde_json::json!({"error":{"code":"BUSY","retryable":true}})),
        )
        .mount(&server)
        .await;
    let transport = HttpTransport::new_with_options(
        Client::new(),
        "context",
        ServiceEndpoint::new(server.uri()),
        ClientOptions {
            request_timeout: Duration::from_millis(50),
            retry: RetryPolicy {
                max_attempts: 2,
                ..RetryPolicy::default()
            },
            ..ClientOptions::default()
        },
    );
    let error = transport
        .get_json::<serde_json::Value>("/busy")
        .await
        .unwrap_err();
    assert!(matches!(error, InfraClientError::DeadlineExceeded { .. }));
    assert_eq!(server.received_requests().await.unwrap().len(), 1);
}

#[test]
fn parses_standard_nested_error_envelope() {
    let parsed = error_envelope(
        br#"{"error":{"code":"VERSION_CONFLICT","message":"stale","retryable":false,"requestId":"req-1"}}"#,
    );
    assert_eq!(parsed.code.as_deref(), Some("VERSION_CONFLICT"));
    assert_eq!(parsed.message, "stale");
    assert_eq!(parsed.retryable, Some(false));
    assert_eq!(parsed.request_id.as_deref(), Some("req-1"));
}

#[test]
fn endpoint_and_credential_debug_never_expose_secret_material() {
    let endpoint = ServiceEndpoint::new(
        "https://user:sensitive-password-value@example.test/base?access_token=sensitive-query-value#sensitive-fragment-value",
    );
    let debug = format!("{endpoint:?}");
    for secret in [
        "user",
        "sensitive-password-value",
        "sensitive-query-value",
        "sensitive-fragment-value",
    ] {
        assert!(!debug.contains(secret));
    }
    assert!(debug.contains("https://example.test/base"));
    let credential_error = CredentialError("provider leaked token-123".into());
    assert!(!format!("{credential_error:?}").contains("token-123"));
    assert!(!credential_error.to_string().contains("token-123"));
    let transport = HttpTransport::new_with_options(
        Client::new(),
        "context",
        endpoint,
        ClientOptions::default(),
    );
    let transport_debug = format!("{transport:?}");
    for secret in [
        "user",
        "sensitive-password-value",
        "sensitive-query-value",
        "sensitive-fragment-value",
    ] {
        assert!(!transport_debug.contains(secret));
    }
    let error = transport.url("/test").unwrap_err();
    for secret in [
        "user",
        "sensitive-password-value",
        "sensitive-query-value",
        "sensitive-fragment-value",
    ] {
        assert!(!format!("{error:?}").contains(secret));
        assert!(!error.to_string().contains(secret));
    }
}

#[test]
fn remote_plaintext_requires_an_explicit_trusted_mesh_boundary() {
    let endpoint = ServiceEndpoint::new("http://context.service:5120")
        .with_bearer_token("must-not-cross-plaintext");
    let transport = HttpTransport::new_with_options(
        Client::new(),
        "context",
        endpoint.clone(),
        ClientOptions::default(),
    );
    assert!(matches!(
        transport.url("/health"),
        Err(InfraClientError::InvalidEndpoint { .. })
    ));

    let trusted = HttpTransport::new_with_options(
        Client::new(),
        "context",
        endpoint,
        ClientOptions {
            trusted_mesh_http: true,
            ..ClientOptions::default()
        },
    );
    assert!(trusted.url("/health").is_ok());
}

#[derive(Debug)]
struct RecordingCredentials(Mutex<Vec<String>>);

#[async_trait]
impl CredentialsProvider for RecordingCredentials {
    async fn credential(&self, audience: &str) -> Result<BearerCredential, CredentialError> {
        self.0.lock().unwrap().push(audience.to_string());
        BearerCredential::new("opaque-token")
    }
}

#[tokio::test]
async fn explicit_gateway_audience_overrides_the_service_label() {
    let provider = Arc::new(RecordingCredentials(Mutex::new(Vec::new())));
    let endpoint = ServiceEndpoint::new("http://127.0.0.1:0")
        .with_credentials(provider.clone())
        .with_credential_audience("agent-infra");
    let transport = HttpTransport::new_with_options(
        Client::new(),
        "context",
        endpoint,
        ClientOptions {
            max_response_bytes: 1024,
            ..ClientOptions::default()
        },
    );
    let _: Result<serde_json::Value, _> = transport.get_json("/unreachable").await;
    assert_eq!(provider.0.lock().unwrap().as_slice(), ["agent-infra"]);
}

#[derive(Debug)]
struct HangingCredentials;

#[async_trait]
impl CredentialsProvider for HangingCredentials {
    async fn credential(&self, _audience: &str) -> Result<BearerCredential, CredentialError> {
        std::future::pending().await
    }
}

#[tokio::test]
async fn request_deadline_includes_dynamic_credential_acquisition() {
    let telemetry = Arc::new(RecordingTelemetry::default());
    let endpoint =
        ServiceEndpoint::new("http://127.0.0.1:0").with_credentials(Arc::new(HangingCredentials));
    let transport = HttpTransport::new_with_options(
        Client::new(),
        "context",
        endpoint,
        ClientOptions {
            request_timeout: Duration::from_millis(20),
            telemetry: telemetry.clone(),
            ..ClientOptions::default()
        },
    );
    let error: InfraClientError = transport
        .get_json::<serde_json::Value>("/never")
        .await
        .unwrap_err();
    assert!(matches!(error, InfraClientError::DeadlineExceeded { .. }));
    assert!(
        telemetry
            .0
            .lock()
            .unwrap()
            .iter()
            .any(|event| { event.phase == TelemetryPhase::Result && event.outcome == "deadline" })
    );
}

#[derive(Debug, Default)]
struct RecordingTelemetry(Mutex<Vec<TelemetryEvent>>);

impl TelemetryObserver for RecordingTelemetry {
    fn observe(&self, event: TelemetryEvent) {
        self.0.lock().unwrap().push(event);
    }
}

#[tokio::test]
async fn telemetry_observer_receives_bounded_attempt_status_and_size_metadata() {
    let server = wiremock::MockServer::start().await;
    wiremock::Mock::given(wiremock::matchers::method("GET"))
        .and(wiremock::matchers::path("/value"))
        .respond_with(
            wiremock::ResponseTemplate::new(200).set_body_json(serde_json::json!({ "ok": true })),
        )
        .mount(&server)
        .await;
    let telemetry = Arc::new(RecordingTelemetry::default());
    let transport = HttpTransport::new_with_options(
        Client::new(),
        "context",
        ServiceEndpoint::new(server.uri()),
        ClientOptions {
            telemetry: telemetry.clone(),
            ..ClientOptions::default()
        },
    );
    let _: serde_json::Value = transport.get_json("/value").await.unwrap();
    let events = telemetry.0.lock().unwrap();
    assert!(
        events
            .iter()
            .any(|event| event.phase == TelemetryPhase::Attempt)
    );
    let result = events
        .iter()
        .find(|event| event.phase == TelemetryPhase::Result)
        .unwrap();
    assert_eq!(result.status, Some(200));
    assert!(result.response_bytes.is_some_and(|bytes| bytes > 0));
    assert_eq!(result.outcome, "success");
}