diff --git a/src/commands/test/labeler.rs b/src/commands/test/labeler.rs index 98a6ebc..64e3a77 100644 --- a/src/commands/test/labeler.rs +++ b/src/commands/test/labeler.rs @@ -1,12 +1,8 @@ //! `atproto-devtool test labeler ` command. -pub mod create_report; -pub mod crypto; -pub mod http; -pub mod identity; pub mod pipeline; pub mod report; -pub mod subscription; +pub mod target; use std::io; use std::process::ExitCode; @@ -15,13 +11,21 @@ use std::time::Duration; use clap::Args; use miette::Report; -use crate::commands::test::labeler::create_report::self_mint::{SelfMintCurve, SelfMintSigner}; +use self::{ + pipeline::{ + LabelerOptions, + create_report::{ + self, + self_mint::{SelfMintCurve, SelfMintSigner}, + }, + run_pipeline, + }, + report::RenderConfig, +}; use crate::common::{ APP_USER_AGENT, identity::{Did, RealDnsResolver, RealHttpClient, is_local_labeler_hostname}, }; -use pipeline::{LabelerOptions, parse_target, run_pipeline}; -use report::RenderConfig; /// Run the labeler conformance suite against a handle, DID, or endpoint URL. #[derive(Debug, Args)] @@ -90,15 +94,15 @@ impl LabelerCmd { pub async fn run(self, no_color: bool) -> Result { // Parse the target. let target = - parse_target(&self.target, self.did.as_deref()).map_err(|e| miette::miette!("{e}"))?; + target::parse(&self.target, self.did.as_deref()).map_err(|e| miette::miette!("{e}"))?; // Determine tentative endpoint for the locality check. When the target is a // DID or handle, the endpoint is known only after identity stage; for the // self-mint signer construction we need it now. We construct the signer // pessimistically (endpoint unknown) only when --force-self-mint is set. let tentative_endpoint: Option = match &target { - pipeline::LabelerTarget::Endpoint { url, .. } => Some(url.clone()), - pipeline::LabelerTarget::Identified { .. } => None, + target::LabelerTarget::Endpoint { url, .. } => Some(url.clone()), + target::LabelerTarget::Identified { .. } => None, }; let tentative_local = tentative_endpoint diff --git a/src/commands/test/labeler/CLAUDE.md b/src/commands/test/labeler/CLAUDE.md index 9d8e1c7..309f846 100644 --- a/src/commands/test/labeler/CLAUDE.md +++ b/src/commands/test/labeler/CLAUDE.md @@ -17,15 +17,16 @@ integration tests can replay fixtures instead of talking to real servers. - `LabelerCmd::run(no_color) -> Result` (in `labeler.rs`) — constructs the shared reqwest client and calls `pipeline::run_pipeline`. - - `pipeline::parse_target(raw, explicit_did) -> LabelerTarget` — the + - `target::parse(raw, explicit_did) -> LabelerTarget` — the accepted target grammar is handle, `did:*`, `https://` URL, or `http://` URL with a local hostname (loopback, RFC 1918, `.local`). Remote HTTP is rejected with a helpful error; raw endpoints with no DID simply skip identity/crypto. - `pipeline::run_pipeline(target, LabelerOptions) -> LabelerReport` — the one orchestrator that every test hits. -- **Per-stage entry points**: `identity::run`, `http::run`, - `subscription::run`, `crypto::run`, `create_report::run`. Each returns a +- **Per-stage entry points**: `pipeline::identity::run`, `pipeline::http::run`, + `pipeline::subscription::run`, `pipeline::crypto::run`, + `pipeline::create_report::run`. Each returns a `*StageOutput` with an `Option<*Facts>` (populated only when the stage succeeds enough to let downstream stages run, or `None` when there are no meaningful facts to carry forward) plus a `Vec`. diff --git a/src/commands/test/labeler/pipeline.rs b/src/commands/test/labeler/pipeline.rs index f0929d8..ff5a0fa 100644 --- a/src/commands/test/labeler/pipeline.rs +++ b/src/commands/test/labeler/pipeline.rs @@ -3,53 +3,26 @@ use std::borrow::Cow; use std::time::Duration; -use miette::Diagnostic; -use thiserror::Error; use url::Url; -use crate::commands::test::labeler::create_report::self_mint::{SelfMintCurve, SelfMintSigner}; -use crate::commands::test::labeler::create_report::{ - self, CreateReportTee, PdsXrpcClient, RealCreateReportTee, +use super::{ + report::{CheckResult, CheckStatus, LabelerReport, ReportHeader, Stage}, + target::{AtIdentifier, LabelerTarget}, }; -use crate::commands::test::labeler::crypto; -use crate::commands::test::labeler::http::{self, RealHttpTee}; -use crate::commands::test::labeler::identity; -use crate::commands::test::labeler::report::{ - CheckResult, CheckStatus, LabelerReport, ReportHeader, Stage, -}; -use crate::commands::test::labeler::subscription::{self, RealWebSocketClient}; use crate::common::identity::{ - Did, DnsResolver, HttpClient, find_service, is_local_labeler_hostname, resolve_did, - resolve_handle, + Did, DnsResolver, HttpClient, find_service, resolve_did, resolve_handle, }; -/// A labeler target: either a resolvable identifier (handle or DID) or a raw endpoint URL. -#[derive(Debug, Clone)] -pub enum LabelerTarget { - /// A handle or DID that can be resolved. - Identified { - /// The handle or DID to resolve. - identifier: AtIdentifier, - /// An optional explicit DID override (for cross-checking). - explicit_did: Option, - }, - /// A raw HTTP endpoint, optionally with a DID for identity checks. - Endpoint { - /// The endpoint URL. - url: Url, - /// An optional DID to cross-check against the endpoint. - did: Option, - }, -} +pub mod create_report; +pub mod crypto; +pub mod http; +pub mod identity; +pub mod subscription; -/// An ATProto identifier: a handle or a DID. -#[derive(Debug, Clone)] -pub enum AtIdentifier { - /// An ATProto handle (e.g., `alice.bsky.social`). - Handle(String), - /// A decentralized identifier (e.g., `did:plc:...` or `did:web:...`). - Did(Did), -} +use self::create_report::self_mint::{SelfMintCurve, SelfMintSigner}; +use self::create_report::{CreateReportTee, PdsXrpcClient, RealCreateReportTee}; +use self::http::RealHttpTee; +use self::subscription::RealWebSocketClient; /// Options for running the labeler pipeline. pub struct LabelerOptions<'a> { @@ -124,147 +97,13 @@ pub struct PdsCredentials { pub app_password: String, } -/// Error from target parsing. -#[derive(Debug, Error, Diagnostic)] -#[error("{message}")] -pub struct TargetParseError { - /// The error message. - pub message: String, -} - -impl TargetParseError { - fn new(message: impl Into) -> Self { - Self { - message: message.into(), - } - } - - fn unrecognized_target(raw: &str) -> Self { - Self::new(format!( - "Unrecognized target '{raw}'. Expected one of:\n - ATProto handle (e.g., alice.bsky.social)\n - DID (e.g., did:plc:abc123 or did:web:example.com)\n - HTTPS endpoint URL (e.g., https://labeler.example.com)\n - HTTP endpoint URL with a local hostname (e.g., http://localhost:8080)" - )) - } - - fn http_not_supported(raw: &str) -> Self { - Self::new(format!( - "HTTP endpoint '{raw}' is not supported for remote hosts. Use HTTPS, or point at a local labeler (localhost / 127.0.0.0/8 / RFC 1918 / .local) to allow plaintext HTTP." - )) - } - - fn ambiguous_did(raw: &str, explicit: &str) -> Self { - Self::new(format!( - "Ambiguous target specification: target '{raw}' is already a DID, but --did {explicit} was also provided. Please use only one." - )) - } -} - -/// Check if a string is a valid ATProto handle. -/// -/// A valid handle: -/// - Contains at least one dot. -/// - Contains only alphanumeric characters, hyphens, and dots. -/// - Does not start or end with a hyphen or dot. -/// - Has no empty segments (no consecutive dots or leading/trailing dots). -fn is_valid_handle(s: &str) -> bool { - if !s.contains('.') { - return false; - } - - // Check for empty string or leading/trailing special chars. - if s.is_empty() - || s.starts_with('-') - || s.starts_with('.') - || s.ends_with('-') - || s.ends_with('.') - { - return false; - } - - // Check all characters are alphanumeric, hyphen, or dot. - for c in s.chars() { - if !c.is_ascii_alphanumeric() && c != '-' && c != '.' { - return false; - } - } - - // Check no empty segments (no consecutive dots). - if s.contains("..") { - return false; - } - - true -} - -/// Parse a labeler target from a string and optional explicit DID. -/// -/// Returns a `LabelerTarget` on success, or a `TargetParseError` on failure. -/// -/// Rules: -/// - If `raw` starts with `did:`, parse as a DID. If `explicit_did` is also provided, return an error. -/// - If `raw` starts with `https://`, parse as a URL. `explicit_did` is carried as an optional DID. -/// - If `raw` starts with `http://`, return an error pointing the user to HTTPS. -/// - If `raw` contains a dot and matches handle grammar, treat as a handle. `explicit_did` is carried. -/// - Otherwise, return an unrecognized target error. -pub fn parse_target( - raw: &str, - explicit_did: Option<&str>, -) -> Result { - // Check for DID. - if raw.starts_with("did:") { - if let Some(ed) = explicit_did { - return Err(TargetParseError::ambiguous_did(raw, ed)); - } - return Ok(LabelerTarget::Identified { - identifier: AtIdentifier::Did(Did(raw.to_string())), - explicit_did: None, - }); - } - - // Check for HTTPS URL. - if raw.starts_with("https://") { - let url = Url::parse(raw) - .map_err(|e| TargetParseError::new(format!("Invalid URL '{raw}': {e}")))?; - return Ok(LabelerTarget::Endpoint { - url, - did: explicit_did.map(|d| Did(d.to_string())), - }); - } - - // Check for HTTP URL. Local hostnames (loopback, RFC 1918, .local, mDNS) - // are accepted so developers can target a labeler running on their - // machine or LAN. Remote HTTP is still rejected to guard against - // accidental plaintext traffic to a production labeler. - if raw.starts_with("http://") { - let url = Url::parse(raw) - .map_err(|e| TargetParseError::new(format!("Invalid URL '{raw}': {e}")))?; - if is_local_labeler_hostname(&url) { - return Ok(LabelerTarget::Endpoint { - url, - did: explicit_did.map(|d| Did(d.to_string())), - }); - } - return Err(TargetParseError::http_not_supported(raw)); - } - - // Check for handle. - if is_valid_handle(raw) { - return Ok(LabelerTarget::Identified { - identifier: AtIdentifier::Handle(raw.to_string()), - explicit_did: explicit_did.map(|d| Did(d.to_string())), - }); - } - - // Unrecognized target. - Err(TargetParseError::unrecognized_target(raw)) -} - /// Run the full labeler conformance pipeline. /// /// This is the main driver that orchestrates all validation stages. pub async fn run_pipeline(target: LabelerTarget, opts: LabelerOptions<'_>) -> LabelerReport { // Build initial header from target. let header = ReportHeader { - target: format_target(&target), + target: target.to_string(), resolved_did: None, pds_endpoint: None, labeler_endpoint: None, @@ -552,159 +391,3 @@ async fn resolve_reporter_pds_endpoint( ) }) } - -/// Format a target for display in the report header. -fn format_target(target: &LabelerTarget) -> String { - match target { - LabelerTarget::Identified { - identifier, - explicit_did, - } => { - let id_str = match identifier { - AtIdentifier::Handle(h) => h.clone(), - AtIdentifier::Did(d) => d.0.clone(), - }; - if explicit_did.is_some() { - format!("{id_str} (with explicit DID)") - } else { - id_str - } - } - LabelerTarget::Endpoint { url, did } => { - if did.is_some() { - format!("{url} (with explicit DID)") - } else { - url.to_string() - } - } - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn parse_target_handle() { - let target = parse_target("alice.bsky.social", None).expect("should parse"); - match target { - LabelerTarget::Identified { - identifier, - explicit_did, - } => { - assert!( - matches!(identifier, AtIdentifier::Handle(ref h) if h == "alice.bsky.social") - ); - assert!(explicit_did.is_none()); - } - _ => panic!("expected Identified variant"), - } - } - - #[test] - fn parse_target_did_plc() { - let target = parse_target("did:plc:abc123", None).expect("should parse"); - match target { - LabelerTarget::Identified { - identifier, - explicit_did, - } => { - assert!(matches!(identifier, AtIdentifier::Did(ref d) if d.0 == "did:plc:abc123")); - assert!(explicit_did.is_none()); - } - _ => panic!("expected Identified variant"), - } - } - - #[test] - fn parse_target_did_web() { - let target = parse_target("did:web:example.com", None).expect("should parse"); - match target { - LabelerTarget::Identified { - identifier, - explicit_did, - } => { - assert!( - matches!(identifier, AtIdentifier::Did(ref d) if d.0 == "did:web:example.com") - ); - assert!(explicit_did.is_none()); - } - _ => panic!("expected Identified variant"), - } - } - - #[test] - fn parse_target_endpoint_https() { - let target = parse_target("https://example.com/labeler", None).expect("should parse"); - match target { - LabelerTarget::Endpoint { url, did } => { - assert_eq!(url.as_str(), "https://example.com/labeler"); - assert!(did.is_none()); - } - _ => panic!("expected Endpoint variant"), - } - } - - #[test] - fn parse_target_endpoint_with_explicit_did() { - let target = - parse_target("https://example.com/labeler", Some("did:plc:xyz")).expect("should parse"); - match target { - LabelerTarget::Endpoint { url, did } => { - assert_eq!(url.as_str(), "https://example.com/labeler"); - assert_eq!(did.map(|d| d.0.clone()), Some("did:plc:xyz".to_string())); - } - _ => panic!("expected Endpoint variant"), - } - } - - #[test] - fn parse_target_endpoint_http_remote_rejected() { - let err = parse_target("http://evil.example", None).expect_err("should reject http"); - assert!(err.message.contains("HTTP")); - assert!(err.message.contains("local")); - } - - #[test] - fn parse_target_endpoint_http_local_accepted() { - // Each of these hostnames is classified as local by - // `is_local_labeler_hostname`, so plaintext HTTP is allowed. - let cases = &[ - "http://localhost:8080", - "http://127.0.0.1:5000", - "http://127.1.2.3/", - "http://[::1]:8080/", - "http://10.0.0.1/", - "http://192.168.1.100:8080", - "http://172.16.0.1/", - "http://mybox.local:8080", - ]; - for raw in cases { - let target = parse_target(raw, None) - .unwrap_or_else(|e| panic!("expected {raw} to parse, got: {}", e.message)); - match target { - LabelerTarget::Endpoint { url, did } => { - assert_eq!( - url.as_str().trim_end_matches('/'), - raw.trim_end_matches('/') - ); - assert!(did.is_none()); - } - _ => panic!("expected Endpoint variant for {raw}"), - } - } - } - - #[test] - fn parse_target_unrecognised() { - let err = parse_target("not a handle or did", None).expect_err("should fail"); - assert!(err.message.contains("Unrecognized target")); - } - - #[test] - fn parse_target_did_with_conflicting_flag() { - let err = parse_target("did:plc:abc", Some("did:web:example.com")) - .expect_err("should reject ambiguous target"); - assert!(err.message.contains("Ambiguous")); - } -} diff --git a/src/commands/test/labeler/create_report.rs b/src/commands/test/labeler/pipeline/create_report.rs similarity index 99% rename from src/commands/test/labeler/create_report.rs rename to src/commands/test/labeler/pipeline/create_report.rs index 6f433da..71868c3 100644 --- a/src/commands/test/labeler/create_report.rs +++ b/src/commands/test/labeler/pipeline/create_report.rs @@ -14,7 +14,7 @@ use miette::{Diagnostic, NamedSource, SourceSpan}; use reqwest::StatusCode; use thiserror::Error; -use crate::commands::test::labeler::identity::IdentityFacts; +use super::identity::IdentityFacts; use crate::commands::test::labeler::report::{CheckResult, CheckStatus, Stage}; use crate::common::diagnostics::pretty_json_for_display; use crate::common::identity::{Did, is_local_labeler_hostname}; @@ -645,7 +645,7 @@ pub struct CreateReportRunOptions<'a> { /// always emits exactly 10 `report::*` CheckResults (AC7.1) in canonical /// order (AC7.2), regardless of gating decisions. pub async fn run( - identity_facts: Option<&crate::commands::test::labeler::identity::IdentityFacts>, + identity_facts: Option<&super::identity::IdentityFacts>, report_tee: &dyn CreateReportTee, opts: &CreateReportRunOptions<'_>, ) -> CreateReportStageOutput { diff --git a/src/commands/test/labeler/create_report/did_doc_server.rs b/src/commands/test/labeler/pipeline/create_report/did_doc_server.rs similarity index 100% rename from src/commands/test/labeler/create_report/did_doc_server.rs rename to src/commands/test/labeler/pipeline/create_report/did_doc_server.rs diff --git a/src/commands/test/labeler/create_report/pollution.rs b/src/commands/test/labeler/pipeline/create_report/pollution.rs similarity index 100% rename from src/commands/test/labeler/create_report/pollution.rs rename to src/commands/test/labeler/pipeline/create_report/pollution.rs diff --git a/src/commands/test/labeler/create_report/self_mint.rs b/src/commands/test/labeler/pipeline/create_report/self_mint.rs similarity index 100% rename from src/commands/test/labeler/create_report/self_mint.rs rename to src/commands/test/labeler/pipeline/create_report/self_mint.rs diff --git a/src/commands/test/labeler/create_report/sentinel.rs b/src/commands/test/labeler/pipeline/create_report/sentinel.rs similarity index 100% rename from src/commands/test/labeler/create_report/sentinel.rs rename to src/commands/test/labeler/pipeline/create_report/sentinel.rs diff --git a/src/commands/test/labeler/crypto.rs b/src/commands/test/labeler/pipeline/crypto.rs similarity index 99% rename from src/commands/test/labeler/crypto.rs rename to src/commands/test/labeler/pipeline/crypto.rs index 62635b0..c8642fc 100644 --- a/src/commands/test/labeler/crypto.rs +++ b/src/commands/test/labeler/pipeline/crypto.rs @@ -459,7 +459,7 @@ struct FailedLabel { /// 6. Else if `did:plc`: fetch PLC audit log and retry against historic keys. /// 7. Else (`did:web`): emit `crypto::rollup` SpecViolation with no rotation history. pub async fn run( - identity: &crate::commands::test::labeler::identity::IdentityFacts, + identity: &super::identity::IdentityFacts, labels: &[Label], http: &dyn crate::common::identity::HttpClient, ) -> CryptoStageOutput { @@ -788,8 +788,8 @@ fn parse_signature( #[cfg(test)] mod tests { + use super::super::identity::IdentityFacts; use super::*; - use crate::commands::test::labeler::identity::IdentityFacts; use crate::common::identity::{ AnySignature, AnyVerifyingKey, Did, DidDocument, IdentityError, RawDidDocument, encode_multikey, diff --git a/src/commands/test/labeler/http.rs b/src/commands/test/labeler/pipeline/http.rs similarity index 100% rename from src/commands/test/labeler/http.rs rename to src/commands/test/labeler/pipeline/http.rs diff --git a/src/commands/test/labeler/identity.rs b/src/commands/test/labeler/pipeline/identity.rs similarity index 100% rename from src/commands/test/labeler/identity.rs rename to src/commands/test/labeler/pipeline/identity.rs diff --git a/src/commands/test/labeler/subscription.rs b/src/commands/test/labeler/pipeline/subscription.rs similarity index 100% rename from src/commands/test/labeler/subscription.rs rename to src/commands/test/labeler/pipeline/subscription.rs diff --git a/src/commands/test/labeler/target.rs b/src/commands/test/labeler/target.rs new file mode 100644 index 0000000..3954a19 --- /dev/null +++ b/src/commands/test/labeler/target.rs @@ -0,0 +1,324 @@ +use std::fmt; + +use miette::Diagnostic; +use thiserror::Error; +use url::Url; + +use crate::common::identity::{Did, is_local_labeler_hostname}; + +/// A labeler target: either a resolvable identifier (handle or DID) or a raw endpoint URL. +#[derive(Debug, Clone)] +pub enum LabelerTarget { + /// A handle or DID that can be resolved. + Identified { + /// The handle or DID to resolve. + identifier: AtIdentifier, + /// An optional explicit DID override (for cross-checking). + explicit_did: Option, + }, + /// A raw HTTP endpoint, optionally with a DID for identity checks. + Endpoint { + /// The endpoint URL. + url: Url, + /// An optional DID to cross-check against the endpoint. + did: Option, + }, +} + +/// An ATProto identifier: a handle or a DID. +#[derive(Debug, Clone)] +pub enum AtIdentifier { + /// An ATProto handle (e.g., `alice.bsky.social`). + Handle(String), + /// A decentralized identifier (e.g., `did:plc:...` or `did:web:...`). + Did(Did), +} + +/// Parse a labeler target from a string and optional explicit DID. +/// +/// Returns a `LabelerTarget` on success, or a `TargetParseError` on failure. +/// +/// Rules: +/// - If `raw` starts with `did:`, parse as a DID. If `explicit_did` is also provided, return an error. +/// - If `raw` starts with `https://`, parse as a URL. `explicit_did` is carried as an optional DID. +/// - If `raw` starts with `http://`, return an error pointing the user to HTTPS. +/// - If `raw` contains a dot and matches handle grammar, treat as a handle. `explicit_did` is carried. +/// - Otherwise, return an unrecognized target error. +pub fn parse(raw: &str, explicit_did: Option<&str>) -> Result { + // Check for DID. + if raw.starts_with("did:") { + if let Some(ed) = explicit_did { + return Err(TargetParseError::ambiguous_did(raw, ed)); + } + return Ok(LabelerTarget::Identified { + identifier: AtIdentifier::Did(Did(raw.to_string())), + explicit_did: None, + }); + } + + // Check for HTTPS URL. + if raw.starts_with("https://") { + let url = Url::parse(raw) + .map_err(|e| TargetParseError::new(format!("Invalid URL '{raw}': {e}")))?; + return Ok(LabelerTarget::Endpoint { + url, + did: explicit_did.map(|d| Did(d.to_string())), + }); + } + + // Check for HTTP URL. Local hostnames (loopback, RFC 1918, .local, mDNS) + // are accepted so developers can target a labeler running on their + // machine or LAN. Remote HTTP is still rejected to guard against + // accidental plaintext traffic to a production labeler. + if raw.starts_with("http://") { + let url = Url::parse(raw) + .map_err(|e| TargetParseError::new(format!("Invalid URL '{raw}': {e}")))?; + if is_local_labeler_hostname(&url) { + return Ok(LabelerTarget::Endpoint { + url, + did: explicit_did.map(|d| Did(d.to_string())), + }); + } + return Err(TargetParseError::http_not_supported(raw)); + } + + // Check for handle. + if is_valid_handle(raw) { + return Ok(LabelerTarget::Identified { + identifier: AtIdentifier::Handle(raw.to_string()), + explicit_did: explicit_did.map(|d| Did(d.to_string())), + }); + } + + // Unrecognized target. + Err(TargetParseError::unrecognized_target(raw)) +} + +impl fmt::Display for LabelerTarget { + /// Format a target for display in the report header. + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + LabelerTarget::Identified { + identifier, + explicit_did, + } => { + let id_str = match identifier { + AtIdentifier::Handle(h) => h.clone(), + AtIdentifier::Did(d) => d.0.clone(), + }; + if explicit_did.is_some() { + write!(f, "{id_str} (with explicit DID)") + } else { + id_str.fmt(f) + } + } + LabelerTarget::Endpoint { url, did } => { + if did.is_some() { + write!(f, "{url} (with explicit DID)") + } else { + url.fmt(f) + } + } + } + } +} + +/// Check if a string is a valid ATProto handle. +/// +/// A valid handle: +/// - Contains at least one dot. +/// - Contains only alphanumeric characters, hyphens, and dots. +/// - Does not start or end with a hyphen or dot. +/// - Has no empty segments (no consecutive dots or leading/trailing dots). +fn is_valid_handle(s: &str) -> bool { + if !s.contains('.') { + return false; + } + + // Check for empty string or leading/trailing special chars. + if s.is_empty() + || s.starts_with('-') + || s.starts_with('.') + || s.ends_with('-') + || s.ends_with('.') + { + return false; + } + + // Check all characters are alphanumeric, hyphen, or dot. + for c in s.chars() { + if !c.is_ascii_alphanumeric() && c != '-' && c != '.' { + return false; + } + } + + // Check no empty segments (no consecutive dots). + if s.contains("..") { + return false; + } + + true +} + +/// Error from target parsing. +#[derive(Debug, Error, Diagnostic)] +#[error("{message}")] +pub struct TargetParseError { + /// The error message. + pub message: String, +} + +impl TargetParseError { + fn new(message: impl Into) -> Self { + Self { + message: message.into(), + } + } + + fn unrecognized_target(raw: &str) -> Self { + Self::new(format!( + "Unrecognized target '{raw}'. Expected one of:\n - ATProto handle (e.g., alice.bsky.social)\n - DID (e.g., did:plc:abc123 or did:web:example.com)\n - HTTPS endpoint URL (e.g., https://labeler.example.com)\n - HTTP endpoint URL with a local hostname (e.g., http://localhost:8080)" + )) + } + + fn http_not_supported(raw: &str) -> Self { + Self::new(format!( + "HTTP endpoint '{raw}' is not supported for remote hosts. Use HTTPS, or point at a local labeler (localhost / 127.0.0.0/8 / RFC 1918 / .local) to allow plaintext HTTP." + )) + } + + fn ambiguous_did(raw: &str, explicit: &str) -> Self { + Self::new(format!( + "Ambiguous target specification: target '{raw}' is already a DID, but --did {explicit} was also provided. Please use only one." + )) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parse_target_handle() { + let target = parse("alice.bsky.social", None).expect("should parse"); + match target { + LabelerTarget::Identified { + identifier, + explicit_did, + } => { + assert!( + matches!(identifier, AtIdentifier::Handle(ref h) if h == "alice.bsky.social") + ); + assert!(explicit_did.is_none()); + } + _ => panic!("expected Identified variant"), + } + } + + #[test] + fn parse_target_did_plc() { + let target = parse("did:plc:abc123", None).expect("should parse"); + match target { + LabelerTarget::Identified { + identifier, + explicit_did, + } => { + assert!(matches!(identifier, AtIdentifier::Did(ref d) if d.0 == "did:plc:abc123")); + assert!(explicit_did.is_none()); + } + _ => panic!("expected Identified variant"), + } + } + + #[test] + fn parse_target_did_web() { + let target = parse("did:web:example.com", None).expect("should parse"); + match target { + LabelerTarget::Identified { + identifier, + explicit_did, + } => { + assert!( + matches!(identifier, AtIdentifier::Did(ref d) if d.0 == "did:web:example.com") + ); + assert!(explicit_did.is_none()); + } + _ => panic!("expected Identified variant"), + } + } + + #[test] + fn parse_target_endpoint_https() { + let target = parse("https://example.com/labeler", None).expect("should parse"); + match target { + LabelerTarget::Endpoint { url, did } => { + assert_eq!(url.as_str(), "https://example.com/labeler"); + assert!(did.is_none()); + } + _ => panic!("expected Endpoint variant"), + } + } + + #[test] + fn parse_target_endpoint_with_explicit_did() { + let target = + parse("https://example.com/labeler", Some("did:plc:xyz")).expect("should parse"); + match target { + LabelerTarget::Endpoint { url, did } => { + assert_eq!(url.as_str(), "https://example.com/labeler"); + assert_eq!(did.map(|d| d.0.clone()), Some("did:plc:xyz".to_string())); + } + _ => panic!("expected Endpoint variant"), + } + } + + #[test] + fn parse_target_endpoint_http_remote_rejected() { + let err = parse("http://evil.example", None).expect_err("should reject http"); + assert!(err.message.contains("HTTP")); + assert!(err.message.contains("local")); + } + + #[test] + fn parse_target_endpoint_http_local_accepted() { + // Each of these hostnames is classified as local by + // `is_local_labeler_hostname`, so plaintext HTTP is allowed. + let cases = &[ + "http://localhost:8080", + "http://127.0.0.1:5000", + "http://127.1.2.3/", + "http://[::1]:8080/", + "http://10.0.0.1/", + "http://192.168.1.100:8080", + "http://172.16.0.1/", + "http://mybox.local:8080", + ]; + for raw in cases { + let target = parse(raw, None) + .unwrap_or_else(|e| panic!("expected {raw} to parse, got: {}", e.message)); + match target { + LabelerTarget::Endpoint { url, did } => { + assert_eq!( + url.as_str().trim_end_matches('/'), + raw.trim_end_matches('/') + ); + assert!(did.is_none()); + } + _ => panic!("expected Endpoint variant for {raw}"), + } + } + } + + #[test] + fn parse_target_unrecognised() { + let err = parse("not a handle or did", None).expect_err("should fail"); + assert!(err.message.contains("Unrecognized target")); + } + + #[test] + fn parse_target_did_with_conflicting_flag() { + let err = parse("did:plc:abc", Some("did:web:example.com")) + .expect_err("should reject ambiguous target"); + assert!(err.message.contains("Ambiguous")); + } +} diff --git a/tests/common/mod.rs b/tests/common/mod.rs index fa83298..142476c 100644 --- a/tests/common/mod.rs +++ b/tests/common/mod.rs @@ -8,13 +8,13 @@ #![allow(dead_code)] use async_trait::async_trait; -use atproto_devtool::commands::test::labeler::create_report::{ - CreateReportStageError, CreateReportTee, PdsXrpcClient, RawCreateReportResponse, - RawPdsXrpcResponse, -}; -use atproto_devtool::commands::test::labeler::http::{HttpStageError, RawHttpTee, RawXrpcResponse}; -use atproto_devtool::commands::test::labeler::subscription::{ - FrameStream, SubscriptionStageError, WebSocketClient, +use atproto_devtool::commands::test::labeler::pipeline::{ + create_report::{ + CreateReportStageError, CreateReportTee, PdsXrpcClient, RawCreateReportResponse, + RawPdsXrpcResponse, + }, + http::{HttpStageError, RawHttpTee, RawXrpcResponse}, + subscription::{FrameStream, SubscriptionStageError, WebSocketClient}, }; use reqwest::StatusCode; use std::collections::HashMap; diff --git a/tests/common_fakes.rs b/tests/common_fakes.rs index 1155856..e9f43d0 100644 --- a/tests/common_fakes.rs +++ b/tests/common_fakes.rs @@ -2,7 +2,7 @@ mod common; -use atproto_devtool::commands::test::labeler::create_report::CreateReportTee; +use atproto_devtool::commands::test::labeler::pipeline::create_report::CreateReportTee; use common::*; use reqwest::StatusCode; diff --git a/tests/labeler_endtoend.rs b/tests/labeler_endtoend.rs index b26caea..c5a331e 100644 --- a/tests/labeler_endtoend.rs +++ b/tests/labeler_endtoend.rs @@ -3,12 +3,16 @@ mod common; use async_trait::async_trait; -use atproto_devtool::commands::test::labeler::create_report::self_mint::SelfMintCurve; -use atproto_devtool::commands::test::labeler::crypto::canonicalize_label_for_signing; -use atproto_devtool::commands::test::labeler::pipeline::{ - CreateReportTeeKind, HttpTee, LabelerOptions, parse_target, run_pipeline, +use atproto_devtool::{ + commands::test::labeler::{ + pipeline::{ + CreateReportTeeKind, HttpTee, LabelerOptions, create_report::self_mint::SelfMintCurve, + crypto::canonicalize_label_for_signing, run_pipeline, + }, + target, + }, + common::identity::{DnsResolver, HttpClient, IdentityError}, }; -use atproto_devtool::common::identity::{DnsResolver, HttpClient, IdentityError}; use atrium_api::com::atproto::label::defs::{Label, LabelData}; use atrium_api::com::atproto::label::query_labels::{Output, OutputData}; use atrium_api::types::string::Datetime; @@ -284,7 +288,7 @@ async fn all_pass_exits_zero_and_renders_all_ok() { labeler_record_json.to_vec(), ); - let target = parse_target(did, None).expect("parse failed"); + let target = target::parse(did, None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); fake_tee.add_response(None, 200, labels_response); let fake_ws = common::FakeWebSocketClient::empty(); @@ -349,7 +353,8 @@ async fn identity_only_failure_exits_one_with_severity_breakdown() { serde_json::to_vec(&did_json).unwrap(), ); - let target = parse_target("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); + let target = + target::parse("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); let fake_ws = common::FakeWebSocketClient::empty(); let fake_report_tee = common::FakeCreateReportTee::new(); @@ -404,7 +409,8 @@ async fn http_decode_failure_exits_one() { labeler_record_json.to_vec(), ); - let target = parse_target("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); + let target = + target::parse("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); fake_tee.add_response(None, 200, b"{malformed json".to_vec()); let fake_ws = common::FakeWebSocketClient::empty(); @@ -457,7 +463,8 @@ async fn subscription_transport_error_exits_two() { labeler_record_json.to_vec(), ); - let target = parse_target("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); + let target = + target::parse("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); fake_tee.add_response(None, 200, br#"{"cursor": null, "labels": []}"#.to_vec()); @@ -561,7 +568,7 @@ async fn current_key_fail_history_pass_exits_zero_with_advisory() { serde_json::to_vec(&plc_audit_log).expect("serialize plc audit log"), ); - let target = parse_target(did, None).expect("parse failed"); + let target = target::parse(did, None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); fake_tee.add_response(None, 200, labels_response); let fake_ws = common::FakeWebSocketClient::empty(); @@ -654,7 +661,8 @@ async fn current_key_fail_history_also_fail_exits_one() { serde_json::to_vec(&plc_audit_log).unwrap(), ); - let target = parse_target("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); + let target = + target::parse("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); fake_tee.add_response(None, 200, labels_response); let fake_ws = common::FakeWebSocketClient::empty(); @@ -744,7 +752,7 @@ async fn did_web_current_key_fail_exits_one() { labeler_record_json.to_vec(), ); - let target = parse_target("did:web:web-labeler.example", None).expect("parse failed"); + let target = target::parse("did:web:web-labeler.example", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); fake_tee.add_response(None, 200, labels_response); let fake_ws = common::FakeWebSocketClient::empty(); @@ -799,7 +807,8 @@ async fn canonicalization_error_distinct_diagnostic_code() { labeler_record_json.to_vec(), ); - let target = parse_target("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); + let target = + target::parse("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); fake_tee.add_response(None, 200, labels_response); let fake_ws = common::FakeWebSocketClient::empty(); @@ -879,7 +888,8 @@ async fn plc_directory_unreachable_network_error_no_false_fail() { "https://plc.directory/did:plc:test123456789abcdefghijklmnop/log/audit", ); - let target = parse_target("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); + let target = + target::parse("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); fake_tee.add_response(None, 200, labels_response); let fake_ws = common::FakeWebSocketClient::empty(); @@ -935,7 +945,8 @@ async fn empty_labeler_skipped_crypto_only_advisory_exits_zero() { labeler_record_json.to_vec(), ); - let target = parse_target("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); + let target = + target::parse("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); fake_tee.add_response(None, 200, br#"{"cursor": null, "labels": []}"#.to_vec()); let fake_ws = common::FakeWebSocketClient::empty(); @@ -995,7 +1006,8 @@ async fn exit_code_summary_for_network_only_run_is_two() { labeler_record_json.to_vec(), ); - let target = parse_target("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); + let target = + target::parse("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); fake_tee.set_transport_error(); // Force HTTP stage transport error. let fake_ws = common::FakeWebSocketClient::empty(); @@ -1040,7 +1052,7 @@ async fn skipped_reasons_rendered() { let http = FakeHttpClient::new(); let dns = FakeDnsResolver::new(); - let target = parse_target("https://labeler.example.com", None).expect("parse failed"); + let target = target::parse("https://labeler.example.com", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); fake_tee.add_response(None, 200, br#"{"cursor": null, "labels": []}"#.to_vec()); let fake_ws = common::FakeWebSocketClient::empty(); diff --git a/tests/labeler_http.rs b/tests/labeler_http.rs index 0526745..f55e0ae 100644 --- a/tests/labeler_http.rs +++ b/tests/labeler_http.rs @@ -2,7 +2,7 @@ mod common; -use atproto_devtool::commands::test::labeler::http::run; +use atproto_devtool::commands::test::labeler::pipeline::http; use common::FakeRawHttpTee; /// Helper to render a report to a string for snapshot testing. @@ -29,7 +29,7 @@ async fn http_healthy_renders_all_ok() { tee.add_response(None, 200, first_page); tee.add_response(Some("cursor1"), 200, second_page); - let output = run(&tee).await; + let output = http::run(&tee).await; // Build a minimal report for snapshot testing. let mut report = atproto_devtool::commands::test::labeler::report::LabelerReport::new( @@ -57,7 +57,7 @@ async fn http_empty_labeler_emits_advisory() { tee.add_response(None, 200, empty_page); - let output = run(&tee).await; + let output = http::run(&tee).await; let mut report = atproto_devtool::commands::test::labeler::report::LabelerReport::new( atproto_devtool::commands::test::labeler::report::ReportHeader { @@ -84,7 +84,7 @@ async fn http_malformed_schema_fails_with_source_span() { tee.add_response(None, 200, malformed.clone()); - let output = run(&tee).await; + let output = http::run(&tee).await; let mut report = atproto_devtool::commands::test::labeler::report::LabelerReport::new( atproto_devtool::commands::test::labeler::report::ReportHeader { @@ -116,7 +116,7 @@ async fn http_ignored_cursor_fails() { tee.add_response(None, 200, first_page); tee.add_response(Some("cursor1"), 200, second_page); - let output = run(&tee).await; + let output = http::run(&tee).await; let mut report = atproto_devtool::commands::test::labeler::report::LabelerReport::new( atproto_devtool::commands::test::labeler::report::ReportHeader { @@ -141,7 +141,7 @@ async fn http_transport_error_renders_network_error() { let tee = FakeRawHttpTee::new(); tee.set_transport_error(); - let output = run(&tee).await; + let output = http::run(&tee).await; let mut report = atproto_devtool::commands::test::labeler::report::LabelerReport::new( atproto_devtool::commands::test::labeler::report::ReportHeader { diff --git a/tests/labeler_identity.rs b/tests/labeler_identity.rs index 27369ae..32f6752 100644 --- a/tests/labeler_identity.rs +++ b/tests/labeler_identity.rs @@ -2,15 +2,20 @@ mod common; -use async_trait::async_trait; -use atproto_devtool::commands::test::labeler::create_report::self_mint::SelfMintCurve; -use atproto_devtool::commands::test::labeler::pipeline::{ - CreateReportTeeKind, HttpTee, LabelerOptions, parse_target, run_pipeline, -}; -use atproto_devtool::common::identity::{DnsResolver, HttpClient, IdentityError}; use std::collections::HashMap; use std::sync::{Arc, Mutex}; +use async_trait::async_trait; +use atproto_devtool::{ + commands::test::labeler::{ + pipeline::{ + CreateReportTeeKind, HttpTee, LabelerOptions, create_report::self_mint::SelfMintCurve, + run_pipeline, + }, + target, + }, + common::identity::{DnsResolver, HttpClient, IdentityError}, +}; use url::Url; /// Type alias for the response map in FakeHttpClient. @@ -143,7 +148,8 @@ async fn healthy_plc_renders_all_ok() { labeler_record_json, ); - let target = parse_target("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); + let target = + target::parse("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); fake_tee.add_response(None, 200, healthy_labels_response()); let fake_ws = common::FakeWebSocketClient::empty(); @@ -179,7 +185,7 @@ async fn endpoint_only_no_did_skips_identity() { let http = FakeHttpClient::new(); let dns = FakeDnsResolver::new(); - let target = parse_target("https://example.com/labeler", None).expect("parse failed"); + let target = target::parse("https://example.com/labeler", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); let empty_response = br#"{"cursor":null,"labels":[]}"#.to_vec(); fake_tee.add_response(None, 200, empty_response); @@ -240,7 +246,7 @@ async fn handle_resolution_happy_path() { labeler_record_json, ); - let target = parse_target("alice.example", None).expect("parse failed"); + let target = target::parse("alice.example", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); let empty_response = br#"{"cursor":null,"labels":[]}"#.to_vec(); fake_tee.add_response(None, 200, empty_response); @@ -295,7 +301,8 @@ async fn did_plc_direct_happy_path() { labeler_record_json, ); - let target = parse_target("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); + let target = + target::parse("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); fake_tee.add_response(None, 200, healthy_labels_response()); let fake_ws = common::FakeWebSocketClient::empty(); @@ -348,7 +355,7 @@ async fn did_web_direct_happy_path() { labeler_record_json, ); - let target = parse_target("did:web:web-labeler.example", None).expect("parse failed"); + let target = target::parse("did:web:web-labeler.example", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); let empty_response = br#"{"cursor":null,"labels":[]}"#.to_vec(); fake_tee.add_response(None, 200, empty_response); @@ -388,7 +395,7 @@ async fn plc_directory_unreachable_renders_network_error() { // Don't add any response for PLC directory or DNS resolver - causes network error. - let target = parse_target("alice.test", None).expect("parse failed"); + let target = target::parse("alice.test", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); let empty_response = br#"{"cursor":null,"labels":[]}"#.to_vec(); fake_tee.add_response(None, 200, empty_response); @@ -441,7 +448,8 @@ async fn missing_labeler_record_renders_404_distinct_from_transport() { b"Not found".to_vec(), ); - let target = parse_target("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); + let target = + target::parse("did:plc:test123456789abcdefghijklmnop", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); fake_tee.add_response(None, 200, healthy_labels_response()); let fake_ws = common::FakeWebSocketClient::empty(); @@ -495,7 +503,7 @@ async fn missing_service_renders_spec_violation_with_span() { ); let target = - parse_target("did:plc:missing_service_test_123456789", None).expect("parse failed"); + target::parse("did:plc:missing_service_test_123456789", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); let empty_response = br#"{"cursor":null,"labels":[]}"#.to_vec(); fake_tee.add_response(None, 200, empty_response); @@ -553,7 +561,7 @@ async fn missing_signing_key_renders_spec_violation() { ); let target = - parse_target("did:plc:missing_signing_key_test_12345", None).expect("parse failed"); + target::parse("did:plc:missing_signing_key_test_12345", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); let empty_response = br#"{"cursor":null,"labels":[]}"#.to_vec(); fake_tee.add_response(None, 200, empty_response); @@ -609,7 +617,7 @@ async fn non_https_endpoint_renders_spec_violation() { ); let target = - parse_target("did:plc:non_https_endpoint_test_123456", None).expect("parse failed"); + target::parse("did:plc:non_https_endpoint_test_123456", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); let empty_response = br#"{"cursor":null,"labels":[]}"#.to_vec(); fake_tee.add_response(None, 200, empty_response); @@ -665,7 +673,7 @@ async fn empty_policies_renders_spec_violation_with_span() { ); let target = - parse_target("did:plc:empty_policies_test_123456789ab", None).expect("parse failed"); + target::parse("did:plc:empty_policies_test_123456789ab", None).expect("parse failed"); let fake_tee = common::FakeRawHttpTee::new(); let empty_response = br#"{"cursor":null,"labels":[]}"#.to_vec(); fake_tee.add_response(None, 200, empty_response); @@ -721,7 +729,7 @@ async fn endpoint_mismatch_spec_violation() { ); // The target is the mismatched URL; the DID says a different endpoint. - let target = parse_target( + let target = target::parse( "https://other-labeler.example/", Some("did:plc:endpoint_mismatch_test_123456789"), ) @@ -791,7 +799,7 @@ async fn local_http_override_mismatch_is_advisory() { ); // Target is a local labeler copy; the DID document points somewhere else. - let target = parse_target( + let target = target::parse( "http://localhost:8080", Some("did:plc:endpoint_mismatch_test_123456789"), ) diff --git a/tests/labeler_report.rs b/tests/labeler_report.rs index 92938b8..73eb05e 100644 --- a/tests/labeler_report.rs +++ b/tests/labeler_report.rs @@ -11,19 +11,24 @@ use std::borrow::Cow; use std::sync::Arc; use std::time::Duration; -use atproto_devtool::commands::test::labeler::create_report; -use atproto_devtool::commands::test::labeler::create_report::Check; -use atproto_devtool::commands::test::labeler::create_report::self_mint::{ - SelfMintCurve, SelfMintSigner, -}; -use atproto_devtool::commands::test::labeler::identity::IdentityFacts; -use atproto_devtool::commands::test::labeler::pipeline::{ - CreateReportTeeKind, HttpTee, LabelerOptions, parse_target, run_pipeline, -}; -use atproto_devtool::commands::test::labeler::report::{ - CheckResult, CheckStatus, LabelerReport, RenderConfig, ReportHeader, +use atproto_devtool::{ + commands::test::labeler::{ + report::{CheckResult, CheckStatus, LabelerReport, RenderConfig, ReportHeader}, + { + pipeline::{ + CreateReportTeeKind, HttpTee, LabelerOptions, + create_report::{ + self, + self_mint::{SelfMintCurve, SelfMintSigner}, + }, + identity::IdentityFacts, + run_pipeline, + }, + target, + }, + }, + common::identity::{Did, DnsResolver, HttpClient, IdentityError}, }; -use atproto_devtool::common::identity::{Did, DnsResolver, HttpClient, IdentityError}; use atrium_api::app::bsky::labeler::defs::LabelerPolicies; @@ -478,7 +483,7 @@ async fn pipeline_integration_happy_path_via_endpoint() { let fake_create_report_tee = common::FakeCreateReportTee::new(); // Use endpoint target (no identity stage needed). - let target = parse_target("https://labeler.example.com", None).expect("parse endpoint target"); + let target = target::parse("https://labeler.example.com", None).expect("parse endpoint target"); let opts = LabelerOptions { http: &StubHttpClient, @@ -1522,7 +1527,7 @@ async fn ac7_1_row_count_is_always_10() { "AC7.1 failed: expected 10 rows with commit={commit}, force={force}, pds={pds}" ); // Verify Check::ORDER is respected. - let expected_ids: Vec<_> = Check::ORDER.iter().map(|c| c.id()).collect(); + let expected_ids: Vec<_> = create_report::Check::ORDER.iter().map(|c| c.id()).collect(); let actual_ids: Vec<_> = results.iter().map(|r| r.id).collect(); assert_eq!(actual_ids, expected_ids, "Check order mismatch"); } diff --git a/tests/labeler_subscription.rs b/tests/labeler_subscription.rs index 6aaca29..a7b09e9 100644 --- a/tests/labeler_subscription.rs +++ b/tests/labeler_subscription.rs @@ -7,15 +7,22 @@ mod common; -use atproto_devtool::commands::test::labeler::create_report::self_mint::SelfMintCurve; -use atproto_devtool::commands::test::labeler::pipeline::{ - CreateReportTeeKind, HttpTee, LabelerOptions, parse_target, run_pipeline, -}; -use atproto_devtool::commands::test::labeler::subscription::{FrameHeader, SubscribeLabelsPayload}; -use atproto_devtool::common::identity::{DnsResolver, HttpClient, IdentityError}; use std::collections::HashMap; use std::sync::{Arc, Mutex}; use std::time::Duration; + +use atproto_devtool::{ + commands::test::labeler::{ + pipeline::{ + CreateReportTeeKind, HttpTee, LabelerOptions, + create_report::self_mint::SelfMintCurve, + run_pipeline, + subscription::{FrameHeader, SubscribeLabelsPayload}, + }, + target, + }, + common::identity::{DnsResolver, HttpClient, IdentityError}, +}; use url::Url; /// Type alias for the response map in FakeHttpClient. @@ -284,7 +291,7 @@ async fn backfill_completes_within_budget_passes() { let http = FakeHttpClient::new(); let dns = FakeDnsResolver::new(); let fake_tee = make_passing_http_tee(); - let target = parse_target("https://example.com/labeler", None).expect("parse failed"); + let target = target::parse("https://example.com/labeler", None).expect("parse failed"); let opts = LabelerOptions { http: &http, @@ -343,7 +350,7 @@ async fn backfill_exceeds_budget_triggers_live_tail() { let http = FakeHttpClient::new(); let dns = FakeDnsResolver::new(); let fake_tee = make_passing_http_tee(); - let target = parse_target("https://example.com/labeler", None).expect("parse failed"); + let target = target::parse("https://example.com/labeler", None).expect("parse failed"); let opts = LabelerOptions { http: &http, @@ -388,7 +395,7 @@ async fn empty_stream_advisories() { let http = FakeHttpClient::new(); let dns = FakeDnsResolver::new(); let fake_tee = make_passing_http_tee(); - let target = parse_target("https://example.com/labeler", None).expect("parse failed"); + let target = target::parse("https://example.com/labeler", None).expect("parse failed"); let opts = LabelerOptions { http: &http, @@ -444,7 +451,7 @@ async fn malformed_frame_emits_spec_violation() { let http = FakeHttpClient::new(); let dns = FakeDnsResolver::new(); let fake_tee = make_passing_http_tee(); - let target = parse_target("https://example.com/labeler", None).expect("parse failed"); + let target = target::parse("https://example.com/labeler", None).expect("parse failed"); let opts = LabelerOptions { http: &http, @@ -501,7 +508,7 @@ async fn error_frame_malformed_payload_spec_violation() { let http = FakeHttpClient::new(); let dns = FakeDnsResolver::new(); let fake_tee = make_passing_http_tee(); - let target = parse_target("https://example.com/labeler", None).expect("parse failed"); + let target = target::parse("https://example.com/labeler", None).expect("parse failed"); let opts = LabelerOptions { http: &http, @@ -545,7 +552,7 @@ async fn unreachable_endpoint_network_error() { let http = FakeHttpClient::new(); let dns = FakeDnsResolver::new(); let fake_tee = make_passing_http_tee(); - let target = parse_target("https://example.com/labeler", None).expect("parse failed"); + let target = target::parse("https://example.com/labeler", None).expect("parse failed"); let opts = LabelerOptions { http: &http, @@ -601,7 +608,7 @@ async fn live_tail_connect_failure_emits_network_error() { let http = FakeHttpClient::new(); let dns = FakeDnsResolver::new(); let fake_tee = make_passing_http_tee(); - let target = parse_target("https://example.com/labeler", None).expect("parse failed"); + let target = target::parse("https://example.com/labeler", None).expect("parse failed"); let opts = LabelerOptions { http: &http, @@ -668,7 +675,7 @@ async fn mid_stream_transport_error_does_not_reset_idle_gap() { let http = FakeHttpClient::new(); let dns = FakeDnsResolver::new(); let fake_tee = make_passing_http_tee(); - let target = parse_target("https://example.com/labeler", None).expect("parse failed"); + let target = target::parse("https://example.com/labeler", None).expect("parse failed"); let opts = LabelerOptions { http: &http,