Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion crates/machine/src/machine_bus.rs
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ use crate::virtio::blk::VirtioBlk;
use crate::virtio::console::VirtioConsole;
use crate::virtio::fs::{Mount, VirtioFs};
use crate::virtio::net::VirtioNet;
use crate::virtio::secrets::SecretBinding;
use crate::virtio::slirp::SlirpBackend;

use crate::{
Expand Down Expand Up @@ -73,7 +74,11 @@ impl MachineBus {
}

pub fn attach_net(&mut self) {
let backend = SlirpBackend::new(GUEST_MAC);
self.attach_net_with_secrets(Vec::new());
}

pub fn attach_net_with_secrets(&mut self, secrets: Vec<SecretBinding>) {
let backend = SlirpBackend::with_secrets(GUEST_MAC, secrets);
self.net = Some(VirtioNet::new(backend, GUEST_MAC));
}

Expand Down
113 changes: 110 additions & 3 deletions crates/machine/src/virtio/https_gateway.rs
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ use std::collections::VecDeque;
use rustls::ClientConfig;
use std::sync::Arc;

use super::secrets::{SecretBinding, Substitution};
use super::tls_proxy::{Timing, TlsContext, TlsProxy};
use super::upstream::{PREAMBLE_PREFIX, Upstream, UpstreamMode, UpstreamStatus};
use crate::trace::{HttpObserver, Tracer};
Expand All @@ -25,6 +26,7 @@ pub struct HttpsGateway {
dst_ip: [u8; 4],
timing: Option<Timing>,
tracer: Option<Tracer>,
secrets: Vec<SecretBinding>,
}

impl HttpsGateway {
Expand All @@ -36,9 +38,14 @@ impl HttpsGateway {
dst_ip,
timing: Timing::new(),
tracer: None,
secrets: Vec::new(),
}
}

pub fn carry_secrets(&mut self, secrets: Vec<SecretBinding>) {
self.secrets = secrets;
}

pub fn observe_http(&mut self, tracer: Tracer) {
self.tracer = Some(tracer);
}
Expand All @@ -60,6 +67,7 @@ impl HttpsGateway {
dst_ip,
timing: None,
tracer: None,
secrets: Vec::new(),
}
}

