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
6 changes: 3 additions & 3 deletions crates/skilld-command/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -23,10 +23,10 @@ pub use local_store::{
pub use output::{CommandPlatform, OutputContext};
pub use remote::{
Cancellation, HeaderValue, HttpAdapter, HttpHeader, HttpMethod, HttpRequest, HttpResponse,
NativeRemoteConfig, NeverCancelled, NoTokenProvider, PreparedRemoteSkill,
NativeRemoteConfig, NeverCancelled, NoRemoteProgress, NoTokenProvider, PreparedRemoteSkill,
RemoteComparisonAccess, RemoteComparisonOutcome, RemoteComparisonRelation, RemoteLatestCommit,
RemoteProvider, RemoteSourceState, RemoteUpdateComparison, RemoteUpdateResult, SecretValue,
SkilldRemote, Sleeper, ThreadSleeper, TokenProvider,
RemoteProgress, RemoteProgressStage, RemoteProvider, RemoteSourceState, RemoteUpdateComparison,
RemoteUpdateResult, SecretValue, SkilldRemote, Sleeper, ThreadSleeper, TokenProvider,
};
pub use run::{
FileContent, FileKind, PulledFile, RunOutcome, SkillOrigin, SupportingFile, TransientSkill,
Expand Down
61 changes: 58 additions & 3 deletions crates/skilld-command/src/remote.rs
Original file line number Diff line number Diff line change
Expand Up @@ -649,11 +649,50 @@ pub enum NativeRemoteConfig {
Unconfigured,
}

#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq)]
#[serde(rename_all = "kebab-case")]
pub enum RemoteProgressStage {
RequestingResolution,
Requested,
Resolving,
Fetching,
Checking,
Packaging,
Encrypting,
Signing,
Publishing,
RetryWait,
VerifyingAttestation,
RequestingDownload,
DownloadingArtifact,
VerifyingArtifact,
}

pub trait RemoteProgress: Send + Sync {
fn stage(&self, stage: RemoteProgressStage);
}

#[derive(Debug, Deserialize)]
#[serde(untagged)]
enum RemotePendingStage {
Known(RemoteProgressStage),
#[allow(dead_code)]
Unknown(String),
}

#[derive(Default)]
pub struct NoRemoteProgress;

impl RemoteProgress for NoRemoteProgress {
fn stage(&self, _stage: RemoteProgressStage) {}
}

pub struct SkilldRemote {
adapter: Arc<dyn HttpAdapter>,
tokens: Arc<dyn TokenProvider>,
cancellation: Arc<dyn Cancellation>,
sleeper: Arc<dyn Sleeper>,
progress: Arc<dyn RemoteProgress>,
endpoint: Url,
root_pin: NativeRemoteConfig,
}
Expand All @@ -669,6 +708,7 @@ impl SkilldRemote {
tokens,
cancellation: Arc::new(NeverCancelled),
sleeper: Arc::new(ThreadSleeper),
progress: Arc::new(NoRemoteProgress),
endpoint: Url::parse("https://skilld.dev").expect("the fixed endpoint is valid"),
root_pin,
}
Expand All @@ -684,6 +724,11 @@ impl SkilldRemote {
self
}

pub fn with_progress(mut self, progress: Arc<dyn RemoteProgress>) -> Self {
self.progress = progress;
self
}

pub fn with_endpoint(mut self, endpoint: &str) -> Result<Self, RemoteError> {
let endpoint = Url::parse(endpoint)
.map_err(|_| RemoteError::new("INVALID_ENDPOINT", "the API endpoint is invalid"))?;
Expand Down Expand Up @@ -843,6 +888,8 @@ impl SkilldRemote {

fn resolve(&self, source: &SourceRequest) -> Result<ArtifactDescriptor, RemoteError> {
let mut deadline = ResolutionDeadline::new();
self.progress
.stage(RemoteProgressStage::RequestingResolution);
let body = serde_json::to_vec(&json!({ "source": source })).map_err(|_| {
RemoteError::new("INVALID_SOURCE", "the source request cannot be encoded")
})?;
Expand Down Expand Up @@ -891,9 +938,12 @@ impl SkilldRemote {
}
Resolution::Pending {
resolution_id,
stage,
poll_after_ms,
..
} => {
if let RemotePendingStage::Known(stage) = stage {
self.progress.stage(stage);
}
if !(250..=60_000).contains(&poll_after_ms) {
return Err(RemoteError::new(
"INVALID_RESPONSE",
Expand Down Expand Up @@ -1468,10 +1518,16 @@ impl RemoteProvider for SkilldRemote {
return self.direct(selector);
}
let descriptor = self.resolve(selector.source())?;
self.progress
.stage(RemoteProgressStage::VerifyingAttestation);
let root = self.verified_root()?;
verify_attestation(&descriptor.attestation, &root)?;
self.progress.stage(RemoteProgressStage::RequestingDownload);
let grant = self.grant(&descriptor.artifact_id)?;
self.progress
.stage(RemoteProgressStage::DownloadingArtifact);
let archive = self.download_grant(&descriptor, grant)?;
self.progress.stage(RemoteProgressStage::VerifyingArtifact);
let verified = verify_artifact(descriptor.attestation, &root, &archive)?;
if matches!(
&selector.source().selector,
Expand Down Expand Up @@ -2261,8 +2317,7 @@ enum Resolution {
Pending {
#[serde(rename = "resolutionId")]
resolution_id: String,
#[serde(rename = "stage")]
_stage: String,
stage: RemotePendingStage,
#[serde(rename = "pollAfterMs")]
poll_after_ms: u64,
},
Expand Down
119 changes: 117 additions & 2 deletions crates/skilld-command/tests/remote.rs
Original file line number Diff line number Diff line change
Expand Up @@ -11,8 +11,9 @@ use sha2::{Digest, Sha256};
use skilld_command::{
Cancellation, HeaderValue, Host, HttpAdapter, HttpRequest, HttpResponse, LocalHost,
NativeRemoteConfig, NoTokenProvider, PreparedRemoteSkill, RemoteComparisonAccess,
RemoteComparisonOutcome, RemoteComparisonRelation, RemoteProvider, RemoteSourceState,
RemoteUpdateComparison, SecretValue, SkilldRemote, Sleeper, TokenProvider, run,
RemoteComparisonOutcome, RemoteComparisonRelation, RemoteProgress, RemoteProgressStage,
RemoteProvider, RemoteSourceState, RemoteUpdateComparison, SecretValue, SkilldRemote, Sleeper,
TokenProvider, run,
};
use skilld_core::{
AgentTargetId, ArtifactAttestation, ArtifactFile, AttestationSignature, CheckOutcome,
Expand Down Expand Up @@ -141,6 +142,15 @@ impl Sleeper for NoSleep {
}
}

#[derive(Default)]
struct RecordingProgress(Mutex<Vec<RemoteProgressStage>>);

impl RemoteProgress for RecordingProgress {
fn stage(&self, stage: RemoteProgressStage) {
self.0.lock().unwrap().push(stage);
}
}

#[derive(Default)]
struct RecordingSleeper {
elapsed: Mutex<Duration>,
Expand Down Expand Up @@ -1020,6 +1030,111 @@ fn a_resolution_cannot_change_its_identity_while_polling() {
assert_eq!(error.code, "INVALID_RESPONSE");
}

#[test]
fn a_hosted_resolution_reports_each_service_stage() {
let resolution_id = "018f47a4-2d38-7c5f-8d3e-1c5a6b7d8e9f";
let stages = [
"requested",
"resolving",
"fetching",
"checking",
"packaging",
"encrypting",
"signing",
"publishing",
"retry-wait",
];
let mut responses = stages
.iter()
.map(|stage| {
response(
200,
serde_json::to_vec(&json!({
"state": "pending",
"resolutionId": resolution_id,
"stage": stage,
"pollAfterMs": 250
}))
.unwrap(),
)
})
.collect::<Vec<_>>();
responses.push(response(
200,
serde_json::to_vec(&json!({
"state": "blocked",
"resolutionId": resolution_id,
"checkResults": []
}))
.unwrap(),
));
let progress = Arc::new(RecordingProgress::default());
let remote = search_remote(Arc::new(FakeHttp::with(responses))).with_progress(progress.clone());

let error = remote.prepare(&skilld_selector(), false).unwrap_err();

assert_eq!(error.code, "CHECK_BLOCKED");
assert_eq!(
*progress.0.lock().unwrap(),
[
RemoteProgressStage::RequestingResolution,
RemoteProgressStage::Requested,
RemoteProgressStage::Resolving,
RemoteProgressStage::Fetching,
RemoteProgressStage::Checking,
RemoteProgressStage::Packaging,
RemoteProgressStage::Encrypting,
RemoteProgressStage::Signing,
RemoteProgressStage::Publishing,
RemoteProgressStage::RetryWait,
]
);
}

#[test]
fn an_unknown_pending_stage_does_not_abort_the_resolution() {
let (pin, mut responses) = verified_remote_responses();
let pending = response(
200,
serde_json::to_vec(&json!({
"state": "pending",
"resolutionId": "018f47a4-2d38-7c5f-8d3e-1c5a6b7d8e9f",
"stage": "queueing",
"pollAfterMs": 250
}))
.unwrap(),
);
responses.insert(0, pending);
let progress = Arc::new(RecordingProgress::default());
let remote = SkilldRemote::new(
Arc::new(FakeHttp::with(responses)),
Arc::new(NoTokenProvider),
NativeRemoteConfig::Pinned(pin),
)
.with_endpoint("http://127.0.0.1:8787")
.unwrap()
.with_sleeper(Arc::new(NoSleep))
.with_progress(progress.clone());
let selector = RemoteSelector::parse("skilld:skilld-dev/skills/example").unwrap();

let prepared = remote.prepare(&selector, false).unwrap();

assert!(matches!(
prepared.source_status,
SourceStatus::Verified { .. }
));
assert_eq!(
*progress.0.lock().unwrap(),
[
RemoteProgressStage::RequestingResolution,
RemoteProgressStage::VerifyingAttestation,
RemoteProgressStage::RequestingDownload,
RemoteProgressStage::DownloadingArtifact,
RemoteProgressStage::VerifyingArtifact,
]
);
}

#[test]
fn a_pending_resolution_times_out_after_at_most_sixty_seconds() {
let resolution_id = "018f47a4-2d38-7c5f-8d3e-1c5a6b7d8e9f";
Expand Down
46 changes: 27 additions & 19 deletions crates/skilld-native/src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -60,17 +60,39 @@ fn main() -> ExitCode {
};
let global_root = global_root();
let detection = detection_environment();
let output = OutputContext::auto(
std::io::stdout().is_terminal(),
active_agent_detected(),
environment_enabled("CI"),
environment_present("NO_COLOR"),
env::var("TERM").is_ok_and(|term| term.eq_ignore_ascii_case("dumb")),
terminal_width(),
CommandPlatform::current(),
);
let label = if interactive {
None
} else {
status::status_label(args.iter().map(|arg| arg.to_string_lossy()))
};
let status = match label {
Some(label) => StatusLine::for_terminal(label, output),
None => StatusLine::disabled(),
};
let remote_progress = status.remote_progress();
let account = Arc::new(NativeAccount::new());
let host = LocalHost::new(project_root, global_root)
.with_target_roots(target_roots())
.with_detection_environment(detection.clone())
.with_bundled_provider(Arc::new(EmbeddedSkilld::new()))
.with_account_provider(account.clone())
.with_remote_provider(Arc::new(SkilldRemote::new(
Arc::new(NativeHttpAdapter::new()),
account,
native_remote_config(),
)));
.with_remote_provider(Arc::new(
SkilldRemote::new(
Arc::new(NativeHttpAdapter::new()),
account,
native_remote_config(),
)
.with_progress(remote_progress),
));
let host = if args.iter().skip(1).any(|arg| arg == "outdated")
&& !args.iter().any(|arg| arg == "--json" || arg == "--plain")
{
Expand Down Expand Up @@ -115,20 +137,6 @@ fn main() -> ExitCode {

let mut stdout = std::io::stdout().lock();
let mut stderr = std::io::stderr();
let output = OutputContext::auto(
stdout.is_terminal(),
active_agent_detected(),
environment_enabled("CI"),
environment_present("NO_COLOR"),
env::var("TERM").is_ok_and(|term| term.eq_ignore_ascii_case("dumb")),
terminal_width(),
CommandPlatform::current(),
);
let label = status::status_label(args.iter().map(|arg| arg.to_string_lossy()));
let status = match label {
Some(label) => StatusLine::for_terminal(label, output),
None => StatusLine::disabled(),
};
let mut gated = status::GatedStderr::new(&mut stderr, status);
let result = run_with_output(args, host.as_ref(), output, &mut stdout, &mut gated);
gated.finish_status();
Expand Down
Loading