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
97pub fn refuse_insecure_transport(id: &str, url: &str) -> Result<(), HostError> {
130 let parsed = reqwest::Url::parse(url).map_err(|e| HostError::Transport {
131 id: id.to_string(),
132 message: format!("invalid provider url: {e}"),
133 })?;
134 if parsed.scheme() == "http" {
135 let host = parsed.host_str().unwrap_or("");
136 if !is_loopback_host(host) {
137 return Err(HostError::InsecureTransport {
138 id: id.to_string(),
139 host: host.to_string(),
140 });
141 }
142 }
143 Ok(())
144}
145
146pub struct HttpProvider {
149 id: String,
150 url: String,
151 client: reqwest::Client,
152 info: ProviderInfo,
153 capabilities: Capabilities,
154 credential: Option<Credential>,
157}
158
159impl HttpProvider {
160 pub async fn connect(id: impl Into<String>, url: impl Into<String>) -> Result<Self, HostError> {
165 Self::connect_with_auth(id, url, None).await
166 }
167
168 pub async fn connect_with_auth(
175 id: impl Into<String>,
176 url: impl Into<String>,
177 credential: Option<Credential>,
178 ) -> Result<Self, HostError> {
179 let id = id.into();
180 let url = url.into();
181 refuse_insecure_transport(&id, &url)?;
184 let client = reqwest::Client::builder()
185 .timeout(HTTP_TIMEOUT)
186 .build()
187 .map_err(|e| HostError::Transport {
188 id: id.clone(),
189 message: format!("building HTTP client: {e}"),
190 })?;
191
192 let ack = post_envelope(
193 &client,
194 &url,
195 &Envelope::Handshake {
196 protocol_version: PROTOCOL_VERSION.to_string(),
197 },
198 &id,
199 credential.as_ref(),
200 )
201 .await?;
202
203 match ack {
204 Envelope::HandshakeAck {
205 protocol_version,
206 provider,
207 capabilities,
208 } => {
209 if !versions_compatible(PROTOCOL_VERSION, &protocol_version) {
210 return Err(HostError::VersionMismatch {
211 host: PROTOCOL_VERSION.to_string(),
212 provider: provider.name,
213 provider_version: protocol_version,
214 });
215 }
216 let mut info = provider;
224 info.data_flow.egress = true;
225 Ok(Self {
226 id,
227 url,
228 client,
229 info,
230 capabilities,
231 credential,
232 })
233 }
234 other => Err(HostError::UnexpectedEnvelope {
235 id,
236 expected: "handshake_ack".into(),
237 got: envelope_kind(&other).into(),
238 }),
239 }
240 }
241}
242
243async fn post_envelope(
251 client: &reqwest::Client,
252 url: &str,
253 env: &Envelope,
254 id: &str,
255 credential: Option<&Credential>,
256) -> Result<Envelope, HostError> {
257 let mut request = client.post(url).json(env);
258 if let Some(credential) = credential {
259 request = request.bearer_auth(credential.expose());
260 }
261 let response = request.send().await.map_err(|e| HostError::Transport {
262 id: id.to_string(),
263 message: e.to_string(),
264 })?;
265
266 if response.status() == reqwest::StatusCode::UNAUTHORIZED {
270 return Err(HostError::Unauthorized { id: id.to_string() });
271 }
272
273 if !response.status().is_success() {
274 let status = response.status();
275 let body = response.text().await.unwrap_or_default();
276 return Err(HostError::Transport {
277 id: id.to_string(),
278 message: format!("HTTP {status}: {body}"),
279 });
280 }
281
282 response.json::<Envelope>().await.map_err(|e| {
283 HostError::Wire(format!(
284 "provider {id} returned a non-envelope HTTP body: {e}"
285 ))
286 })
287}
288
289#[async_trait]
290impl ContextProvider for HttpProvider {
291 fn id(&self) -> &str {
292 &self.id
293 }
294
295 fn info(&self) -> &ProviderInfo {
296 &self.info
297 }
298
299 fn capabilities(&self) -> &Capabilities {
300 &self.capabilities
301 }
302
303 async fn query(&self, query: &ContextQuery) -> Result<ContextQueryResult, HostError> {
304 let sent_id = self.capabilities.correlation.then(next_correlation_id);
305 let reply = post_envelope(
306 &self.client,
307 &self.url,
308 &Envelope::Query {
309 id: sent_id.clone(),
310 query: query.clone(),
311 },
312 &self.id,
313 self.credential.as_ref(),
314 )
315 .await?;
316 match reply {
317 Envelope::Frames { id: echoed, result } => {
318 verify_correlation(&self.id, sent_id.as_deref(), echoed.as_deref())?;
319 Ok(result)
320 }
321 Envelope::Error { message, code, .. } => Err(HostError::Provider {
322 id: self.id.clone(),
323 code,
324 message,
325 }),
326 other => Err(HostError::UnexpectedEnvelope {
327 id: self.id.clone(),
328 expected: "frames".into(),
329 got: envelope_kind(&other).into(),
330 }),
331 }
332 }
333
334 async fn verify(&self, request: &VerifyRequest) -> Result<VerifyResponse, HostError> {
335 let reply = post_envelope(
336 &self.client,
337 &self.url,
338 &Envelope::Verify {
339 request: request.clone(),
340 },
341 &self.id,
342 self.credential.as_ref(),
343 )
344 .await?;
345 match reply {
346 Envelope::Verified { response } => Ok(response),
347 Envelope::Error { message, code, .. } => Err(HostError::Provider {
348 id: self.id.clone(),
349 code,
350 message,
351 }),
352 other => Err(HostError::UnexpectedEnvelope {
353 id: self.id.clone(),
354 expected: "verified".into(),
355 got: envelope_kind(&other).into(),
356 }),
357 }
358 }
359
360 async fn shutdown(&self) -> Result<(), HostError> {
361 let _ = post_envelope(
363 &self.client,
364 &self.url,
365 &Envelope::Shutdown,
366 &self.id,
367 self.credential.as_ref(),
368 )
369 .await;
370 Ok(())
371 }
372}
373
374#[cfg(test)]
375mod tests {
376 use super::*;
377 use contextgraph_types::capability::QueryCapability;
378 use contextgraph_types::{ContextFrame, DataFlow, FrameKind};
379 use wiremock::matchers::{header, method};
380 use wiremock::{Mock, MockServer, ResponseTemplate};
381
382 fn ack_body(version: &str) -> serde_json::Value {
383 serde_json::to_value(Envelope::HandshakeAck {
384 protocol_version: version.to_string(),
385 provider: ProviderInfo {
386 name: "remote-docs".into(),
387 version: "0.1.0".into(),
388 data_flow: DataFlow {
389 reads: true,
390 writes: false,
391 egress: true,
392 egress_scopes: vec![],
393 },
394 },
395 capabilities: Capabilities {
396 query: QueryCapability {
397 kinds: vec!["doc".into()],
398 },
399 ..Capabilities::default()
400 },
401 })
402 .unwrap()
403 }
404
405 fn frames_body() -> serde_json::Value {
406 serde_json::to_value(Envelope::Frames {
407 id: None,
408 result: ContextQueryResult {
409 frames: vec![ContextFrame {
410 id: "frm_h".into(),
411 kind: FrameKind::Doc,
412 title: "remote doc".into(),
413 content: Some("remote content".into()),
414 content_digest: None,
415 uri: Some("https://example.test/doc".into()),
416 representation: Default::default(),
417 content_fidelity: None,
418 canonical_content_hash: None,
419 content_ref: None,
420 transform: None,
421 minimum_content_fidelity: None,
422 inline_content_requirement: None,
423 score: 0.6,
424 token_cost: 20,
425 canonical_token_cost: None,
426 tokenizer_ref: None,
427 valid_from: None,
428 valid_to: None,
429 recorded_at: None,
430 provenance: vec![],
431 citation_label: Some("remote doc".into()),
432 embedding: None,
433 relations: vec![],
434 }],
435 truncated: false,
436 dropped_estimate: None,
437 },
438 })
439 .unwrap()
440 }
441
442 fn sample_query() -> ContextQuery {
443 ContextQuery {
444 goal: "g".into(),
445 query_text: None,
446 embedding: None,
447 kinds: vec![],
448 anchors: vec![],
449 max_frames: 5,
450 max_tokens: 4000,
451 as_of: None,
452 representation_preferences: vec![],
453 }
454 }
455
456 #[tokio::test]
457 async fn http_handshake_then_query_round_trips_via_wiremock() {
458 let server = MockServer::start().await;
459 Mock::given(method("POST"))
462 .respond_with(|req: &wiremock::Request| {
463 let body = match serde_json::from_slice::<Envelope>(&req.body) {
464 Ok(Envelope::Handshake { .. }) => ack_body(PROTOCOL_VERSION),
465 Ok(Envelope::Query { .. }) => frames_body(),
466 _ => serde_json::to_value(Envelope::Error {
467 id: None,
468 code: None,
469 message: "unexpected request".into(),
470 })
471 .unwrap(),
472 };
473 ResponseTemplate::new(200).set_body_json(body)
474 })
475 .mount(&server)
476 .await;
477
478 let provider = HttpProvider::connect("remote", server.uri())
479 .await
480 .expect("handshake ok");
481 assert_eq!(provider.info().name, "remote-docs");
482 assert!(provider.info().data_flow.egress);
483
484 let result = provider.query(&sample_query()).await.expect("query ok");
485 assert_eq!(result.frames.len(), 1);
486 assert_eq!(result.frames[0].title, "remote doc");
487 }
488
489 #[tokio::test]
490 async fn http_version_mismatch_rejects_the_provider() {
491 let server = MockServer::start().await;
492 Mock::given(method("POST"))
493 .respond_with(ResponseTemplate::new(200).set_body_json(ack_body("contextgraph/2.0")))
494 .mount(&server)
495 .await;
496
497 let err = match HttpProvider::connect("remote", server.uri()).await {
498 Ok(_) => panic!("incompatible version must reject"),
499 Err(e) => e,
500 };
501 assert!(matches!(err, HostError::VersionMismatch { .. }));
502 }
503
504 #[tokio::test]
505 async fn a_non_envelope_http_body_is_a_clean_wire_error() {
506 let server = MockServer::start().await;
507 Mock::given(method("POST"))
508 .respond_with(
509 ResponseTemplate::new(200).set_body_string("<html>not contextgraph</html>"),
510 )
511 .mount(&server)
512 .await;
513
514 let err = match HttpProvider::connect("remote", server.uri()).await {
515 Ok(_) => panic!("garbage body must not panic the host"),
516 Err(e) => e,
517 };
518 assert!(matches!(err, HostError::Wire(_)));
519 }
520
521 #[tokio::test]
522 async fn http_transport_forces_egress_even_when_the_remote_claims_local() {
523 let server = MockServer::start().await;
528 let sneaky_ack = serde_json::to_value(Envelope::HandshakeAck {
529 protocol_version: PROTOCOL_VERSION.to_string(),
530 provider: ProviderInfo {
531 name: "sneaky-remote".into(),
532 version: "0.1.0".into(),
533 data_flow: DataFlow {
534 reads: true,
535 writes: false,
536 egress: false, egress_scopes: vec![],
538 },
539 },
540 capabilities: Capabilities {
541 query: QueryCapability {
542 kinds: vec!["doc".into()],
543 },
544 ..Capabilities::default()
545 },
546 })
547 .unwrap();
548 Mock::given(method("POST"))
549 .respond_with(ResponseTemplate::new(200).set_body_json(sneaky_ack))
550 .mount(&server)
551 .await;
552
553 let provider = HttpProvider::connect("remote", server.uri())
554 .await
555 .expect("handshake ok");
556
557 assert!(
558 provider.info().data_flow.egress,
559 "an HTTP transport must be treated as egress regardless of the remote's claim"
560 );
561 assert!(
562 crate::consent::ConsentStore::requires_consent(provider.info()),
563 "an HTTP provider must always require consent, even claiming egress:false"
564 );
565 }
566
567 #[tokio::test]
570 async fn a_plaintext_non_loopback_transport_is_refused_before_any_bytes_leave() {
571 let err = match HttpProvider::connect("remote", "http://example.com:9/cgp").await {
578 Ok(_) => panic!("a plaintext non-loopback transport must be refused (C7)"),
579 Err(e) => e,
580 };
581 match err {
582 HostError::InsecureTransport { id, host } => {
583 assert_eq!(id, "remote");
584 assert_eq!(host, "example.com");
585 }
586 other => panic!("expected InsecureTransport, got {other:?}"),
587 }
588 }
589
590 #[test]
597 fn the_exported_c7_rule_classifies_every_loopback_spelling() {
598 for allowed in [
602 "https://example.com/cgp",
603 "http://localhost:8080/cgp",
604 "http://LOCALHOST:8080/cgp",
605 "http://127.0.0.1/cgp",
606 "http://127.0.0.2/cgp",
607 "http://[::1]:8080/cgp",
608 ] {
609 assert!(
610 refuse_insecure_transport("p", allowed).is_ok(),
611 "C7 must allow {allowed}"
612 );
613 }
614
615 for refused in [
619 "http://example.com/cgp",
620 "http://127.0.0.1.example.com/cgp",
621 "http://[2001:db8::1]/cgp",
622 "http://10.0.0.5/cgp",
623 ] {
624 assert!(
625 matches!(
626 refuse_insecure_transport("p", refused),
627 Err(HostError::InsecureTransport { .. })
628 ),
629 "C7 must refuse {refused}"
630 );
631 }
632
633 assert!(matches!(
637 refuse_insecure_transport("p", "not a url"),
638 Err(HostError::Transport { .. })
639 ));
640 }
641
642 #[tokio::test]
643 async fn a_plaintext_loopback_transport_is_allowed() {
644 let server = MockServer::start().await;
649 assert!(
650 server.uri().starts_with("http://"),
651 "wiremock serves plaintext http on loopback"
652 );
653 Mock::given(method("POST"))
654 .respond_with(ResponseTemplate::new(200).set_body_json(ack_body(PROTOCOL_VERSION)))
655 .mount(&server)
656 .await;
657 let provider = HttpProvider::connect("remote", server.uri())
658 .await
659 .expect("a plaintext loopback (127.0.0.1) transport is allowed");
660 assert_eq!(provider.info().name, "remote-docs");
661 }
662
663 #[tokio::test]
664 async fn a_supplied_credential_is_attached_as_a_bearer_header() {
665 const TOKEN: &str = "s3cr3t-bearer-token-value";
666 let server = MockServer::start().await;
667 let auth_value = format!("Bearer {TOKEN}");
668 Mock::given(method("POST"))
673 .and(header("authorization", auth_value.as_str()))
674 .respond_with(|req: &wiremock::Request| {
675 let body = match serde_json::from_slice::<Envelope>(&req.body) {
676 Ok(Envelope::Handshake { .. }) => ack_body(PROTOCOL_VERSION),
677 Ok(Envelope::Query { .. }) => frames_body(),
678 _ => serde_json::to_value(Envelope::Error {
679 id: None,
680 code: None,
681 message: "unexpected request".into(),
682 })
683 .unwrap(),
684 };
685 ResponseTemplate::new(200).set_body_json(body)
686 })
687 .mount(&server)
688 .await;
689
690 let provider = HttpProvider::connect_with_auth(
691 "remote",
692 server.uri(),
693 Some(Credential::bearer(TOKEN)),
694 )
695 .await
696 .expect("handshake carries the bearer credential");
697 let result = provider.query(&sample_query()).await.expect("query ok");
699 assert_eq!(result.frames.len(), 1);
700 }
701
702 #[test]
703 fn a_credential_is_redacted_in_every_rendering_and_never_in_an_error() {
704 const SECRET: &str = "ghp_this_must_never_appear_in_a_log_0xDEADBEEF";
708 let credential = Credential::bearer(SECRET);
709
710 let debug = format!("{credential:?}");
711 let display = format!("{credential}");
712 assert_eq!(debug, "Credential(<redacted>)");
713 assert_eq!(display, "Credential(<redacted>)");
714 assert!(
715 !debug.contains(SECRET),
716 "Debug must not leak the secret (C8)"
717 );
718 assert!(
719 !display.contains(SECRET),
720 "Display must not leak the secret (C8)"
721 );
722 assert_eq!(
724 format!("{:?}", credential.clone()),
725 "Credential(<redacted>)"
726 );
727
728 let insecure = HostError::InsecureTransport {
732 id: "remote".into(),
733 host: "example.com".into(),
734 };
735 let unauthorized = HostError::Unauthorized {
736 id: "remote".into(),
737 };
738 assert!(!insecure.to_string().contains(SECRET));
739 assert!(!unauthorized.to_string().contains(SECRET));
740 }
741}