Skip to content
Draft
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: 14 additions & 5 deletions crates/now-package-broker/src/auth.rs
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,10 @@ impl PipeClient {
/// unauthenticated work a connection flood can trigger.
pub(crate) fn from_connected_pipe(server: &NamedPipeServer) -> anyhow::Result<Self> {
let process_id = connected_pipe_client_process_id(server).context("failed to query pipe client process id")?;
Self::from_process_id(process_id)
}

fn from_process_id(process_id: u32) -> anyhow::Result<Self> {
let process = Process::get_by_pid(process_id, PROCESS_QUERY_LIMITED_INFORMATION)
.with_context(|| format!("failed to open pipe client process {process_id}"))?;
let executable_path = process
Expand All @@ -50,6 +54,11 @@ impl PipeClient {
})
}

#[cfg(test)]
pub(crate) fn from_current_process() -> anyhow::Result<Self> {
Self::from_process_id(std::process::id())
}

/// Security identifier of the authenticated pipe client user, captured at connect.
pub(crate) fn user_sid(&self) -> &Sid {
&self.user_sid
Expand All @@ -61,7 +70,7 @@ impl PipeClient {
skip_signature_validation: bool,
) -> anyhow::Result<()> {
self.validate_client_context(&request.client)?;
self.validate_signature(skip_signature_validation)
self.validate_connection(skip_signature_validation)
}

pub(crate) fn validate_status_request(
Expand All @@ -70,7 +79,7 @@ impl PipeClient {
skip_signature_validation: bool,
) -> anyhow::Result<()> {
self.validate_client_context(&request.client)?;
self.validate_signature(skip_signature_validation)
self.validate_connection(skip_signature_validation)
}

pub(crate) fn validate_cancel_request(
Expand All @@ -79,15 +88,15 @@ impl PipeClient {
skip_signature_validation: bool,
) -> anyhow::Result<()> {
self.validate_client_context(&request.client)?;
self.validate_signature(skip_signature_validation)
self.validate_connection(skip_signature_validation)
}

fn validate_client_context(&self, client: &ClientContext) -> anyhow::Result<()> {
self.validate_effective_user(&client.effective_user)?;
self.validate_executable_path(&client.client_executable_path)
}

fn validate_signature(&self, skip_signature_validation: bool) -> anyhow::Result<()> {
pub(crate) fn validate_connection(&self, skip_signature_validation: bool) -> anyhow::Result<()> {
if signature_validation_skipped(skip_signature_validation) {
warn!("DEBUG MODE: Skipping package broker client signature validation");
return Ok(());
Expand Down Expand Up @@ -364,7 +373,7 @@ mod tests {
user_sid: client_user_sid(),
};

assert!(client.validate_signature(true).is_err());
assert!(client.validate_connection(true).is_err());
}
}

Expand Down
238 changes: 224 additions & 14 deletions crates/now-package-broker/src/server/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -11,8 +11,8 @@ use now_policy_api::{
CancelRequest, CancelResponse, CancelResponseKind, CapabilitiesResponse, CapabilitiesResponseKind, Decision,
DecisionInfo, Elevation, ErrorCode, ErrorResponse, EvaluationResponse, EvaluationResponseKind, ExecutionResponse,
ExecutionResponseKind, HealthResponse, HealthResponseKind, HealthStatus, ManagerCapability, ManagerName,
OperationStatus, OperationSubmission, PackageRequest, Scope, StatusRequest, StatusResponse, StatusResponseKind,
Transport,
OperationStatus, OperationSubmission, PackageRequest, PolicyResponse, PolicyResponseKind, Scope, StatusRequest,
StatusResponse, StatusResponseKind, Transport,
};
use now_policy_server_template::{MAX_REQUEST_BODY_BYTES, PackageBrokerServer, SharedPackageBrokerServer};
use tracing::{info, trace, warn};
Expand Down Expand Up @@ -115,6 +115,17 @@ impl PackageBrokerServer for BrokerConnection {
self.state.capabilities(self.client.user_sid()).await
}

async fn policy(&self) -> Result<PolicyResponse, ErrorResponse> {
self.client
.validate_connection(self.state.skip_signature_validation)
.map_err(|error| {
warn!(error = format!("{error:#}"), "Rejected package broker policy request");
error_response(ErrorCode::Unauthorized, "pipe client authentication failed")
})?;

self.state.policy_response()
}

async fn evaluate(&self, request: PackageRequest) -> Result<EvaluationResponse, ErrorResponse> {
self.client
.validate_request(&request, self.state.skip_signature_validation)
Expand Down Expand Up @@ -163,6 +174,25 @@ impl PackageBrokerServer for BrokerConnection {
}

impl BrokerState {
fn active_policy(&self) -> Result<Arc<PolicyDocument>, ErrorResponse> {
let guard = self.policy.read().expect("policy lock poisoned");
guard
.as_ref()
.map(Arc::clone)
.ok_or_else(|| error_response(ErrorCode::BrokerPaused, "active policy is unavailable"))
}

fn policy_response(&self) -> Result<PolicyResponse, ErrorResponse> {
let policy = self.active_policy()?;

Ok(PolicyResponse {
response_kind: PolicyResponseKind,
response_version: api_version(),
server: server_context(),
policy: (*policy).clone(),
})
}

async fn health(&self) -> HealthResponse {
let policy_guard = self.policy.read().expect("policy lock poisoned");
let (status, policy_id) = match policy_guard.as_ref() {
Expand Down Expand Up @@ -419,18 +449,7 @@ impl BrokerState {
}

let received_at = Utc::now();
let policy = {
let guard = self.policy.read().expect("policy lock poisoned");
match guard.as_ref() {
Some(policy) => Arc::clone(policy),
None => {
return Err(error_response(
ErrorCode::BrokerPaused,
"policy file is unavailable or corrupted",
));
}
}
};
let policy = self.active_policy()?;

if let Some(reason) = policy_validity_failure(&policy, received_at) {
warn!(%reason, "Rejecting request: policy outside validity window");
Expand Down Expand Up @@ -521,12 +540,15 @@ mod tests {

use std::sync::atomic::{AtomicUsize, Ordering};

use axum::body::{Body, to_bytes};
use axum::http::{Method, Request, StatusCode};
use chrono::Utc;
use now_policy::{
PackageBrokerPolicy, PolicyEnforcement, PolicyMetadata, PolicySchemaUri, ResourceId, RulePrecedence,
SemanticVersion,
};
use now_policy_api as api;
use tower_service::Service as _;

use super::*;
use crate::executor::{ExecutionOutput, OperationCanceled, ProcessStartedCallback};
Expand Down Expand Up @@ -601,6 +623,194 @@ mod tests {
}
}

fn shared_state(policy: Option<PolicyDocument>) -> Arc<BrokerState> {
let mut state = state();
state.policy = RwLock::new(policy.map(Arc::new));
Arc::new(state)
}

async fn route_request(state: Arc<BrokerState>, method: Method, uri: &str) -> axum::response::Response {
let client = PipeClient::from_current_process().expect("capture current test process");
let mut router = build_router_for_client(state, client);
router
.call(
Request::builder()
.method(method)
.uri(uri)
.body(Body::empty())
.expect("valid test request"),
)
.await
.expect("router is infallible")
}

async fn response_json(response: axum::response::Response) -> serde_json::Value {
let body = to_bytes(response.into_body(), usize::MAX)
.await
.expect("read response body");
serde_json::from_slice(&body).expect("response is valid JSON")
}

#[cfg(feature = "dev-skip-broker-signature")]
#[tokio::test]
async fn policy_route_serializes_active_policy_with_empty_rules() {
let expected = permissive_policy();
let response = route_request(shared_state(Some(expected.clone())), Method::GET, "/v1/policy").await;

assert_eq!(response.status(), StatusCode::OK);
assert_eq!(response.headers().get("content-type").unwrap(), "application/json");

let response: PolicyResponse =
serde_json::from_value(response_json(response).await).expect("deserialize policy response");
assert_eq!(response.response_kind, PolicyResponseKind);
assert_eq!(&*response.response_version, api::API_VERSION_STR);
assert_eq!(response.server.transport, Transport::HttpNamedPipe);
assert_eq!(
serde_json::to_value(response.policy).unwrap(),
serde_json::to_value(expected).unwrap()
);
}

#[cfg(feature = "dev-skip-broker-signature")]
#[tokio::test]
async fn policy_route_serializes_full_policy_matches_and_constraints() {
let expected =
now_policy::schema::parse_policy_json(include_str!("../assets/samples/corporate-allowlist.policy.json"))
.expect("sample policy is valid");
let response = route_request(shared_state(Some(expected.clone())), Method::GET, "/v1/policy").await;

assert_eq!(response.status(), StatusCode::OK);

let response: PolicyResponse =
serde_json::from_value(response_json(response).await).expect("deserialize policy response");
assert_eq!(
serde_json::to_value(response.policy).unwrap(),
serde_json::to_value(expected).unwrap()
);
}

#[cfg(feature = "dev-skip-broker-signature")]
#[tokio::test]
async fn policy_route_returns_structured_service_unavailable_without_active_policy() {
let response = route_request(shared_state(None), Method::GET, "/v1/policy").await;

assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE);

let body = response_json(response).await;
let error: ErrorResponse = serde_json::from_value(body.clone()).expect("deserialize error response");
assert_eq!(error.code, ErrorCode::BrokerPaused);
assert_eq!(error.message, "active policy is unavailable");
assert!(error.details.is_empty());
assert!(body.get("Policy").is_none());
}

#[cfg(not(feature = "dev-skip-broker-signature"))]
#[tokio::test]
async fn policy_route_rejects_unsigned_client() {
let response = route_request(shared_state(Some(permissive_policy())), Method::GET, "/v1/policy").await;

assert_eq!(response.status(), StatusCode::UNAUTHORIZED);

let body = response_json(response).await;
let error: ErrorResponse = serde_json::from_value(body.clone()).expect("deserialize error response");
assert_eq!(error.code, ErrorCode::Unauthorized);
assert_eq!(error.message, "pipe client authentication failed");
assert!(body.get("Policy").is_none());
}

#[cfg(feature = "dev-skip-broker-signature")]
#[tokio::test]
async fn policy_route_preserves_existing_routes_and_method_restrictions() {
let state = shared_state(Some(permissive_policy()));

for uri in ["/v1/health", "/v1/capabilities"] {
let response = route_request(Arc::clone(&state), Method::GET, uri).await;
assert_eq!(response.status(), StatusCode::OK, "unexpected status for {uri}");
}

let response = route_request(Arc::clone(&state), Method::HEAD, "/v1/policy").await;
assert_eq!(response.status(), StatusCode::OK);
assert!(
to_bytes(response.into_body(), usize::MAX)
.await
.expect("read HEAD response")
.is_empty()
);

for method in [
Method::POST,
Method::PUT,
Method::PATCH,
Method::DELETE,
Method::OPTIONS,
Method::TRACE,
Method::CONNECT,
] {
let response = route_request(Arc::clone(&state), method.clone(), "/v1/policy").await;
assert_eq!(
response.status(),
StatusCode::METHOD_NOT_ALLOWED,
"unexpected status for {method}"
);
}

let response = route_request(state, Method::GET, "/v1/not-a-route").await;
assert_eq!(response.status(), StatusCode::NOT_FOUND);
}

#[test]
fn concurrent_policy_replacement_returns_only_complete_snapshots() {
let policy_a = permissive_policy();
let mut policy_b =
now_policy::schema::parse_policy_json(include_str!("../assets/samples/corporate-allowlist.policy.json"))
.expect("sample policy is valid");
policy_b.metadata.id = ResourceId::from("replacement-policy");
policy_b.metadata.revision = 42;

let current_policy_json = serde_json::to_value(&policy_a).unwrap();
let replacement_policy_json = serde_json::to_value(&policy_b).unwrap();
let policy_a = Arc::new(policy_a);
let policy_b = Arc::new(policy_b);
let state = shared_state(None);
*state.policy.write().expect("policy lock") = Some(Arc::clone(&policy_a));

const READER_COUNT: usize = 4;
const ITERATIONS: usize = 1_000;
let barrier = Arc::new(std::sync::Barrier::new(READER_COUNT + 1));

std::thread::scope(|scope| {
for _ in 0..READER_COUNT {
let state = Arc::clone(&state);
let barrier = Arc::clone(&barrier);
let current_policy_json = &current_policy_json;
let replacement_policy_json = &replacement_policy_json;
scope.spawn(move || {
barrier.wait();
for _ in 0..ITERATIONS {
let response = state.policy_response().expect("active policy response");
let actual = serde_json::to_value(response.policy).unwrap();
assert!(
actual == *current_policy_json || actual == *replacement_policy_json,
"response mixed two policy snapshots"
);
std::thread::yield_now();
}
});
}

barrier.wait();
for index in 0..ITERATIONS {
let replacement = if index % 2 == 0 {
Arc::clone(&policy_b)
} else {
Arc::clone(&policy_a)
};
*state.policy.write().expect("policy lock") = Some(replacement);
std::thread::yield_now();
}
});
}

fn request() -> PackageRequest {
PackageRequest {
request_kind: api::PackageRequestKind,
Expand Down
Loading