use std::fmt;
use std::net::IpAddr;
use std::time::Duration;
use async_trait::async_trait;
use contextgraph_types::{
Capabilities, ContextQuery, ContextQueryResult, PROTOCOL_VERSION, ProviderInfo, VerifyRequest,
VerifyResponse,
};
use crate::error::HostError;
use crate::provider::ContextProvider;
use crate::wire::{
Envelope, envelope_kind, next_correlation_id, verify_correlation, versions_compatible,
};
const HTTP_TIMEOUT: Duration = Duration::from_secs(30);
#[derive(Clone)]
pub struct Credential {
token: String,
}
impl Credential {
pub fn bearer(token: impl Into<String>) -> Self {
Self {
token: token.into(),
}
}
fn expose(&self) -> &str {
&self.token
}
}
impl fmt::Debug for Credential {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("Credential(<redacted>)")
}
}
impl fmt::Display for Credential {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("Credential(<redacted>)")
}
}
fn is_loopback_host(host: &str) -> bool {
if host.eq_ignore_ascii_case("localhost") {
return true;
}
let bare = host
.strip_prefix('[')
.and_then(|inner| inner.strip_suffix(']'))
.unwrap_or(host);
matches!(bare.parse::<IpAddr>(), Ok(ip) if ip.is_loopback())
}
fn refuse_insecure_transport(id: &str, url: &str) -> Result<(), HostError> {
let parsed = reqwest::Url::parse(url).map_err(|e| HostError::Transport {
id: id.to_string(),
message: format!("invalid provider url: {e}"),
})?;
if parsed.scheme() == "http" {
let host = parsed.host_str().unwrap_or("");
if !is_loopback_host(host) {
return Err(HostError::InsecureTransport {
id: id.to_string(),
host: host.to_string(),
});
}
}
Ok(())
}
pub struct HttpProvider {
id: String,
url: String,
client: reqwest::Client,
info: ProviderInfo,
capabilities: Capabilities,
credential: Option<Credential>,
}
impl HttpProvider {
pub async fn connect(id: impl Into<String>, url: impl Into<String>) -> Result<Self, HostError> {
Self::connect_with_auth(id, url, None).await
}
pub async fn connect_with_auth(
id: impl Into<String>,
url: impl Into<String>,
credential: Option<Credential>,
) -> Result<Self, HostError> {
let id = id.into();
let url = url.into();
refuse_insecure_transport(&id, &url)?;
let client = reqwest::Client::builder()
.timeout(HTTP_TIMEOUT)
.build()
.map_err(|e| HostError::Transport {
id: id.clone(),
message: format!("building HTTP client: {e}"),
})?;
let ack = post_envelope(
&client,
&url,
&Envelope::Handshake {
protocol_version: PROTOCOL_VERSION.to_string(),
},
&id,
credential.as_ref(),
)
.await?;
match ack {
Envelope::HandshakeAck {
protocol_version,
provider,
capabilities,
} => {
if !versions_compatible(PROTOCOL_VERSION, &protocol_version) {
return Err(HostError::VersionMismatch {
host: PROTOCOL_VERSION.to_string(),
provider: provider.name,
provider_version: protocol_version,
});
}
let mut info = provider;
info.data_flow.egress = true;
Ok(Self {
id,
url,
client,
info,
capabilities,
credential,
})
}
other => Err(HostError::UnexpectedEnvelope {
id,
expected: "handshake_ack".into(),
got: envelope_kind(&other).into(),
}),
}
}
}
async fn post_envelope(
client: &reqwest::Client,
url: &str,
env: &Envelope,
id: &str,
credential: Option<&Credential>,
) -> Result<Envelope, HostError> {
let mut request = client.post(url).json(env);
if let Some(credential) = credential {
request = request.bearer_auth(credential.expose());
}
let response = request.send().await.map_err(|e| HostError::Transport {
id: id.to_string(),
message: e.to_string(),
})?;
if response.status() == reqwest::StatusCode::UNAUTHORIZED {
return Err(HostError::Unauthorized { id: id.to_string() });
}
if !response.status().is_success() {
let status = response.status();
let body = response.text().await.unwrap_or_default();
return Err(HostError::Transport {
id: id.to_string(),
message: format!("HTTP {status}: {body}"),
});
}
response.json::<Envelope>().await.map_err(|e| {
HostError::Wire(format!(
"provider {id} returned a non-envelope HTTP body: {e}"
))
})
}
#[async_trait]
impl ContextProvider for HttpProvider {
fn id(&self) -> &str {
&self.id
}
fn info(&self) -> &ProviderInfo {
&self.info
}
fn capabilities(&self) -> &Capabilities {
&self.capabilities
}
async fn query(&self, query: &ContextQuery) -> Result<ContextQueryResult, HostError> {
let sent_id = self.capabilities.correlation.then(next_correlation_id);
let reply = post_envelope(
&self.client,
&self.url,
&Envelope::Query {
id: sent_id.clone(),
query: query.clone(),
},
&self.id,
self.credential.as_ref(),
)
.await?;
match reply {
Envelope::Frames { id: echoed, result } => {
verify_correlation(&self.id, sent_id.as_deref(), echoed.as_deref())?;
Ok(result)
}
Envelope::Error { message, code, .. } => Err(HostError::Provider {
id: self.id.clone(),
code,
message,
}),
other => Err(HostError::UnexpectedEnvelope {
id: self.id.clone(),
expected: "frames".into(),
got: envelope_kind(&other).into(),
}),
}
}
async fn verify(&self, request: &VerifyRequest) -> Result<VerifyResponse, HostError> {
let reply = post_envelope(
&self.client,
&self.url,
&Envelope::Verify {
request: request.clone(),
},
&self.id,
self.credential.as_ref(),
)
.await?;
match reply {
Envelope::Verified { response } => Ok(response),
Envelope::Error { message, code, .. } => Err(HostError::Provider {
id: self.id.clone(),
code,
message,
}),
other => Err(HostError::UnexpectedEnvelope {
id: self.id.clone(),
expected: "verified".into(),
got: envelope_kind(&other).into(),
}),
}
}
async fn shutdown(&self) -> Result<(), HostError> {
let _ = post_envelope(
&self.client,
&self.url,
&Envelope::Shutdown,
&self.id,
self.credential.as_ref(),
)
.await;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use contextgraph_types::capability::QueryCapability;
use contextgraph_types::{ContextFrame, DataFlow, FrameKind};
use wiremock::matchers::{header, method};
use wiremock::{Mock, MockServer, ResponseTemplate};
fn ack_body(version: &str) -> serde_json::Value {
serde_json::to_value(Envelope::HandshakeAck {
protocol_version: version.to_string(),
provider: ProviderInfo {
name: "remote-docs".into(),
version: "0.1.0".into(),
data_flow: DataFlow {
reads: true,
writes: false,
egress: true,
egress_scopes: vec![],
},
},
capabilities: Capabilities {
query: QueryCapability {
kinds: vec!["doc".into()],
},
..Capabilities::default()
},
})
.unwrap()
}
fn frames_body() -> serde_json::Value {
serde_json::to_value(Envelope::Frames {
id: None,
result: ContextQueryResult {
frames: vec![ContextFrame {
id: "frm_h".into(),
kind: FrameKind::Doc,
title: "remote doc".into(),
content: Some("remote content".into()),
content_digest: None,
uri: Some("https://example.test/doc".into()),
representation: Default::default(),
content_fidelity: None,
canonical_content_hash: None,
content_ref: None,
transform: None,
minimum_content_fidelity: None,
inline_content_requirement: None,
score: 0.6,
token_cost: 20,
canonical_token_cost: None,
tokenizer_ref: None,
valid_from: None,
valid_to: None,
recorded_at: None,
provenance: vec![],
citation_label: Some("remote doc".into()),
embedding: None,
relations: vec![],
}],
truncated: false,
dropped_estimate: None,
},
})
.unwrap()
}
fn sample_query() -> ContextQuery {
ContextQuery {
goal: "g".into(),
query_text: None,
embedding: None,
kinds: vec![],
anchors: vec![],
max_frames: 5,
max_tokens: 4000,
as_of: None,
representation_preferences: vec![],
}
}
#[tokio::test]
async fn http_handshake_then_query_round_trips_via_wiremock() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(|req: &wiremock::Request| {
let body = match serde_json::from_slice::<Envelope>(&req.body) {
Ok(Envelope::Handshake { .. }) => ack_body(PROTOCOL_VERSION),
Ok(Envelope::Query { .. }) => frames_body(),
_ => serde_json::to_value(Envelope::Error {
id: None,
code: None,
message: "unexpected request".into(),
})
.unwrap(),
};
ResponseTemplate::new(200).set_body_json(body)
})
.mount(&server)
.await;
let provider = HttpProvider::connect("remote", server.uri())
.await
.expect("handshake ok");
assert_eq!(provider.info().name, "remote-docs");
assert!(provider.info().data_flow.egress);
let result = provider.query(&sample_query()).await.expect("query ok");
assert_eq!(result.frames.len(), 1);
assert_eq!(result.frames[0].title, "remote doc");
}
#[tokio::test]
async fn http_version_mismatch_rejects_the_provider() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200).set_body_json(ack_body("contextgraph/2.0")))
.mount(&server)
.await;
let err = match HttpProvider::connect("remote", server.uri()).await {
Ok(_) => panic!("incompatible version must reject"),
Err(e) => e,
};
assert!(matches!(err, HostError::VersionMismatch { .. }));
}
#[tokio::test]
async fn a_non_envelope_http_body_is_a_clean_wire_error() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(
ResponseTemplate::new(200).set_body_string("<html>not contextgraph</html>"),
)
.mount(&server)
.await;
let err = match HttpProvider::connect("remote", server.uri()).await {
Ok(_) => panic!("garbage body must not panic the host"),
Err(e) => e,
};
assert!(matches!(err, HostError::Wire(_)));
}
#[tokio::test]
async fn http_transport_forces_egress_even_when_the_remote_claims_local() {
let server = MockServer::start().await;
let sneaky_ack = serde_json::to_value(Envelope::HandshakeAck {
protocol_version: PROTOCOL_VERSION.to_string(),
provider: ProviderInfo {
name: "sneaky-remote".into(),
version: "0.1.0".into(),
data_flow: DataFlow {
reads: true,
writes: false,
egress: false, egress_scopes: vec![],
},
},
capabilities: Capabilities {
query: QueryCapability {
kinds: vec!["doc".into()],
},
..Capabilities::default()
},
})
.unwrap();
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200).set_body_json(sneaky_ack))
.mount(&server)
.await;
let provider = HttpProvider::connect("remote", server.uri())
.await
.expect("handshake ok");
assert!(
provider.info().data_flow.egress,
"an HTTP transport must be treated as egress regardless of the remote's claim"
);
assert!(
crate::consent::ConsentStore::requires_consent(provider.info()),
"an HTTP provider must always require consent, even claiming egress:false"
);
}
#[tokio::test]
async fn a_plaintext_non_loopback_transport_is_refused_before_any_bytes_leave() {
let err = match HttpProvider::connect("remote", "http://example.com:9/cgp").await {
Ok(_) => panic!("a plaintext non-loopback transport must be refused (C7)"),
Err(e) => e,
};
match err {
HostError::InsecureTransport { id, host } => {
assert_eq!(id, "remote");
assert_eq!(host, "example.com");
}
other => panic!("expected InsecureTransport, got {other:?}"),
}
}
#[tokio::test]
async fn a_plaintext_loopback_transport_is_allowed() {
let server = MockServer::start().await;
assert!(
server.uri().starts_with("http://"),
"wiremock serves plaintext http on loopback"
);
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200).set_body_json(ack_body(PROTOCOL_VERSION)))
.mount(&server)
.await;
let provider = HttpProvider::connect("remote", server.uri())
.await
.expect("a plaintext loopback (127.0.0.1) transport is allowed");
assert_eq!(provider.info().name, "remote-docs");
}
#[tokio::test]
async fn a_supplied_credential_is_attached_as_a_bearer_header() {
const TOKEN: &str = "s3cr3t-bearer-token-value";
let server = MockServer::start().await;
let auth_value = format!("Bearer {TOKEN}");
Mock::given(method("POST"))
.and(header("authorization", auth_value.as_str()))
.respond_with(|req: &wiremock::Request| {
let body = match serde_json::from_slice::<Envelope>(&req.body) {
Ok(Envelope::Handshake { .. }) => ack_body(PROTOCOL_VERSION),
Ok(Envelope::Query { .. }) => frames_body(),
_ => serde_json::to_value(Envelope::Error {
id: None,
code: None,
message: "unexpected request".into(),
})
.unwrap(),
};
ResponseTemplate::new(200).set_body_json(body)
})
.mount(&server)
.await;
let provider = HttpProvider::connect_with_auth(
"remote",
server.uri(),
Some(Credential::bearer(TOKEN)),
)
.await
.expect("handshake carries the bearer credential");
let result = provider.query(&sample_query()).await.expect("query ok");
assert_eq!(result.frames.len(), 1);
}
#[test]
fn a_credential_is_redacted_in_every_rendering_and_never_in_an_error() {
const SECRET: &str = "ghp_this_must_never_appear_in_a_log_0xDEADBEEF";
let credential = Credential::bearer(SECRET);
let debug = format!("{credential:?}");
let display = format!("{credential}");
assert_eq!(debug, "Credential(<redacted>)");
assert_eq!(display, "Credential(<redacted>)");
assert!(
!debug.contains(SECRET),
"Debug must not leak the secret (C8)"
);
assert!(
!display.contains(SECRET),
"Display must not leak the secret (C8)"
);
assert_eq!(
format!("{:?}", credential.clone()),
"Credential(<redacted>)"
);
let insecure = HostError::InsecureTransport {
id: "remote".into(),
host: "example.com".into(),
};
let unauthorized = HostError::Unauthorized {
id: "remote".into(),
};
assert!(!insecure.to_string().contains(SECRET));
assert!(!unauthorized.to_string().contains(SECRET));
}
}