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
19 changes: 18 additions & 1 deletion Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

2 changes: 1 addition & 1 deletion Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ members = [
"crates/persisting-cli",
"crates/persisting-dlcapt",
"crates/persisting-compute",
"crates/persisting-pvisor",
]

[workspace.package]
Expand All @@ -33,4 +34,3 @@ codegen-units = 1

[workspace.dependencies]
pulsing-actor = { version = "0.1.2", default-features = false }

3 changes: 2 additions & 1 deletion crates/persisting-capture/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,8 @@ pulsing-actor = { workspace = true }
dashmap = "6"
blake3 = "1"
hostname = "0.4"
ipnet = "2"
persisting-proto = { path = "../persisting-proto" }
persisting-pvisor = { path = "../persisting-pvisor" }

[dev-dependencies]
tempfile = "3"
Expand Down
25 changes: 25 additions & 0 deletions crates/persisting-capture/src/proxy/common.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ use std::sync::Arc;

use axum::http::HeaderMap;
use bytes::Bytes;
use persisting_proto::ModelAccessPolicy;
use serde_json::Value;

use super::state::ProxyState;
Expand All @@ -22,6 +23,30 @@ pub(crate) fn effective_config(state: &ProxyState, route: &CaptureRoute) -> Arc<
.unwrap_or_else(|| Arc::clone(&state.config))
}

pub(crate) fn model_access_policy(config: &ProxyConfig) -> ModelAccessPolicy {
let allowed_models = config
.models
.iter()
.map(|route| route.name.clone())
.collect();
let providers: Vec<String> = config
.models
.iter()
.filter_map(|route| route.provider.clone())
.collect();
// An inferred/custom provider must remain representable during migration.
// An empty provider list means model identity is enforced but provider is open.
let allowed_providers = if providers.len() == config.models.len() {
providers
} else {
Vec::new()
};
ModelAccessPolicy {
allowed_models,
allowed_providers,
}
}