Expand Down Expand Up @@ -134,6 +142,8 @@ impl HttpsGateway {
if let Some(tracer) = self.tracer.take() {
proxy.observe_http(HttpObserver::new(tracer, "https", self.dst_ip, 443));
}

proxy.carry_secrets(std::mem::take(&mut self.secrets));
proxy.push_from_guest(&buffered);
self.state = GatewayState::Tls(Box::new(proxy));
}
Expand Down Expand Up @@ -179,6 +189,7 @@ impl HttpsGateway {
port,
&host,
self.timing.take(),
&self.secrets,
) {
Ok(mut bridge) => {
if let Some(tracer) = self.tracer.take() {
Expand Down Expand Up @@ -223,6 +234,7 @@ struct PlainBridge {
upstream_hs_done: bool,
marked_upstream_hs: bool,
first_reply: bool,
substitution: Option<Substitution>,
}

impl PlainBridge {
Expand All @@ -233,6 +245,7 @@ impl PlainBridge {
port: u16,
host: &str,
timing: Option<Timing>,
secrets: &[SecretBinding],
) -> Result<Self, String> {
let upstream = Upstream::connect(mode, client_config, dst_ip, port, host)?;

Expand All @@ -247,6 +260,7 @@ impl PlainBridge {
upstream_hs_done: false,
marked_upstream_hs: false,
first_reply: false,
substitution: Substitution::for_host(secrets, host),
};
if let Some(t) = &bridge.timing {
t.mark("upstream TCP connected");
Expand All @@ -259,7 +273,27 @@ impl PlainBridge {
http.observe(bytes);
}

if self.upstream.send_plaintext(bytes).is_err() {
let outbound = match &mut self.substitution {
None => bytes.to_vec(),
Some(substitution) => match substitution.feed(bytes) {
Ok(bytes) => bytes,
Err(refused) => {
log::warn!(
"secrets: refusing a credential for {}, which is not on its allowlist",
refused.host
);
self.failed = true;
return;
}
},
};

if outbound.is_empty() {
self.pump();
return;
}

if self.upstream.send_plaintext(&outbound).is_err() {
self.failed = true;
return;
}
Expand Down Expand Up @@ -332,8 +366,8 @@ impl PlainBridge {
#[cfg(test)]
mod tests {
use super::super::tls_proxy::{
ca_cert_pem, client_config_trusting, spawn_plaintext_host, spawn_test_upstream,
spawn_test_upstream_rst, spawn_test_upstream_streaming,
ca_cert_pem, client_config_trusting, spawn_capturing_upstream, spawn_plaintext_host,
spawn_test_upstream, spawn_test_upstream_rst, spawn_test_upstream_streaming,
};
use super::*;
use std::time::Duration;
Expand All @@ -347,6 +381,79 @@ mod tests {
}
}

fn secret_for(hosts: &[&str]) -> Vec<SecretBinding> {
vec![SecretBinding {
placeholder: "vpod-secret-key-a1b2c3d4".to_string(),
value: "sk-ant-the-real-thing".to_string(),
hosts: hosts.iter().map(|host| host.to_string()).collect(),
}]
}

#[test]
fn the_preamble_path_swaps_a_credential_too() {
let (port, up_ca, seen, up) = spawn_capturing_upstream(UPSTREAM_REPLY);
let ctx = TlsContext::new().unwrap();
let mut gateway =
HttpsGateway::new_test(&ctx, [127, 0, 0, 1], client_config_trusting(&up_ca));
gateway.carry_secrets(secret_for(&["localhost"]));

let wire = format!(
"VPOD-CONNECT localhost {port}\nGET / HTTP/1.1\r\nHost: localhost\r\n\
x-api-key: vpod-secret-key-a1b2c3d4\r\nContent-Length: 0\r\n\r\n"
);
gateway.push_from_guest(wire.as_bytes());

let mut got = Vec::new();
for _ in 0..2000 {
drain(&mut gateway, &mut got);
if gateway.eof() {
break;
}
std::thread::sleep(Duration::from_millis(1));
}
let _ = up.join();

let received =
String::from_utf8_lossy(&seen.recv().expect("upstream saw no request")).into_owned();

assert!(
received.contains("sk-ant-the-real-thing"),
"the preamble path sent the stand-in: {received}"
);
assert!(!received.contains("vpod-secret-key-a1b2c3d4"));
}

#[test]
fn the_preamble_path_refuses_a_credential_bound_elsewhere() {
let (port, up_ca, seen, up) = spawn_capturing_upstream(UPSTREAM_REPLY);
let ctx = TlsContext::new().unwrap();
let mut gateway =
HttpsGateway::new_test(&ctx, [127, 0, 0, 1], client_config_trusting(&up_ca));
gateway.carry_secrets(secret_for(&["api.anthropic.com"]));

let wire = format!(
"VPOD-CONNECT localhost {port}\nGET / HTTP/1.1\r\nHost: api.anthropic.com\r\n\
x-api-key: vpod-secret-key-a1b2c3d4\r\nContent-Length: 0\r\n\r\n"
);
gateway.push_from_guest(wire.as_bytes());

let mut got = Vec::new();
for _ in 0..200 {
drain(&mut gateway, &mut got);
if gateway.failed() {
break;
}
std::thread::sleep(Duration::from_millis(1));
}
drop(up);

assert!(gateway.failed(), "the connection survived the refusal");
assert!(
seen.try_recv().is_err(),
"a request reached a host outside the allowlist"
);
}

#[test]
fn preamble_bridges_plaintext_to_real_tls_upstream() {
let (port, up_ca, up) = spawn_test_upstream(UPSTREAM_REPLY);
Expand Down
1 change: 1 addition & 0 deletions crates/machine/src/virtio/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ pub mod console;
pub mod fs;
pub mod https_gateway;
pub mod net;
pub mod secrets;
pub mod slirp;
pub mod tls_proxy;
pub mod upstream;
Expand Down
Loading
Loading