1use std::fmt;
13use std::net::IpAddr;
14use std::time::Duration;
15
16use async_trait::async_trait;
17use contextgraph_types::{
18 Capabilities, ContextQuery, ContextQueryResult, PROTOCOL_VERSION, ProviderInfo, VerifyRequest,
19 VerifyResponse,
20};
21
22use crate::error::HostError;
23use crate::provider::ContextProvider;
24use crate::wire::{
25 Envelope, envelope_kind, next_correlation_id, verify_correlation, versions_compatible,
26};
27
28const HTTP_TIMEOUT: Duration = Duration::from_secs(30);
30
31#[derive(Clone)]
40pub struct Credential {
41 token: String,
44}
45
46impl Credential {
47 pub fn bearer(token: impl Into<String>) -> Self {
50 Self {
51 token: token.into(),
52 }
53 }
54
55 fn expose(&self) -> &str {
58 &self.token
59 }
60}
61
62impl fmt::Debug for Credential {
65 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
66 f.write_str("Credential(<redacted>)")
67 }
68}
69
70impl fmt::Display for Credential {
73 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
74 f.write_str("Credential(<redacted>)")
75 }
76}
77
78fn is_loopback_host(host: &str) -> bool {
86 if host.eq_ignore_ascii_case("localhost") {
87 return true;
88 }
89 let bare = host
90 .strip_prefix('[')
91 .and_then(|inner| inner.strip_suffix(']'))
92 .unwrap_or(host);
93 matches!(bare.parse::<IpAddr>(), Ok(ip) if ip.is_loopback())
95}
96
97fn refuse_insecure_transport(id: &str, url: &str) -> Result<(), HostError> {
105 let parsed = reqwest::Url::parse(url).map_err(|e| HostError::Transport {
106 id: id.to_string(),
107 message: format!("invalid provider url: {e}"),
108 })?;
109 if parsed.scheme() == "http" {
110 let host = parsed.host_str().unwrap_or("");
111 if !is_loopback_host(host) {
112 return Err(HostError::InsecureTransport {
113 id: id.to_string(),
114 host: host.to_string(),
115 });
116 }
117 }
118 Ok(())
119}
120
121pub struct HttpProvider {
124 id: String,
125 url: String,
126 client: reqwest::Client,
127 info: ProviderInfo,
128 capabilities: Capabilities,
129 credential: Option<Credential>,
132}
133
134impl HttpProvider {
135 pub async fn connect(id: impl Into<String>, url: impl Into<String>) -> Result<Self, HostError> {
140 Self::connect_with_auth(id, url, None).await
141 }
142
143 pub async fn connect_with_auth(
150 id: impl Into<String>,
151 url: impl Into<String>,
152 credential: Option<Credential>,
153 ) -> Result<Self, HostError> {
154 let id = id.into();
155 let url = url.into();
156 refuse_insecure_transport(&id, &url)?;
159 let client = reqwest::Client::builder()
160 .timeout(HTTP_TIMEOUT)
161 .build()
162 .map_err(|e| HostError::Transport {
163 id: id.clone(),
164 message: format!("building HTTP client: {e}"),
165 })?;
166
167 let ack = post_envelope(
168 &client,
169 &url,
170 &Envelope::Handshake {
171 protocol_version: PROTOCOL_VERSION.to_string(),
172 },
173 &id,
174 credential.as_ref(),
175 )
176 .await?;
177
178 match ack {
179 Envelope::HandshakeAck {
180 protocol_version,
181 provider,
182 capabilities,
183 } => {
184 if !versions_compatible(PROTOCOL_VERSION, &protocol_version) {
185 return Err(HostError::VersionMismatch {
186 host: PROTOCOL_VERSION.to_string(),
187 provider: provider.name,
188 provider_version: protocol_version,
189 });
190 }
191 let mut info = provider;
199 info.data_flow.egress = true;
200 Ok(Self {
201 id,
202 url,
203 client,
204 info,
205 capabilities,
206 credential,
207 })
208 }
209 other => Err(HostError::UnexpectedEnvelope {
210 id,
211 expected: "handshake_ack".into(),
212 got: envelope_kind(&other).into(),
213 }),
214 }
215 }
216}
217
218async fn post_envelope(
226 client: &reqwest::Client,
227 url: &str,
228 env: &Envelope,
229 id: &str,
230 credential: Option<&Credential>,
231) -> Result<Envelope, HostError> {
232 let mut request = client.post(url).json(env);
233 if let Some(credential) = credential {
234 request = request.bearer_auth(credential.expose());
235 }
236 let response = request.send().await.map_err(|e| HostError::Transport {
237 id: id.to_string(),
238 message: e.to_string(),
239 })?;
240
241 if response.status() == reqwest::StatusCode::UNAUTHORIZED {
245 return Err(HostError::Unauthorized { id: id.to_string() });
246 }
247
248 if !response.status().is_success() {
249 let status = response.status();
250 let body = response.text().await.unwrap_or_default();
251 return Err(HostError::Transport {
252 id: id.to_string(),
253 message: format!("HTTP {status}: {body}"),
254 });
255 }
256
257 response.json::<Envelope>().await.map_err(|e| {
258 HostError::Wire(format!(
259 "provider {id} returned a non-envelope HTTP body: {e}"
260 ))
261 })
262}
263
264#[async_trait]
265impl ContextProvider for HttpProvider {
266 fn id(&self) -> &str {
267 &self.id
268 }
269
270 fn info(&self) -> &ProviderInfo {
271 &self.info
272 }
273
274 fn capabilities(&self) -> &Capabilities {
275 &self.capabilities
276 }
277
278 async fn query(&self, query: &ContextQuery) -> Result<ContextQueryResult, HostError> {
279 let sent_id = self.capabilities.correlation.then(next_correlation_id);
280 let reply = post_envelope(
281 &self.client,
282 &self.url,
283 &Envelope::Query {
284 id: sent_id.clone(),
285 query: query.clone(),
286 },
287 &self.id,
288 self.credential.as_ref(),
289 )
290 .await?;
291 match reply {
292 Envelope::Frames { id: echoed, result } => {
293 verify_correlation(&self.id, sent_id.as_deref(), echoed.as_deref())?;
294 Ok(result)
295 }
296 Envelope::Error { message, code, .. } => Err(HostError::Provider {
297 id: self.id.clone(),
298 code,
299 message,
300 }),
301 other => Err(HostError::UnexpectedEnvelope {
302 id: self.id.clone(),
303 expected: "frames".into(),
304 got: envelope_kind(&other).into(),
305 }),
306 }
307 }
308
309 async fn verify(&self, request: &VerifyRequest) -> Result<VerifyResponse, HostError> {
310 let reply = post_envelope(
311 &self.client,
312 &self.url,
313 &Envelope::Verify {
314 request: request.clone(),
315 },
316 &self.id,
317 self.credential.as_ref(),
318 )
319 .await?;
320 match reply {
321 Envelope::Verified { response } => Ok(response),
322 Envelope::Error { message, code, .. } => Err(HostError::Provider {
323 id: self.id.clone(),
324 code,
325 message,
326 }),
327 other => Err(HostError::UnexpectedEnvelope {
328 id: self.id.clone(),
329 expected: "verified".into(),
330 got: envelope_kind(&other).into(),
331 }),
332 }
333 }
334
335 async fn shutdown(&self) -> Result<(), HostError> {
336 let _ = post_envelope(
338 &self.client,
339 &self.url,
340 &Envelope::Shutdown,
341 &self.id,
342 self.credential.as_ref(),
343 )
344 .await;
345 Ok(())
346 }
347}
348
349#[cfg(test)]
350mod tests {
351 use super::*;
352 use contextgraph_types::capability::QueryCapability;
353 use contextgraph_types::{ContextFrame, DataFlow, FrameKind};
354 use wiremock::matchers::{header, method};
355 use wiremock::{Mock, MockServer, ResponseTemplate};
356
357 fn ack_body(version: &str) -> serde_json::Value {
358 serde_json::to_value(Envelope::HandshakeAck {
359 protocol_version: version.to_string(),
360 provider: ProviderInfo {
361 name: "remote-docs".into(),
362 version: "0.1.0".into(),
363 data_flow: DataFlow {
364 reads: true,
365 writes: false,
366 egress: true,
367 egress_scopes: vec![],
368 },
369 },
370 capabilities: Capabilities {
371 query: QueryCapability {
372 kinds: vec!["doc".into()],
373 },
374 ..Capabilities::default()
375 },
376 })
377 .unwrap()
378 }
379
380 fn frames_body() -> serde_json::Value {
381 serde_json::to_value(Envelope::Frames {
382 id: None,
383 result: ContextQueryResult {
384 frames: vec![ContextFrame {
385 id: "frm_h".into(),
386 kind: FrameKind::Doc,
387 title: "remote doc".into(),
388 content: Some("remote content".into()),
389 content_digest: None,
390 uri: Some("https://example.test/doc".into()),
391 representation: Default::default(),
392 content_fidelity: None,
393 canonical_content_hash: None,
394 content_ref: None,
395 transform: None,
396 minimum_content_fidelity: None,
397 inline_content_requirement: None,
398 score: 0.6,
399 token_cost: 20,
400 canonical_token_cost: None,
401 tokenizer_ref: None,
402 valid_from: None,
403 valid_to: None,
404 recorded_at: None,
405 provenance: vec![],
406 citation_label: Some("remote doc".into()),
407 embedding: None,
408 relations: vec![],
409 }],
410 truncated: false,
411 dropped_estimate: None,
412 },
413 })
414 .unwrap()
415 }
416
417 fn sample_query() -> ContextQuery {
418 ContextQuery {
419 goal: "g".into(),
420 query_text: None,
421 embedding: None,
422 kinds: vec![],
423 anchors: vec![],
424 max_frames: 5,
425 max_tokens: 4000,
426 as_of: None,
427 representation_preferences: vec![],
428 }
429 }
430
431 #[tokio::test]
432 async fn http_handshake_then_query_round_trips_via_wiremock() {
433 let server = MockServer::start().await;
434 Mock::given(method("POST"))
437 .respond_with(|req: &wiremock::Request| {
438 let body = match serde_json::from_slice::<Envelope>(&req.body) {
439 Ok(Envelope::Handshake { .. }) => ack_body(PROTOCOL_VERSION),
440 Ok(Envelope::Query { .. }) => frames_body(),
441 _ => serde_json::to_value(Envelope::Error {
442 id: None,
443 code: None,
444 message: "unexpected request".into(),
445 })
446 .unwrap(),
447 };
448 ResponseTemplate::new(200).set_body_json(body)
449 })
450 .mount(&server)
451 .await;
452
453 let provider = HttpProvider::connect("remote", server.uri())
454 .await
455 .expect("handshake ok");
456 assert_eq!(provider.info().name, "remote-docs");
457 assert!(provider.info().data_flow.egress);
458
459 let result = provider.query(&sample_query()).await.expect("query ok");
460 assert_eq!(result.frames.len(), 1);
461 assert_eq!(result.frames[0].title, "remote doc");
462 }
463
464 #[tokio::test]
465 async fn http_version_mismatch_rejects_the_provider() {
466 let server = MockServer::start().await;
467 Mock::given(method("POST"))
468 .respond_with(ResponseTemplate::new(200).set_body_json(ack_body("contextgraph/2.0")))
469 .mount(&server)
470 .await;
471
472 let err = match HttpProvider::connect("remote", server.uri()).await {
473 Ok(_) => panic!("incompatible version must reject"),
474 Err(e) => e,
475 };
476 assert!(matches!(err, HostError::VersionMismatch { .. }));
477 }
478
479 #[tokio::test]
480 async fn a_non_envelope_http_body_is_a_clean_wire_error() {
481 let server = MockServer::start().await;
482 Mock::given(method("POST"))
483 .respond_with(
484 ResponseTemplate::new(200).set_body_string("<html>not contextgraph</html>"),
485 )
486 .mount(&server)
487 .await;
488
489 let err = match HttpProvider::connect("remote", server.uri()).await {
490 Ok(_) => panic!("garbage body must not panic the host"),
491 Err(e) => e,
492 };
493 assert!(matches!(err, HostError::Wire(_)));
494 }
495
496 #[tokio::test]
497 async fn http_transport_forces_egress_even_when_the_remote_claims_local() {
498 let server = MockServer::start().await;
503 let sneaky_ack = serde_json::to_value(Envelope::HandshakeAck {
504 protocol_version: PROTOCOL_VERSION.to_string(),
505 provider: ProviderInfo {
506 name: "sneaky-remote".into(),
507 version: "0.1.0".into(),
508 data_flow: DataFlow {
509 reads: true,
510 writes: false,
511 egress: false, egress_scopes: vec![],
513 },
514 },
515 capabilities: Capabilities {
516 query: QueryCapability {
517 kinds: vec!["doc".into()],
518 },
519 ..Capabilities::default()
520 },
521 })
522 .unwrap();
523 Mock::given(method("POST"))
524 .respond_with(ResponseTemplate::new(200).set_body_json(sneaky_ack))
525 .mount(&server)
526 .await;
527
528 let provider = HttpProvider::connect("remote", server.uri())
529 .await
530 .expect("handshake ok");
531
532 assert!(
533 provider.info().data_flow.egress,
534 "an HTTP transport must be treated as egress regardless of the remote's claim"
535 );
536 assert!(
537 crate::consent::ConsentStore::requires_consent(provider.info()),
538 "an HTTP provider must always require consent, even claiming egress:false"
539 );
540 }
541
542 #[tokio::test]
545 async fn a_plaintext_non_loopback_transport_is_refused_before_any_bytes_leave() {
546 let err = match HttpProvider::connect("remote", "http://example.com:9/cgp").await {
553 Ok(_) => panic!("a plaintext non-loopback transport must be refused (C7)"),
554 Err(e) => e,
555 };
556 match err {
557 HostError::InsecureTransport { id, host } => {
558 assert_eq!(id, "remote");
559 assert_eq!(host, "example.com");
560 }
561 other => panic!("expected InsecureTransport, got {other:?}"),
562 }
563 }
564
565 #[tokio::test]
566 async fn a_plaintext_loopback_transport_is_allowed() {
567 let server = MockServer::start().await;
572 assert!(
573 server.uri().starts_with("http://"),
574 "wiremock serves plaintext http on loopback"
575 );
576 Mock::given(method("POST"))
577 .respond_with(ResponseTemplate::new(200).set_body_json(ack_body(PROTOCOL_VERSION)))
578 .mount(&server)
579 .await;
580 let provider = HttpProvider::connect("remote", server.uri())
581 .await
582 .expect("a plaintext loopback (127.0.0.1) transport is allowed");
583 assert_eq!(provider.info().name, "remote-docs");
584 }
585
586 #[tokio::test]
587 async fn a_supplied_credential_is_attached_as_a_bearer_header() {
588 const TOKEN: &str = "s3cr3t-bearer-token-value";
589 let server = MockServer::start().await;
590 let auth_value = format!("Bearer {TOKEN}");
591 Mock::given(method("POST"))
596 .and(header("authorization", auth_value.as_str()))
597 .respond_with(|req: &wiremock::Request| {
598 let body = match serde_json::from_slice::<Envelope>(&req.body) {
599 Ok(Envelope::Handshake { .. }) => ack_body(PROTOCOL_VERSION),
600 Ok(Envelope::Query { .. }) => frames_body(),
601 _ => serde_json::to_value(Envelope::Error {
602 id: None,
603 code: None,
604 message: "unexpected request".into(),
605 })
606 .unwrap(),
607 };
608 ResponseTemplate::new(200).set_body_json(body)
609 })
610 .mount(&server)
611 .await;
612
613 let provider = HttpProvider::connect_with_auth(
614 "remote",
615 server.uri(),
616 Some(Credential::bearer(TOKEN)),
617 )
618 .await
619 .expect("handshake carries the bearer credential");
620 let result = provider.query(&sample_query()).await.expect("query ok");
622 assert_eq!(result.frames.len(), 1);
623 }
624
625 #[test]
626 fn a_credential_is_redacted_in_every_rendering_and_never_in_an_error() {
627 const SECRET: &str = "ghp_this_must_never_appear_in_a_log_0xDEADBEEF";
631 let credential = Credential::bearer(SECRET);
632
633 let debug = format!("{credential:?}");
634 let display = format!("{credential}");
635 assert_eq!(debug, "Credential(<redacted>)");
636 assert_eq!(display, "Credential(<redacted>)");
637 assert!(
638 !debug.contains(SECRET),
639 "Debug must not leak the secret (C8)"
640 );
641 assert!(
642 !display.contains(SECRET),
643 "Display must not leak the secret (C8)"
644 );
645 assert_eq!(
647 format!("{:?}", credential.clone()),
648 "Credential(<redacted>)"
649 );
650
651 let insecure = HostError::InsecureTransport {
655 id: "remote".into(),
656 host: "example.com".into(),
657 };
658 let unauthorized = HostError::Unauthorized {
659 id: "remote".into(),
660 };
661 assert!(!insecure.to_string().contains(SECRET));
662 assert!(!unauthorized.to_string().contains(SECRET));
663 }
664}