#[allow(clippy::too_many_arguments)]
pub(crate) fn call_context(
state: &ProxyState,
Expand Down
41 changes: 35 additions & 6 deletions crates/persisting-capture/src/proxy/dispatch.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,14 +10,16 @@ use axum::routing::any;
use axum::Router;
use bytes::Bytes;
use http_body_util::BodyExt;
use persisting_proto::{NetworkAccessRequest, NetworkTransport, RunId, StorylineId};

use super::common::effective_config;
use super::forward::{
handle_connect, is_forward_proxy_request, is_llm_capture_path, transparent_forward,
handle_connect_authorized, is_forward_proxy_request, is_llm_capture_path,
transparent_forward_authorized,
};
use super::llm_capture::llm_capture;
use super::network_policy::{
assert_egress_allowed, forbidden_response, host_from_authority, NetworkPolicy,
authorize_egress, forbidden_response, host_from_authority, NetworkPolicy,
};
use super::state::ProxyState;
use crate::debug::{self, is_debug_enabled};
Expand Down Expand Up @@ -98,7 +100,18 @@ async fn dispatch_impl(
.map(|a| a.to_string())
.unwrap_or_else(|| uri.clone());
let host = host_from_authority(&authority);
if let Err(reason) = assert_egress_allowed(&policy, &host) {
if let Err(reason) = authorize_egress(
state.access_controller.as_ref(),
&policy,
&NetworkAccessRequest {
run_id: log_route.root_session.clone().map(RunId),
attempt_id: None,
storyline_id: Some(StorylineId(log_route.session_id.clone())),
host: host.clone(),
port: req.uri().port_u16(),
transport: NetworkTransport::TcpTunnel,
},
) {
return Ok(deny_egress(
&state,
&policy,
Expand All @@ -111,13 +124,29 @@ async fn dispatch_impl(
if debug_on {
debug::log_connect(state.storage.as_path(), &authority, &session_id);
}
return Ok(handle_connect(req, &policy).await);
return Ok(handle_connect_authorized(req).await);
}

let path = req.uri().path().to_string();
if is_forward_proxy_request(req.method(), req.uri()) {
let host = req.uri().host().map(str::to_string).unwrap_or_default();
if let Err(reason) = assert_egress_allowed(&policy, &host) {
let transport = if req.uri().scheme_str() == Some("https") {
NetworkTransport::Https
} else {
NetworkTransport::Http
};
if let Err(reason) = authorize_egress(
state.access_controller.as_ref(),
&policy,
&NetworkAccessRequest {
run_id: log_route.root_session.clone().map(RunId),
attempt_id: None,
storyline_id: Some(StorylineId(log_route.session_id.clone())),
host: host.clone(),
port: req.uri().port_u16(),
transport,
},
) {
return Ok(deny_egress(
&state,
&policy,
Expand Down Expand Up @@ -148,7 +177,7 @@ async fn dispatch_impl(
"forward",
);
}
let resp = transparent_forward(&state.client, req, &policy).await?;
let resp = transparent_forward_authorized(&state.client, req).await?;
if debug_on {
let status = resp.status();
let headers = resp.headers().clone();
Expand Down
20 changes: 18 additions & 2 deletions crates/persisting-capture/src/proxy/forward.rs
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,14 @@ pub async fn handle_connect(req: Request, policy: &NetworkPolicy) -> Response {
let (status, msg) = forbidden_response(&host, &reason);
return (status, msg).into_response();
}
handle_connect_authorized(req).await
}

/// Execute a CONNECT request after the caller's pVisor controller authorized it.
pub async fn handle_connect_authorized(req: Request) -> Response {
let Some(authority) = req.uri().authority().map(|a| a.to_string()) else {
return StatusCode::BAD_REQUEST.into_response();
};
let target = connect_target(&authority);
let on_upgrade: OnUpgrade = hyper::upgrade::on(req);
tokio::spawn(async move {
Expand Down Expand Up @@ -85,7 +93,15 @@ pub async fn transparent_forward(
.body(Body::from(msg))
.expect("403 body"));
}
transparent_forward_authorized(client, req).await
}

/// Forward an absolute-URI request after the caller's pVisor controller
/// authorized it.
pub async fn transparent_forward_authorized(
client: &reqwest::Client,
req: Request,
) -> anyhow::Result<Response<Body>> {
let (parts, body) = req.into_parts();
let url = parts.uri.to_string();
let body_bytes = body
Expand Down Expand Up @@ -120,9 +136,9 @@ pub async fn transparent_forward(
}
builder = builder.header(name, value);
}
Ok(builder
builder
.body(Body::from(bytes))
.map_err(|e| anyhow::anyhow!("build response: {e}"))?)
.map_err(|e| anyhow::anyhow!("build response: {e}"))
}

#[cfg(test)]
Expand Down
38 changes: 38 additions & 0 deletions crates/persisting-capture/src/proxy/llm_capture.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,11 +8,13 @@ use axum::extract::Request;
use axum::http::{Method, StatusCode};
use axum::response::{IntoResponse, Response};
use http_body_util::BodyExt;
use persisting_proto::{ModelCallRequest, RunId, StorylineId};
use serde_json::Value;

use super::auth::{apply_upstream_headers, resolve_upstream_api_key};
use super::common::{
attach_capture_headers, call_context, effective_config, extract_model, is_models_list_path,
model_access_policy,
};
use super::http_headers::{is_websocket_upgrade, skip_response_header_when_body_changed};
use super::models_list::build_models_response;
Expand Down Expand Up @@ -134,6 +136,42 @@ pub(super) async fn llm_capture(
upstream_url.set_query(Some(q));
}

let access_decision = state.access_controller.authorize_model(
&model_access_policy(&cfg),
&ModelCallRequest {
run_id: capture_route.root_session.clone().map(RunId::new),
attempt_id: None,
storyline_id: Some(StorylineId::new(capture_route.session_id.clone())),
call_id: call.call_id.clone(),
client_model: client_model.clone(),
upstream_model: upstream_model.clone(),
provider: provider.as_str().to_string(),
protocol: protocol.as_str().to_string(),
upstream_host: upstream_url.host_str().unwrap_or_default().to_string(),
},
);
if !access_decision.is_allowed() {
tracing::warn!(
target: "persisting_capture",
run_id = capture_route.root_session.as_deref().unwrap_or("-"),
storyline_id = %capture_route.session_id,
call_id = %call.call_id,
client_model = %client_model,
upstream_model = %upstream_model,
provider = provider.as_str(),
reason = access_decision.reason.code(),
"pVisor denied model call"
);
return Ok((
StatusCode::FORBIDDEN,
format!(
"persisting-proxy: pVisor denied model `{client_model}` ({})",
access_decision.reason.code()
),
)
.into_response());
}

if debug_on {
let body_preview = truncate_body_bytes(&upstream_body);
debug::log_llm_request(
Expand Down
5 changes: 4 additions & 1 deletion crates/persisting-capture/src/proxy/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -19,5 +19,8 @@ pub mod upstream;
pub use auth::{apply_upstream_headers, resolve_upstream_api_key};
pub use model::rewrite_model_in_body;
pub use reasoning::ReasoningCacheHandle;
pub use state::{serve, serve_with_shutdown, serve_with_shutdown_and_ready, ProxyState};
pub use state::{
serve, serve_with_runtime_control, serve_with_shutdown, serve_with_shutdown_and_ready,
ProxyState,
};
pub use upstream::prepare_upstream_body;
Loading
Loading