Skip to content
Open
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
119 changes: 113 additions & 6 deletions crates/openshell-supervisor/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -653,7 +653,7 @@ pub async fn run_sandbox(
loaded_policy_origin,
initial_agent_proposals_enabled,
initial_extension_authentication_enabled,
captured_provider_credentials,
captured_provider_environment,
) = load_policy_with_gateway(
sandbox_id.clone(),
sandbox.clone(),
Expand All @@ -678,8 +678,8 @@ pub async fn run_sandbox(
let workspace = workdir;

let provider_readiness = ProviderReadinessTracker::new();
let provider_credentials = if let Some(credentials) = captured_provider_credentials {
credentials
let provider_credentials = if let Some(environment) = captured_provider_environment {
environment.install(&provider_readiness)
} else {
// Fetch provider environment variables from the server.
// This is done after loading the policy so the sandbox can still start
Expand Down Expand Up @@ -1987,6 +1987,35 @@ enum LocalPolicyIdentity {
EndpointOnly,
}

struct CapturedProviderEnvironment {
credentials: ProviderCredentialState,
expires_at_ms: Option<i64>,
identity: EnvironmentIdentity,
}

impl CapturedProviderEnvironment {
fn install(self, readiness: &ProviderReadinessTracker) -> ProviderCredentialState {
readiness.credentials_installed(self.identity, &self.credentials, self.expires_at_ms);
self.credentials
}

fn new(
credentials: ProviderCredentialState,
provider: &openshell_core::grpc_client::ProviderEnvironmentResult,
) -> Self {
Self {
credentials,
expires_at_ms: provider
.credential_expires_at_ms
.values()
.copied()
.filter(|expiry| *expiry > 0)
.min(),
identity: EnvironmentIdentity::from_environment(provider),
}
}
}

async fn load_policy(
sandbox_id: Option<String>,
sandbox: Option<String>,
Expand All @@ -2003,7 +2032,7 @@ async fn load_policy(
LoadedPolicyOrigin,
bool,
bool,
Option<ProviderCredentialState>,
Option<CapturedProviderEnvironment>,
)> {
load_policy_with_gateway(
sandbox_id,
Expand Down Expand Up @@ -2043,7 +2072,7 @@ async fn load_policy_with_gateway(
LoadedPolicyOrigin,
bool,
bool,
Option<ProviderCredentialState>,
Option<CapturedProviderEnvironment>,
)> {
use openshell_core::proto::ConfigurationAdmissionState;
// File mode: load OPA engine from rego rules + YAML data (dev override)
Expand Down Expand Up @@ -2406,7 +2435,10 @@ async fn load_policy_with_gateway(
},
agent_proposals_enabled_from_settings(&snapshot.settings),
snapshot.extension_authentication_enabled,
Some(captured_provider_credentials),
Some(CapturedProviderEnvironment::new(
captured_provider_credentials,
&provider,
)),
));
}
}
Expand Down Expand Up @@ -5443,6 +5475,22 @@ network_policies:
assert_eq!(credentials.revision(), 10);
}

#[test]
fn startup_environment_seeds_provider_readiness() {
let mut provider = startup_provider(10);
provider.provider_attachment_epoch = "epoch".to_string();
provider.policy_hash = "policy".to_string();
let identity = EnvironmentIdentity::from_environment(&provider);
let credentials = prepare_provider_environment(&provider).unwrap();
let readiness = ProviderReadinessTracker::new();

let credentials =
CapturedProviderEnvironment::new(credentials, &provider).install(&readiness);

assert_eq!(credentials.revision(), 10);
assert!(!readiness.needs_environment(&identity));
}

#[test]
fn startup_configuration_accepts_fail_closed_provider_environment() {
let policy = proto_policy_fixture();
Expand Down Expand Up @@ -6080,6 +6128,65 @@ network_policies:
let _ = task.await;
}

#[tokio::test]
async fn provider_readiness_startup_environment_avoids_unchanged_refresh() {
let policy = proto_policy_fixture();
let mut settings = settings_poll_result(
Some(policy.clone()),
1,
openshell_core::proto::PolicySource::Sandbox,
);
settings.provider_env_revision = 6;
let engine = Arc::new(OpaEngine::from_proto(&policy).unwrap());
let mut ctx = policy_poll_test_context(
engine.clone(),
LoadedPolicyOrigin::Gateway {
revision: Some(LoadedPolicyRevision::from_snapshot(&settings)),
has_last_valid_policy: true,
},
default_middleware_connector(),
);
let mut provider = static_provider_environment(6, Some("initial"));
provider.policy_hash.clone_from(&settings.policy_hash);
let credentials = prepare_provider_environment(&provider).unwrap();
ctx.provider_credentials = CapturedProviderEnvironment::new(credentials, &provider)
.install(&ctx.provider_readiness);
let generation = engine.current_generation();
let guard = engine.generation_guard(generation).unwrap();
let (policy_gateway, polls, mut reports) = scripted_policy_gateway();
let observed_polls = policy_gateway.polled_sandboxes.clone();
let (requests, mut received) = tokio::sync::mpsc::unbounded_channel();
let task = tokio::spawn(run_policy_poll_loop_with_client(
ctx,
ScriptedProviderGateway {
policy: policy_gateway,
requests,
},
));

polls.send(settings.clone()).unwrap();
expect_policy_report(&mut reports, 1).await;
polls.send(settings).unwrap();
timeout(Duration::from_secs(1), async {
while observed_polls.lock().await.len() < 2 {
tokio::task::yield_now().await;
}
})
.await
.unwrap();

assert!(
timeout(Duration::from_millis(50), received.recv())
.await
.is_err(),
"unchanged settings must not refetch the startup environment"
);
assert_eq!(engine.current_generation(), generation);
assert!(!guard.is_stale());
task.abort();
let _ = task.await;
}

#[tokio::test]
async fn provider_poll_installs_fail_closed_environment_and_acknowledges_policy() {
let policy = proto_policy_fixture();
Expand Down
Loading