feat: extract alknet-call into alkcall, vendor core types
Phase 1 of the alknet-call + alknet-channels unification. The call crate
is extracted verbatim from /workspace/@alkdev/alknet/crates/alknet-call
and the needed alknet-core types are vendored into src/core/ — alkcall is
the home for these types going forward (no separate alkcore crate).
Vendored core types (src/core/):
- auth.rs: Identity, AuthToken, AuthContext, IdentityProvider
- ownership.rs: OwnershipProvider, OwnershipStore, InMemoryOwnershipStore,
OwnershipError
- types.rs: ProtocolHandler, Connection, BiStream, BidiStreamSource,
SendStream, RecvStream, HandlerError, StreamError, Capabilities, Secret,
IdentityAlreadySet
- Stripped: quinn/iroh/rustls deps, Connection::from_quinn/from_iroh
(the dial lives in the consumer per ADR-089), config/credentials/
fingerprint/store modules, ConfigIdentityProvider/IdentityStore
Call crate (src/client, src/protocol, src/registry/):
- Copied verbatim from alknet-call; alknet_core::{auth,types,ownership}
rebound to crate::core::{auth,types,ownership}
- ALPN strings unchanged (alknet/call — wire-stable per ADR-006)
- No behavioral changes
Verification:
- cargo test: 343 passed (0 failed)
- cargo clippy --all-targets -- -D warnings: clean
- cargo fmt --check: clean
- cargo doc --no-deps: 1 pre-existing broken intra-doc link
(CallAdapter in dispatch.rs — Phase 2 cleanup)
Next: Phase 2 — channels greenfield implementation + ALPN rename +
architecture doc porting.
This commit is contained in:
1 parent
a779dd0d0d
commit
4bc7a19695
25 files changed
+14023
No files matched your search
@@ -0,0 +1,3 @@
|
||||
target/
|
||||
node_modules/
|
||||
.worktrees/
|
||||
Generated
+1269
File diff suppressed because it is too large.
Load diff
+31
@@ -0,0 +1,31 @@
|
||||
[package]
|
||||
name = "alkcall"
|
||||
version = "0.1.0"
|
||||
edition = "2021"
|
||||
rust-version = "1.85"
|
||||
license = "MIT OR Apache-2.0"
|
||||
description = "Call + channels RPC: structured JSON operations, streaming subscriptions, service discovery, and N-channel multiplexing over one transport stream"
|
||||
repository = "https://git.alk.dev/alkdev/alkcall"
|
||||
readme = "README.md"
|
||||
keywords = ["rpc", "json-rpc", "multiplexing", "wire-format", "alpn"]
|
||||
categories = ["network-programming", "asynchronous", "encoding"]
|
||||
exclude = [".opencode/", "docs/reviews/", "docs/research/", "docs/sdd_process.md", "Cargo.lock"]
|
||||
|
||||
[lib]
|
||||
name = "alkcall"
|
||||
|
||||
[features]
|
||||
default = []
|
||||
|
||||
[dependencies]
|
||||
alktype = "0.1.0"
|
||||
tokio = { version = "1", features = ["full"] }
|
||||
serde = { version = "1", features = ["derive"] }
|
||||
serde_json = "1"
|
||||
async-trait = "0.1"
|
||||
tracing = "0.1"
|
||||
thiserror = "2"
|
||||
uuid = { version = "1", features = ["v4"] }
|
||||
futures = "0.3"
|
||||
parking_lot = "0.12"
|
||||
zeroize = { version = "1", features = ["alloc", "derive"] }
|
||||
@@ -0,0 +1,197 @@
|
||||
//! `CallClient`: the outbound connection opener (ADR-017 §1).
|
||||
//!
|
||||
//! Runs the shared dispatch loop over a pre-established `Connection`
|
||||
//! (delegated to [`crate::protocol::dispatch::Dispatcher`]).
|
||||
//! `CallClient` is the connection-establishment half; `CallAdapter`'s accept
|
||||
//! path is the inbound half; both produce a `CallConnection` and hand it to
|
||||
//! the same `Dispatcher::run_loop` (ADR-017 §1).
|
||||
//!
|
||||
//! After establishment the connection is symmetric (ADR-017 §2): both sides
|
||||
//! can send and receive `call.requested`. The `CallClient` is both a caller
|
||||
//! (initiates outgoing calls via `CallConnection::call()`/`subscribe()`/
|
||||
//! `abort()`) and a callee (dispatches incoming calls against its registry).
|
||||
//!
|
||||
//! Transport-level connection establishment (QUIC dial, TCP+TLS, iroh) is
|
||||
//! handled by `alknet-client`; `CallClient::spawn_dispatch` takes a
|
||||
//! pre-established `Connection` and runs the call protocol over it.
|
||||
//!
|
||||
//! See `docs/architecture/crates/call/client-and-adapters.md` for the spec.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::core::auth::IdentityProvider;
|
||||
use crate::core::types::Connection;
|
||||
|
||||
use crate::protocol::connection::CallConnection;
|
||||
use crate::protocol::dispatch::Dispatcher;
|
||||
use crate::registry::registration::OperationRegistry;
|
||||
|
||||
/// Outbound `alknet/call` connection opener (the #1 gap, ADR-017 §1).
|
||||
///
|
||||
/// Peer authorization flows through the existing `AccessControl::check` gate
|
||||
/// in `OperationRegistry::invoke` (ADR-029 §3) — no parallel `remote_safe`/
|
||||
/// `trusted_peer` gate.
|
||||
pub struct CallClient {
|
||||
registry: Arc<OperationRegistry>,
|
||||
identity_provider: Arc<dyn IdentityProvider>,
|
||||
}
|
||||
|
||||
impl CallClient {
|
||||
pub fn new(
|
||||
registry: Arc<OperationRegistry>,
|
||||
identity_provider: Arc<dyn IdentityProvider>,
|
||||
) -> Self {
|
||||
Self {
|
||||
registry,
|
||||
identity_provider,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn registry(&self) -> &Arc<OperationRegistry> {
|
||||
&self.registry
|
||||
}
|
||||
|
||||
pub fn identity_provider(&self) -> &Arc<dyn IdentityProvider> {
|
||||
&self.identity_provider
|
||||
}
|
||||
|
||||
/// Run the shared dispatch loop over a pre-established `Connection`. The
|
||||
/// `CallClient` spawns the dispatcher task and returns a live
|
||||
/// `CallConnection` the caller can use immediately. Used by the assembly
|
||||
/// layer after `AlknetClient::dial_*` + `spawn_dispatch` and by
|
||||
/// integration tests that wire a mock/loopback `Connection` directly.
|
||||
pub fn spawn_dispatch(&self, connection: Connection) -> CallConnection {
|
||||
let call_connection = Arc::new(CallConnection::new(connection));
|
||||
let dispatcher = Dispatcher::new(
|
||||
Arc::clone(&self.registry),
|
||||
Arc::clone(&self.identity_provider),
|
||||
);
|
||||
let run_conn = Arc::clone(&call_connection);
|
||||
tokio::spawn(async move {
|
||||
dispatcher.run_loop(run_conn).await;
|
||||
});
|
||||
(*call_connection).clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::core::auth::Identity;
|
||||
use crate::core::types::Capabilities;
|
||||
use crate::protocol::connection::CallConnection;
|
||||
use crate::protocol::wire::ResponseEnvelope;
|
||||
use crate::registry::registration::{
|
||||
make_handler, Handler, HandlerKind, HandlerRegistration, OperationProvenance,
|
||||
};
|
||||
use crate::registry::spec::{AccessControl, OperationSpec, OperationType, Visibility};
|
||||
|
||||
use crate::protocol::sink_empty_connection as stub_connection;
|
||||
|
||||
fn external_spec(name: &str) -> OperationSpec {
|
||||
OperationSpec::new(
|
||||
name,
|
||||
OperationType::Query,
|
||||
Visibility::External,
|
||||
serde_json::json!({}),
|
||||
serde_json::json!({}),
|
||||
vec![],
|
||||
AccessControl::default(),
|
||||
None,
|
||||
)
|
||||
}
|
||||
|
||||
fn caps_inspect_handler() -> Handler {
|
||||
make_handler(|_input, context| async move {
|
||||
let has_google = context.capabilities.get("google").is_some();
|
||||
ResponseEnvelope::ok(
|
||||
context.request_id,
|
||||
serde_json::json!({ "has_google_capability": has_google }),
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
struct NoopIdentityProvider;
|
||||
impl crate::core::auth::IdentityProvider for NoopIdentityProvider {
|
||||
fn resolve_from_fingerprint(&self, _fp: &str) -> Option<Identity> {
|
||||
None
|
||||
}
|
||||
fn resolve_from_token(&self, _token: &crate::core::auth::AuthToken) -> Option<Identity> {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
fn registry_with_caps() -> Arc<OperationRegistry> {
|
||||
let mut registry = OperationRegistry::new();
|
||||
registry
|
||||
.register(HandlerRegistration::new(
|
||||
external_spec("pub/run"),
|
||||
HandlerKind::Once(caps_inspect_handler()),
|
||||
OperationProvenance::Local,
|
||||
None,
|
||||
None,
|
||||
Capabilities::new().with_api_key("google", "pub-key".to_string()),
|
||||
))
|
||||
.unwrap();
|
||||
Arc::new(registry)
|
||||
}
|
||||
|
||||
fn dispatcher(registry: &Arc<OperationRegistry>) -> Dispatcher {
|
||||
Dispatcher::new(Arc::clone(registry), Arc::new(NoopIdentityProvider))
|
||||
}
|
||||
|
||||
async fn dispatch(d: &Dispatcher, conn: &Arc<CallConnection>, op: &str) -> ResponseEnvelope {
|
||||
d.dispatch_requested(
|
||||
conn,
|
||||
"req-test".to_string(),
|
||||
serde_json::json!({ "operationId": op, "input": {} }),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn external_op_dispatches_and_populates_capabilities() {
|
||||
let registry = registry_with_caps();
|
||||
let d = dispatcher(®istry);
|
||||
let conn = Arc::new(CallConnection::new(stub_connection()));
|
||||
let response = dispatch(&d, &conn, "pub/run").await;
|
||||
let out = response.result.expect("ok");
|
||||
assert_eq!(
|
||||
out["has_google_capability"],
|
||||
serde_json::json!(true),
|
||||
"an External op's call must populate capabilities for the handler"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn unknown_op_returns_not_found() {
|
||||
let registry = Arc::new(OperationRegistry::new());
|
||||
let d = dispatcher(®istry);
|
||||
let conn = Arc::new(CallConnection::new(stub_connection()));
|
||||
let response = dispatch(&d, &conn, "no/such").await;
|
||||
match response.result {
|
||||
Err(e) => assert_eq!(e.code, "NOT_FOUND"),
|
||||
other => panic!("expected NOT_FOUND, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn spawn_dispatch_returns_live_call_connection() {
|
||||
let registry = registry_with_caps();
|
||||
let client = CallClient::new(Arc::clone(®istry), Arc::new(NoopIdentityProvider));
|
||||
let conn = client.spawn_dispatch(stub_connection());
|
||||
assert_eq!(
|
||||
conn.connection()
|
||||
.expect("quic connection present")
|
||||
.remote_alpn(),
|
||||
b"alknet/call"
|
||||
);
|
||||
std::mem::drop(conn);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn call_client_is_send_sync() {
|
||||
fn assert_send_sync<T: Send + Sync>() {}
|
||||
assert_send_sync::<CallClient>();
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,107 @@
|
||||
//! Client adapters: turn external operation sources (JSON Schema, OpenAPI,
|
||||
//! MCP, remote `from_call` peers) into `HandlerRegistration` bundles.
|
||||
//!
|
||||
//! See `docs/architecture/crates/call/client-and-adapters.md` for the
|
||||
//! OperationAdapter trait and the Adapter Location Map, and
|
||||
//! `docs/architecture/decisions/017-call-protocol-client-and-adapter-contract.md`
|
||||
//! §5 for the trait contract.
|
||||
|
||||
mod call_client;
|
||||
mod from_call;
|
||||
|
||||
pub use call_client::CallClient;
|
||||
pub use from_call::{from_call, FromCallConfig};
|
||||
|
||||
use crate::registry::registration::HandlerRegistration;
|
||||
|
||||
/// Errors produced by [`OperationAdapter::import`].
|
||||
///
|
||||
/// The variant set is the v1 default (two-way-door remainder, OQ-26);
|
||||
/// `#[non_exhaustive]` lets downstream adapters (e.g. `alknet-http`'s
|
||||
/// `from_openapi`/`from_mcp`) extend without breaking match arms. All
|
||||
/// payloads are string messages — kept simple and `Send + Sync` by
|
||||
/// construction.
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
#[non_exhaustive]
|
||||
pub enum AdapterError {
|
||||
/// `from_call` remote unreachable / `services/list` failed.
|
||||
#[error("discovery failed: {message}")]
|
||||
DiscoveryFailed { message: String },
|
||||
|
||||
/// `from_openapi` / `from_jsonschema` couldn't parse the spec.
|
||||
#[error("schema parse error: {message}")]
|
||||
SchemaParse { message: String },
|
||||
|
||||
/// Underlying transport error (QUIC for `from_call`, HTTP for adapters).
|
||||
#[error("transport error: {message}")]
|
||||
Transport { message: String },
|
||||
|
||||
/// HTTP 401 for `from_openapi`/`from_mcp`, auth rejected for `from_call`.
|
||||
#[error("unauthorized: {message}")]
|
||||
Unauthorized { message: String },
|
||||
|
||||
/// Same-peer namespace collision in `from_call` (ADR-029 §5; OQ-26).
|
||||
/// Cross-peer collision dissolves (same name on different peers lives in
|
||||
/// separate sub-overlays); same-peer collision stays an error — a peer
|
||||
/// shouldn't expose two ops with the same name.
|
||||
#[error("same-peer collision: {message}")]
|
||||
SamePeerCollision { message: String },
|
||||
}
|
||||
|
||||
/// Import a set of operations as `HandlerRegistration` bundles.
|
||||
///
|
||||
/// Async because `from_call` requires async discovery (`services/list` +
|
||||
/// `services/schema` over a QUIC connection); sync adapters (e.g.
|
||||
/// `from_openapi` reading a static spec) trivially satisfy
|
||||
/// an async trait — their `import()` bodies contain no `.await` points.
|
||||
///
|
||||
/// See ADR-017 §5 (`docs/architecture/decisions/017-call-protocol-client-and-adapter-contract.md`)
|
||||
/// and `docs/architecture/crates/call/client-and-adapters.md`.
|
||||
#[async_trait::async_trait]
|
||||
pub trait OperationAdapter: Send + Sync {
|
||||
async fn import(&self) -> Result<Vec<HandlerRegistration>, AdapterError>;
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
struct OkAdapter;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl OperationAdapter for OkAdapter {
|
||||
async fn import(&self) -> Result<Vec<HandlerRegistration>, AdapterError> {
|
||||
Ok(vec![])
|
||||
}
|
||||
}
|
||||
|
||||
struct ErrAdapter;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl OperationAdapter for ErrAdapter {
|
||||
async fn import(&self) -> Result<Vec<HandlerRegistration>, AdapterError> {
|
||||
Err(AdapterError::SchemaParse {
|
||||
message: "x".into(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ok_adapter_imports_empty() {
|
||||
let adapter = OkAdapter;
|
||||
match adapter.import().await {
|
||||
Ok(bundles) => assert!(bundles.is_empty()),
|
||||
Err(e) => panic!("expected Ok, got Err: {e}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn err_adapter_returns_schema_parse() {
|
||||
let adapter = ErrAdapter;
|
||||
match adapter.import().await {
|
||||
Ok(_) => panic!("expected Err"),
|
||||
Err(AdapterError::SchemaParse { message }) => assert_eq!(message, "x"),
|
||||
Err(other) => panic!("expected SchemaParse, got {other}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
//! Authentication primitives: `AuthContext`, `Identity`, `AuthToken`,
|
||||
//! `IdentityProvider`.
|
||||
//!
|
||||
//! See `docs/architecture/` for the full specification. The trait-based
|
||||
//! `IdentityProvider` is the integration point — the assembly layer supplies
|
||||
//! the impl (config-backed, vault-backed, or persistence-adapter-backed) and
|
||||
//! the call protocol resolves identity per-request through it.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::net::SocketAddr;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct Identity {
|
||||
pub id: String,
|
||||
pub scopes: Vec<String>,
|
||||
pub resources: HashMap<String, Vec<String>>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct AuthToken {
|
||||
pub raw: Vec<u8>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct AuthContext {
|
||||
pub identity: Option<Identity>,
|
||||
pub alpn: Vec<u8>,
|
||||
pub remote_addr: Option<SocketAddr>,
|
||||
pub tls_client_fingerprint: Option<String>,
|
||||
}
|
||||
|
||||
impl AuthContext {
|
||||
/// Construct an `AuthContext` with no identity, no fingerprint, and no
|
||||
/// remote address — only the ALPN is set. For POCs, tests, and handlers
|
||||
/// that don't require auth.
|
||||
pub fn anonymous(alpn: impl Into<Vec<u8>>) -> Self {
|
||||
Self {
|
||||
identity: None,
|
||||
alpn: alpn.into(),
|
||||
remote_addr: None,
|
||||
tls_client_fingerprint: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub trait IdentityProvider: Send + Sync + 'static {
|
||||
fn resolve_from_fingerprint(&self, fingerprint: &str) -> Option<Identity>;
|
||||
fn resolve_from_token(&self, token: &AuthToken) -> Option<Identity>;
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn identity_fields_and_equality() {
|
||||
let mut resources = HashMap::new();
|
||||
resources.insert("service".to_string(), vec!["gitea".to_string()]);
|
||||
let id = Identity {
|
||||
id: "SHA256:abc123".to_string(),
|
||||
scopes: vec!["relay:connect".to_string()],
|
||||
resources,
|
||||
};
|
||||
let id2 = id.clone();
|
||||
assert_eq!(id, id2);
|
||||
assert_eq!(id.id, "SHA256:abc123");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_token_is_clone() {
|
||||
let token = AuthToken {
|
||||
raw: b"alk_test".to_vec(),
|
||||
};
|
||||
let cloned = token.clone();
|
||||
assert_eq!(token.raw, cloned.raw);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_context_is_clone() {
|
||||
let ctx = AuthContext {
|
||||
identity: None,
|
||||
alpn: b"alknet/test".to_vec(),
|
||||
remote_addr: None,
|
||||
tls_client_fingerprint: None,
|
||||
};
|
||||
let cloned = ctx.clone();
|
||||
assert_eq!(cloned.alpn, b"alknet/test");
|
||||
assert!(cloned.identity.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_context_anonymous_sets_alpn_only() {
|
||||
let ctx = AuthContext::anonymous(b"alknet/test");
|
||||
assert_eq!(ctx.alpn, b"alknet/test");
|
||||
assert!(ctx.identity.is_none());
|
||||
assert!(ctx.remote_addr.is_none());
|
||||
assert!(ctx.tls_client_fingerprint.is_none());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
//! Vendored core types: `ProtocolHandler`, `Connection`, `BiStream`,
|
||||
//! `BidiStreamSource`, `AuthContext`, `IdentityProvider`, `Identity`,
|
||||
//! `AuthToken`, `Capabilities`, `OwnershipProvider`, `HandlerError`,
|
||||
//! `StreamError`.
|
||||
//!
|
||||
//! These types are the home for the former `alknet-core` surface — lean,
|
||||
//! no TLS, no transport coupling, no endpoint/accept-loop. The dial and
|
||||
//! the TLS config are concerns of the consumer, not of this crate. When
|
||||
//! the alknet mono-repo is reworked, it will consume alkcall's versions.
|
||||
|
||||
pub mod auth;
|
||||
pub mod ownership;
|
||||
pub mod types;
|
||||
|
||||
pub use auth::{AuthContext, AuthToken, Identity, IdentityProvider};
|
||||
pub use ownership::{InMemoryOwnershipStore, OwnershipError, OwnershipProvider, OwnershipStore};
|
||||
pub use types::{
|
||||
BiStream, BidiStreamSource, Capabilities, Connection, HandlerError, IdentityAlreadySet,
|
||||
ProtocolHandler, RecvStream, Secret, SendStream, StreamError,
|
||||
};
|
||||
@@ -0,0 +1,276 @@
|
||||
//! Ownership store: `OwnershipProvider` (sync read trait), `OwnershipStore`
|
||||
//! (async write trait), `InMemoryOwnershipStore` default adapter, and
|
||||
//! `OwnershipError`.
|
||||
//!
|
||||
//! Runtime-spawned resources (containers, TTYs, workspace processes) have
|
||||
//! derived ownership — whoever spawned the resource owns it. The static
|
||||
//! `Identity.resources` model can't represent this (the resource didn't
|
||||
//! exist when the identity was resolved), so `AccessControl::check`
|
||||
//! consults `OwnershipProvider` at check time.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::RwLock;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::auth::Identity;
|
||||
|
||||
#[non_exhaustive]
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum OwnershipError {
|
||||
#[error("backend error: {message}")]
|
||||
Backend { message: String },
|
||||
#[error("not found: {entity}")]
|
||||
NotFound { entity: String },
|
||||
}
|
||||
|
||||
/// Read side: consulted by `AccessControl::check` on the dispatch hot path.
|
||||
/// Sync — called in the dispatch loop, no `.await`.
|
||||
pub trait OwnershipProvider: Send + Sync + 'static {
|
||||
/// Does `identity` own `resource_type/resource_id` with `action`?
|
||||
/// The `action` parameter is accepted but not gated on — the base model
|
||||
/// is "owner can do anything they own." Per-action grants are a future
|
||||
/// extension; this preserves the door without building the mechanism.
|
||||
fn owns(
|
||||
&self,
|
||||
identity: &Identity,
|
||||
resource_type: &str,
|
||||
resource_id: &str,
|
||||
action: &str,
|
||||
) -> bool;
|
||||
|
||||
/// What resources of `resource_type` does `identity` own? Returns the
|
||||
/// set of resource IDs the caller owns, for the handler to filter
|
||||
/// against (the result-filter path).
|
||||
fn owned_resources(&self, identity: &Identity, resource_type: &str) -> Vec<String>;
|
||||
|
||||
/// Does `identity` own *any* resource of `resource_type`? The scope-gate
|
||||
/// path.
|
||||
fn owns_any(&self, identity: &Identity, resource_type: &str) -> bool;
|
||||
}
|
||||
|
||||
/// Write side: called by the handler that manages the resource lifecycle.
|
||||
/// Async — not on the dispatch hot path. The handler calls `record` on
|
||||
/// spawn and `revoke` on teardown (handler-driven, not a reaper). The trait
|
||||
/// takes `&self` so it can be shared as `Arc<dyn OwnershipStore>` (interior
|
||||
/// mutability via `RwLock`).
|
||||
#[async_trait]
|
||||
pub trait OwnershipStore: Send + Sync + 'static {
|
||||
/// Record that `identity` spawned `resource_type/resource_id`.
|
||||
async fn record(
|
||||
&self,
|
||||
identity: &Identity,
|
||||
resource_type: &str,
|
||||
resource_id: &str,
|
||||
) -> Result<(), OwnershipError>;
|
||||
|
||||
/// Revoke ownership of `resource_type/resource_id`. Called by the
|
||||
/// handler on resource teardown.
|
||||
async fn revoke(&self, resource_type: &str, resource_id: &str) -> Result<(), OwnershipError>;
|
||||
}
|
||||
|
||||
pub struct InMemoryOwnershipStore {
|
||||
inner: RwLock<HashMap<(String, String), Identity>>,
|
||||
}
|
||||
|
||||
impl InMemoryOwnershipStore {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
inner: RwLock::new(HashMap::new()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for InMemoryOwnershipStore {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl OwnershipProvider for InMemoryOwnershipStore {
|
||||
fn owns(
|
||||
&self,
|
||||
identity: &Identity,
|
||||
resource_type: &str,
|
||||
resource_id: &str,
|
||||
_action: &str,
|
||||
) -> bool {
|
||||
let inner = self.inner.read().unwrap_or_else(|e| e.into_inner());
|
||||
inner
|
||||
.get(&(resource_type.to_string(), resource_id.to_string()))
|
||||
.map(|owner| owner.id == identity.id)
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
fn owned_resources(&self, identity: &Identity, resource_type: &str) -> Vec<String> {
|
||||
let inner = self.inner.read().unwrap_or_else(|e| e.into_inner());
|
||||
inner
|
||||
.iter()
|
||||
.filter(|((rt, _), owner)| rt == resource_type && owner.id == identity.id)
|
||||
.map(|((_, rid), _)| rid.clone())
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn owns_any(&self, identity: &Identity, resource_type: &str) -> bool {
|
||||
let inner = self.inner.read().unwrap_or_else(|e| e.into_inner());
|
||||
inner
|
||||
.iter()
|
||||
.any(|((rt, _), owner)| rt == resource_type && owner.id == identity.id)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl OwnershipStore for InMemoryOwnershipStore {
|
||||
async fn record(
|
||||
&self,
|
||||
identity: &Identity,
|
||||
resource_type: &str,
|
||||
resource_id: &str,
|
||||
) -> Result<(), OwnershipError> {
|
||||
let mut inner = self.inner.write().unwrap_or_else(|e| e.into_inner());
|
||||
inner.insert(
|
||||
(resource_type.to_string(), resource_id.to_string()),
|
||||
identity.clone(),
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Revoking a non-existent resource is a no-op: returns `Ok(())`.
|
||||
/// Teardown paths are idempotent — a handler may call `revoke` on a
|
||||
/// resource that was already removed.
|
||||
async fn revoke(&self, resource_type: &str, resource_id: &str) -> Result<(), OwnershipError> {
|
||||
let mut inner = self.inner.write().unwrap_or_else(|e| e.into_inner());
|
||||
inner.remove(&(resource_type.to_string(), resource_id.to_string()));
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn make_identity(id: &str) -> Identity {
|
||||
Identity {
|
||||
id: id.to_string(),
|
||||
scopes: vec![],
|
||||
resources: HashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn record_owns_revoke_round_trip() {
|
||||
let store = InMemoryOwnershipStore::new();
|
||||
let owner = make_identity("worker-a");
|
||||
|
||||
assert!(!store.owns(&owner, "container", "c1", "exec"));
|
||||
assert!(!store.owns_any(&owner, "container"));
|
||||
assert!(store.owned_resources(&owner, "container").is_empty());
|
||||
|
||||
store.record(&owner, "container", "c1").await.unwrap();
|
||||
assert!(store.owns(&owner, "container", "c1", "exec"));
|
||||
assert!(store.owns(&owner, "container", "c1", "logs"));
|
||||
assert!(store.owns_any(&owner, "container"));
|
||||
assert_eq!(store.owned_resources(&owner, "container"), vec!["c1"]);
|
||||
|
||||
store.revoke("container", "c1").await.unwrap();
|
||||
assert!(!store.owns(&owner, "container", "c1", "exec"));
|
||||
assert!(!store.owns_any(&owner, "container"));
|
||||
assert!(store.owned_resources(&owner, "container").is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn owned_resources_returns_all_for_owner_with_multiple() {
|
||||
let store = InMemoryOwnershipStore::new();
|
||||
let owner = make_identity("worker-a");
|
||||
store.record(&owner, "container", "c1").await.unwrap();
|
||||
store.record(&owner, "container", "c2").await.unwrap();
|
||||
store.record(&owner, "container", "c3").await.unwrap();
|
||||
|
||||
let mut owned = store.owned_resources(&owner, "container");
|
||||
owned.sort();
|
||||
assert_eq!(owned, vec!["c1", "c2", "c3"]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn owned_resources_filters_by_resource_type() {
|
||||
let store = InMemoryOwnershipStore::new();
|
||||
let owner = make_identity("worker-a");
|
||||
store.record(&owner, "container", "c1").await.unwrap();
|
||||
store.record(&owner, "tty", "t1").await.unwrap();
|
||||
|
||||
let owned_containers = store.owned_resources(&owner, "container");
|
||||
assert_eq!(owned_containers, vec!["c1"]);
|
||||
let owned_ttys = store.owned_resources(&owner, "tty");
|
||||
assert_eq!(owned_ttys, vec!["t1"]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn owns_any_returns_false_for_owner_with_no_resources_of_type() {
|
||||
let store = InMemoryOwnershipStore::new();
|
||||
let owner = make_identity("worker-a");
|
||||
store.record(&owner, "container", "c1").await.unwrap();
|
||||
|
||||
assert!(store.owns_any(&owner, "container"));
|
||||
assert!(!store.owns_any(&owner, "tty"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn revoke_on_non_existent_resource_is_no_op() {
|
||||
let store = InMemoryOwnershipStore::new();
|
||||
store.revoke("container", "never-existed").await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn owns_returns_false_for_different_identity() {
|
||||
let store = InMemoryOwnershipStore::new();
|
||||
let owner = make_identity("worker-a");
|
||||
let other = make_identity("worker-b");
|
||||
store.record(&owner, "container", "c1").await.unwrap();
|
||||
|
||||
assert!(store.owns(&owner, "container", "c1", "exec"));
|
||||
assert!(!store.owns(&other, "container", "c1", "exec"));
|
||||
assert!(!store.owns_any(&other, "container"));
|
||||
assert!(store.owned_resources(&other, "container").is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn record_replaces_existing_owner() {
|
||||
let store = InMemoryOwnershipStore::new();
|
||||
let owner_a = make_identity("worker-a");
|
||||
let owner_b = make_identity("worker-b");
|
||||
store.record(&owner_a, "container", "c1").await.unwrap();
|
||||
store.record(&owner_b, "container", "c1").await.unwrap();
|
||||
|
||||
assert!(!store.owns(&owner_a, "container", "c1", "exec"));
|
||||
assert!(store.owns(&owner_b, "container", "c1", "exec"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn default_is_empty_store() {
|
||||
let store = InMemoryOwnershipStore::default();
|
||||
let owner = make_identity("worker-a");
|
||||
assert!(store.owned_resources(&owner, "container").is_empty());
|
||||
assert!(!store.owns_any(&owner, "container"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ownership_error_display_formatting() {
|
||||
let backend = OwnershipError::Backend {
|
||||
message: "disk full".to_string(),
|
||||
};
|
||||
assert_eq!(backend.to_string(), "backend error: disk full");
|
||||
|
||||
let not_found = OwnershipError::NotFound {
|
||||
entity: "container:c1".to_string(),
|
||||
};
|
||||
assert_eq!(not_found.to_string(), "not found: container:c1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ownership_error_is_non_exhaustive() {
|
||||
let err = OwnershipError::Backend {
|
||||
message: "x".to_string(),
|
||||
};
|
||||
let _ = err.to_string();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,943 @@
|
||||
//! Core types: `ProtocolHandler`, `HandlerError`, `Connection`, `BiStream`,
|
||||
//! `SendStream`, `RecvStream`, `StreamError`, `BidiStreamSource`,
|
||||
//! `Capabilities`, `Secret`.
|
||||
//!
|
||||
//! See `docs/architecture/` for the full specification. These types are the
|
||||
//! home for the former `alknet-core` surface — lean, no TLS, no transport
|
||||
//! coupling, no endpoint/accept-loop. The dial and the TLS config are
|
||||
//! concerns of the consumer, not of this crate.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::io;
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::{Mutex, OnceLock};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use tokio::io::{AsyncRead, AsyncWrite};
|
||||
use zeroize::{Zeroize, ZeroizeOnDrop};
|
||||
|
||||
use super::auth::{AuthContext, Identity};
|
||||
|
||||
pub struct Secret<T: Zeroize + Clone> {
|
||||
inner: T,
|
||||
}
|
||||
|
||||
impl<T: Zeroize + Clone> Secret<T> {
|
||||
pub fn new(value: T) -> Self {
|
||||
Self { inner: value }
|
||||
}
|
||||
|
||||
pub fn expose_secret(&self) -> &T {
|
||||
&self.inner
|
||||
}
|
||||
}
|
||||
|
||||
impl<T: Zeroize + Clone> Clone for Secret<T> {
|
||||
fn clone(&self) -> Self {
|
||||
Self {
|
||||
inner: self.inner.clone(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<T: Zeroize + Clone> Zeroize for Secret<T> {
|
||||
fn zeroize(&mut self) {
|
||||
self.inner.zeroize();
|
||||
}
|
||||
}
|
||||
|
||||
impl<T: Zeroize + Clone> Drop for Secret<T> {
|
||||
fn drop(&mut self) {
|
||||
self.inner.zeroize();
|
||||
}
|
||||
}
|
||||
|
||||
impl<T: Zeroize + Clone> std::fmt::Debug for Secret<T> {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.write_str("[REDACTED]")
|
||||
}
|
||||
}
|
||||
|
||||
pub struct Capabilities {
|
||||
entries: HashMap<String, Secret<String>>,
|
||||
}
|
||||
|
||||
impl Zeroize for Capabilities {
|
||||
fn zeroize(&mut self) {
|
||||
for (_, v) in self.entries.iter_mut() {
|
||||
v.zeroize();
|
||||
}
|
||||
self.entries.clear();
|
||||
}
|
||||
}
|
||||
|
||||
impl ZeroizeOnDrop for Capabilities {}
|
||||
|
||||
impl Clone for Capabilities {
|
||||
fn clone(&self) -> Self {
|
||||
Self {
|
||||
entries: self.entries.clone(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Capabilities {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
entries: HashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_api_key(mut self, service: &str, key: String) -> Self {
|
||||
self.entries
|
||||
.insert(format!("api_key:{service}"), Secret::new(key));
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_http_token(mut self, service: &str, token: String) -> Self {
|
||||
self.entries
|
||||
.insert(format!("http_token:{service}"), Secret::new(token));
|
||||
self
|
||||
}
|
||||
|
||||
pub fn get(&self, service: &str) -> Option<&Secret<String>> {
|
||||
self.entries
|
||||
.get(&format!("api_key:{service}"))
|
||||
.or_else(|| self.entries.get(&format!("http_token:{service}")))
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for Capabilities {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for Capabilities {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("Capabilities")
|
||||
.field("entries", &format!("[{} redacted]", self.entries.len()))
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum IdentityAlreadySet {
|
||||
#[error("connection identity already set")]
|
||||
AlreadySet,
|
||||
}
|
||||
|
||||
pub enum HandlerError {
|
||||
ConnectionClosed,
|
||||
StreamError(io::Error),
|
||||
AuthRequired,
|
||||
Internal(Box<dyn std::error::Error + Send + Sync>),
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for HandlerError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::ConnectionClosed => f.write_str("HandlerError::ConnectionClosed"),
|
||||
Self::StreamError(e) => f.debug_tuple("HandlerError::StreamError").field(e).finish(),
|
||||
Self::AuthRequired => f.write_str("HandlerError::AuthRequired"),
|
||||
Self::Internal(e) => f.debug_tuple("HandlerError::Internal").field(e).finish(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Display for HandlerError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::ConnectionClosed => f.write_str("connection closed"),
|
||||
Self::StreamError(e) => write!(f, "stream error: {e}"),
|
||||
Self::AuthRequired => f.write_str("authentication required"),
|
||||
Self::Internal(e) => write!(f, "internal handler error: {e}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for HandlerError {
|
||||
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
|
||||
match self {
|
||||
Self::StreamError(e) => Some(e),
|
||||
Self::Internal(e) => Some(e.as_ref()),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub enum StreamError {
|
||||
ConnectionClosed,
|
||||
StreamClosed,
|
||||
Timeout,
|
||||
Internal(io::Error),
|
||||
}
|
||||
|
||||
impl From<StreamError> for HandlerError {
|
||||
fn from(e: StreamError) -> Self {
|
||||
match e {
|
||||
StreamError::ConnectionClosed => HandlerError::ConnectionClosed,
|
||||
StreamError::StreamClosed => HandlerError::StreamError(io::Error::new(
|
||||
io::ErrorKind::ConnectionReset,
|
||||
"stream closed",
|
||||
)),
|
||||
StreamError::Timeout => HandlerError::StreamError(io::Error::new(
|
||||
io::ErrorKind::TimedOut,
|
||||
"stream timed out",
|
||||
)),
|
||||
StreamError::Internal(e) => HandlerError::StreamError(e),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for StreamError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::ConnectionClosed => f.write_str("StreamError::ConnectionClosed"),
|
||||
Self::StreamClosed => f.write_str("StreamError::StreamClosed"),
|
||||
Self::Timeout => f.write_str("StreamError::Timeout"),
|
||||
Self::Internal(e) => f.debug_tuple("StreamError::Internal").field(e).finish(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Display for StreamError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::ConnectionClosed => f.write_str("connection closed"),
|
||||
Self::StreamClosed => f.write_str("stream closed"),
|
||||
Self::Timeout => f.write_str("stream timed out"),
|
||||
Self::Internal(e) => write!(f, "stream error: {e}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for StreamError {
|
||||
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
|
||||
match self {
|
||||
Self::Internal(e) => Some(e),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait ProtocolHandler: Send + Sync + 'static {
|
||||
fn alpn(&self) -> &'static [u8];
|
||||
async fn handle(&self, connection: Connection, auth: &AuthContext) -> Result<(), HandlerError>;
|
||||
}
|
||||
|
||||
// --- BiStream: the handler leaf (ADR-092) ---------------------------------
|
||||
|
||||
/// Internal helper trait — the union of `AsyncRead + AsyncWrite + Send`.
|
||||
/// Not public; exists only to give `BiStream` a single boxed field.
|
||||
trait AsyncReadWrite: AsyncRead + AsyncWrite + Send {}
|
||||
impl<T: AsyncRead + AsyncWrite + Send> AsyncReadWrite for T {}
|
||||
|
||||
/// The handler leaf — a bidirectional byte stream (ADR-092).
|
||||
///
|
||||
/// `accept_bi`/`open_bi` return a `BiStream`, not a split
|
||||
/// `(SendStream, RecvStream)` pair. Handlers that want the split halves call
|
||||
/// `tokio::io::split(&mut *stream)` (the stdlib idiom `tokio::io::split`
|
||||
/// already provides for `TcpStream` and `TlsStream<TcpStream>`). The
|
||||
/// split is a stdlib call at the handler boundary, not a per-handler trait
|
||||
/// wrapper.
|
||||
///
|
||||
/// `BiStream: AsyncRead + AsyncWrite + Send + Unpin` by construction.
|
||||
pub struct BiStream {
|
||||
inner: Box<dyn AsyncReadWrite + Unpin>,
|
||||
}
|
||||
|
||||
impl BiStream {
|
||||
/// Join a read half and a write half into a single `BiStream`. The join
|
||||
/// happens once, in the `BidiStreamSource` impl — handlers receive the
|
||||
/// joined `BiStream` and never see the pair.
|
||||
///
|
||||
/// Public so that downstream crates (the channels reassembly path, tests
|
||||
/// that construct a `BiStream` from independent halves) can join their
|
||||
/// own halves. The rule this normalizes: **the split never crosses a
|
||||
/// crate boundary as part of a constructor** — `Connection::from_bidi`
|
||||
/// takes a joined `BiStream`, and `BiStream::from_joined` is the join.
|
||||
pub fn from_joined<R, W>(reader: R, writer: W) -> Self
|
||||
where
|
||||
R: AsyncRead + Send + Unpin + 'static,
|
||||
W: AsyncWrite + Send + Unpin + 'static,
|
||||
{
|
||||
Self {
|
||||
inner: Box::new(tokio::io::join(reader, writer)),
|
||||
}
|
||||
}
|
||||
|
||||
/// Wrap a single value that is already `AsyncRead + AsyncWrite` (e.g.
|
||||
/// `tokio::io::DuplexStream`, `TlsStream<TcpStream>`,
|
||||
/// `russh::Channel::into_stream()`). Used by the single-stream
|
||||
/// `BidiStreamSource` impl and by `Connection::from_bidi`.
|
||||
pub(crate) fn from_bidi<S>(stream: S) -> Self
|
||||
where
|
||||
S: AsyncRead + AsyncWrite + Send + Unpin + 'static,
|
||||
{
|
||||
Self {
|
||||
inner: Box::new(stream),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl AsyncRead for BiStream {
|
||||
fn poll_read(
|
||||
mut self: std::pin::Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
buf: &mut tokio::io::ReadBuf<'_>,
|
||||
) -> std::task::Poll<io::Result<()>> {
|
||||
std::pin::Pin::new(self.inner.as_mut()).poll_read(cx, buf)
|
||||
}
|
||||
}
|
||||
|
||||
impl AsyncWrite for BiStream {
|
||||
fn poll_write(
|
||||
mut self: std::pin::Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
buf: &[u8],
|
||||
) -> std::task::Poll<io::Result<usize>> {
|
||||
std::pin::Pin::new(self.inner.as_mut()).poll_write(cx, buf)
|
||||
}
|
||||
|
||||
fn poll_flush(
|
||||
mut self: std::pin::Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
) -> std::task::Poll<io::Result<()>> {
|
||||
std::pin::Pin::new(self.inner.as_mut()).poll_flush(cx)
|
||||
}
|
||||
|
||||
fn poll_shutdown(
|
||||
mut self: std::pin::Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
) -> std::task::Poll<io::Result<()>> {
|
||||
std::pin::Pin::new(self.inner.as_mut()).poll_shutdown(cx)
|
||||
}
|
||||
}
|
||||
|
||||
// --- SendStream / RecvStream: thin newtypes (ADR-092) ---------------------
|
||||
|
||||
pub struct SendStream {
|
||||
inner: Box<dyn AsyncWrite + Send + Unpin>,
|
||||
}
|
||||
|
||||
pub struct RecvStream {
|
||||
inner: Box<dyn AsyncRead + Send + Unpin>,
|
||||
}
|
||||
|
||||
impl SendStream {
|
||||
/// Box a write half into the thin `SendStream` newtype. Used by
|
||||
/// `into_sub_streams()` (ADR-074) and the channels reassembly path.
|
||||
/// Not a constructor that feeds `Connection` — the split never crosses
|
||||
/// a crate boundary as part of a constructor (ADR-092).
|
||||
pub fn from_stream(stream: impl AsyncWrite + Send + Unpin + 'static) -> Self {
|
||||
Self {
|
||||
inner: Box::new(stream),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl RecvStream {
|
||||
/// Box a read half into the thin `RecvStream` newtype. Used by
|
||||
/// `into_sub_streams()` (ADR-074) and the channels reassembly path.
|
||||
/// Not a constructor that feeds `Connection` — the split never crosses
|
||||
/// a crate boundary as part of a constructor (ADR-092).
|
||||
pub fn from_stream(stream: impl AsyncRead + Send + Unpin + 'static) -> Self {
|
||||
Self {
|
||||
inner: Box::new(stream),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl AsyncWrite for SendStream {
|
||||
fn poll_write(
|
||||
mut self: std::pin::Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
buf: &[u8],
|
||||
) -> std::task::Poll<io::Result<usize>> {
|
||||
std::pin::Pin::new(self.inner.as_mut()).poll_write(cx, buf)
|
||||
}
|
||||
|
||||
fn poll_flush(
|
||||
mut self: std::pin::Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
) -> std::task::Poll<io::Result<()>> {
|
||||
std::pin::Pin::new(self.inner.as_mut()).poll_flush(cx)
|
||||
}
|
||||
|
||||
fn poll_shutdown(
|
||||
mut self: std::pin::Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
) -> std::task::Poll<io::Result<()>> {
|
||||
std::pin::Pin::new(self.inner.as_mut()).poll_shutdown(cx)
|
||||
}
|
||||
}
|
||||
|
||||
impl AsyncRead for RecvStream {
|
||||
fn poll_read(
|
||||
mut self: std::pin::Pin<&mut Self>,
|
||||
cx: &mut std::task::Context<'_>,
|
||||
buf: &mut tokio::io::ReadBuf<'_>,
|
||||
) -> std::task::Poll<io::Result<()>> {
|
||||
std::pin::Pin::new(self.inner.as_mut()).poll_read(cx, buf)
|
||||
}
|
||||
}
|
||||
|
||||
/// Yield bidirectional streams to a `Connection`. Downstream crates implement
|
||||
/// this trait to add connection shapes (channels, a future transport, a test
|
||||
/// double beyond the single-stream case) without editing core. See ADR-070
|
||||
/// for the full rationale and ADR-065 for the yield-once contract the
|
||||
/// `StreamBidiStreamSource` impl preserves. The return type is `BiStream`
|
||||
/// (ADR-092) — the join happens once, in the impl, not per-handler.
|
||||
#[async_trait]
|
||||
pub trait BidiStreamSource: Send + Sync + 'static {
|
||||
/// Yield the next bidirectional stream this connection provides.
|
||||
///
|
||||
/// Transport semantics (carried from ADR-065):
|
||||
/// - QUIC (quinn/iroh): returns a new bidi stream on each call,
|
||||
/// `ConnectionClosed` when the underlying connection closes.
|
||||
/// - Single-stream (TCP+TLS, SSH channel, WebTransport stream, wasm):
|
||||
/// yields the underlying stream on the first call, then
|
||||
/// `ConnectionClosed` on all subsequent calls.
|
||||
/// - Channels: yields one bidi stream per channel, `ConnectionClosed`
|
||||
/// when the channels connection closes.
|
||||
async fn accept_bi(&self) -> Result<BiStream, StreamError>;
|
||||
|
||||
/// Open a bidirectional stream to the peer.
|
||||
///
|
||||
/// Single-stream sources return `StreamClosed` (a single stream cannot
|
||||
/// open new application streams — ADR-065). QUIC and channels sources
|
||||
/// open new streams.
|
||||
async fn open_bi(&self) -> Result<BiStream, StreamError>;
|
||||
|
||||
/// The peer's address, if available. Informational (NAT/proxy).
|
||||
fn remote_addr(&self) -> Option<SocketAddr>;
|
||||
|
||||
/// Close the connection. The `code`/`reason` args are QUIC application-
|
||||
/// level close codes; non-QUIC sources ignore them (the drop is the
|
||||
/// close — ADR-065 §"Negative"). See ADR-070 §"REQ-CORE-02" for the
|
||||
/// rationale for keeping the QUIC-shaped signature on the trait.
|
||||
fn close(&self, code: u32, reason: &str);
|
||||
}
|
||||
|
||||
/// Single-stream `BidiStreamSource` (TCP+TLS, SSH channel, WebTransport
|
||||
/// stream, wasm stream — ADR-065). Crate-private; constructed via
|
||||
/// `Connection::from_bidi`. `accept_bi` yields the underlying `BiStream`
|
||||
/// once, then `ConnectionClosed`; `open_bi` returns `StreamClosed`.
|
||||
struct StreamBidiStreamSource {
|
||||
stream: Mutex<Option<BiStream>>,
|
||||
remote_addr: Option<SocketAddr>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl BidiStreamSource for StreamBidiStreamSource {
|
||||
async fn accept_bi(&self) -> Result<BiStream, StreamError> {
|
||||
let mut guard = self.stream.lock().expect("stream mutex poisoned");
|
||||
match guard.take() {
|
||||
Some(stream) => Ok(stream),
|
||||
None => Err(StreamError::ConnectionClosed),
|
||||
}
|
||||
}
|
||||
|
||||
async fn open_bi(&self) -> Result<BiStream, StreamError> {
|
||||
Err(StreamError::StreamClosed)
|
||||
}
|
||||
|
||||
fn remote_addr(&self) -> Option<SocketAddr> {
|
||||
self.remote_addr
|
||||
}
|
||||
|
||||
/// `code`/`reason` are ignored: a single stream has no QUIC-shaped
|
||||
/// application-level close codes. The drop is the close (ADR-065
|
||||
/// §"Negative"). The `_` prefix is intentional — the signature matches
|
||||
/// the public `Connection::close` API (ADR-070 §"REQ-CORE-02").
|
||||
fn close(&self, _code: u32, _reason: &str) {
|
||||
let _ = self.stream.lock().expect("stream mutex poisoned").take();
|
||||
}
|
||||
}
|
||||
|
||||
pub struct Connection {
|
||||
source: Box<dyn BidiStreamSource>,
|
||||
alpn: Vec<u8>,
|
||||
identity: OnceLock<Identity>,
|
||||
}
|
||||
|
||||
impl Connection {
|
||||
/// Construct a `Connection` from a single bidirectional stream (e.g.
|
||||
/// `tokio::io::DuplexStream`, `TlsStream<TcpStream>`,
|
||||
/// `russh::Channel::into_stream()`). The stream is wrapped in a
|
||||
/// `BiStream` (ADR-092) and yielded by `accept_bi` once, then
|
||||
/// `ConnectionClosed`. `open_bi` returns `StreamClosed` (a single
|
||||
/// stream can't open new application streams — ADR-065).
|
||||
///
|
||||
/// This is the only public stream constructor (ADR-092): the split
|
||||
/// never crosses a crate boundary as part of a constructor. Handlers
|
||||
/// that want the split halves call `tokio::io::split(&mut *stream)` on
|
||||
/// the `BiStream` they receive from `accept_bi`.
|
||||
pub fn from_bidi(
|
||||
stream: impl AsyncRead + AsyncWrite + Send + Unpin + 'static,
|
||||
alpn: Vec<u8>,
|
||||
remote_addr: Option<SocketAddr>,
|
||||
) -> Self {
|
||||
Self {
|
||||
source: Box::new(StreamBidiStreamSource {
|
||||
stream: Mutex::new(Some(BiStream::from_bidi(stream))),
|
||||
remote_addr,
|
||||
}),
|
||||
alpn,
|
||||
identity: OnceLock::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Construct from a caller-supplied `BidiStreamSource` impl. The
|
||||
/// extension point for downstream crates — implement the trait and
|
||||
/// construct a `Connection` from it without editing core. See ADR-070.
|
||||
pub fn from_source(source: impl BidiStreamSource, alpn: Vec<u8>) -> Self {
|
||||
Self {
|
||||
source: Box::new(source),
|
||||
alpn,
|
||||
identity: OnceLock::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Yield the next bidirectional stream this connection provides.
|
||||
///
|
||||
/// # Transport semantics
|
||||
///
|
||||
/// - **QUIC (quinn/iroh)**: returns a new bidi stream on each call.
|
||||
/// `ConnectionClosed` when the underlying connection closes.
|
||||
/// - **TCP+TLS / single-stream**: yields the underlying stream on the
|
||||
/// first call, then `ConnectionClosed` on all subsequent calls.
|
||||
/// A single transport stream cannot open new application streams.
|
||||
///
|
||||
/// Handlers that loop `accept_bi` (e.g. `TtyAdapter`) get one session
|
||||
/// per single-stream connection; handlers that call once (e.g.
|
||||
/// `HttpAdapter`) get the stream directly. Both are correct. The
|
||||
/// return type is `BiStream` (ADR-092); handlers that want the split
|
||||
/// halves call `tokio::io::split` on the `BiStream`.
|
||||
pub async fn accept_bi(&self) -> Result<BiStream, StreamError> {
|
||||
self.source.accept_bi().await
|
||||
}
|
||||
|
||||
pub async fn open_bi(&self) -> Result<BiStream, StreamError> {
|
||||
self.source.open_bi().await
|
||||
}
|
||||
|
||||
pub fn remote_alpn(&self) -> &[u8] {
|
||||
&self.alpn
|
||||
}
|
||||
|
||||
pub fn remote_addr(&self) -> Option<SocketAddr> {
|
||||
self.source.remote_addr()
|
||||
}
|
||||
|
||||
pub fn close(&self, code: u32, reason: &str) {
|
||||
self.source.close(code, reason)
|
||||
}
|
||||
|
||||
pub fn set_identity(&self, identity: Identity) -> Result<(), IdentityAlreadySet> {
|
||||
self.identity
|
||||
.set(identity)
|
||||
.map_err(|_| IdentityAlreadySet::AlreadySet)
|
||||
}
|
||||
|
||||
pub fn identity(&self) -> Option<&Identity> {
|
||||
self.identity.get()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod from_source_tests {
|
||||
use super::*;
|
||||
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
|
||||
use std::sync::Arc;
|
||||
|
||||
/// A minimal custom `BidiStreamSource` impl to prove `from_source`
|
||||
/// delegates to a caller-supplied impl. Not a built-in — the whole
|
||||
/// point of `from_source` is that a non-core type can drive `Connection`.
|
||||
struct RecordingSource {
|
||||
stream: Mutex<Option<BiStream>>,
|
||||
addr: Option<SocketAddr>,
|
||||
closed: Arc<Mutex<Option<(u32, String)>>>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl BidiStreamSource for RecordingSource {
|
||||
async fn accept_bi(&self) -> Result<BiStream, StreamError> {
|
||||
match self.stream.lock().expect("mock mutex poisoned").take() {
|
||||
Some(stream) => Ok(stream),
|
||||
None => Err(StreamError::ConnectionClosed),
|
||||
}
|
||||
}
|
||||
|
||||
async fn open_bi(&self) -> Result<BiStream, StreamError> {
|
||||
Err(StreamError::StreamClosed)
|
||||
}
|
||||
|
||||
fn remote_addr(&self) -> Option<SocketAddr> {
|
||||
self.addr
|
||||
}
|
||||
|
||||
fn close(&self, code: u32, reason: &str) {
|
||||
let _ = self
|
||||
.closed
|
||||
.lock()
|
||||
.expect("mock closed mutex poisoned")
|
||||
.replace((code, reason.to_string()));
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn from_source_delegates_to_custom_impl() {
|
||||
use tokio::io::AsyncReadExt;
|
||||
use tokio::io::AsyncWriteExt;
|
||||
|
||||
let (a, b) = tokio::io::duplex(64);
|
||||
let (mut recv_b, mut send_b) = tokio::io::split(b);
|
||||
let addr = Some(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 7777));
|
||||
let recorded = Arc::new(Mutex::new(None));
|
||||
let conn = Connection::from_source(
|
||||
RecordingSource {
|
||||
stream: Mutex::new(Some(BiStream::from_bidi(a))),
|
||||
addr,
|
||||
closed: Arc::clone(&recorded),
|
||||
},
|
||||
b"alknet/test".to_vec(),
|
||||
);
|
||||
|
||||
assert_eq!(conn.remote_alpn(), b"alknet/test");
|
||||
assert_eq!(conn.remote_addr(), addr);
|
||||
|
||||
let mut stream = conn.accept_bi().await.expect("first accept_bi yields");
|
||||
|
||||
stream.write_all(b"hello").await.expect("write round-trips");
|
||||
let mut buf = [0u8; 5];
|
||||
recv_b.read_exact(&mut buf).await.expect("driver reads");
|
||||
assert_eq!(&buf, b"hello");
|
||||
|
||||
send_b
|
||||
.write_all(b"world")
|
||||
.await
|
||||
.expect("driver writes back");
|
||||
let mut buf = [0u8; 5];
|
||||
stream.read_exact(&mut buf).await.expect("read round-trips");
|
||||
assert_eq!(&buf, b"world");
|
||||
|
||||
match conn.accept_bi().await {
|
||||
Err(StreamError::ConnectionClosed) => {}
|
||||
Err(e) => panic!("expected ConnectionClosed on second accept_bi, got {e}"),
|
||||
Ok(_) => panic!("expected ConnectionClosed on second accept_bi, got a stream"),
|
||||
}
|
||||
|
||||
match conn.open_bi().await {
|
||||
Err(StreamError::StreamClosed) => {}
|
||||
Err(e) => panic!("expected StreamClosed from open_bi, got {e}"),
|
||||
Ok(_) => panic!("expected StreamClosed from open_bi, got a stream"),
|
||||
}
|
||||
|
||||
conn.close(42, "shutting down");
|
||||
assert_eq!(
|
||||
recorded.lock().expect("recorded mutex poisoned").take(),
|
||||
Some((42, "shutting down".to_string()))
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
|
||||
use std::pin::Pin;
|
||||
use std::task::{Context, Poll};
|
||||
|
||||
/// A test-only `AsyncRead + AsyncWrite` pair equivalent to
|
||||
/// `tokio::io::sink()` + `tokio::io::empty()`: reads yield EOF
|
||||
/// immediately (zero bytes), writes discard. Exists because
|
||||
/// `Connection::from_bidi` requires a single value that implements
|
||||
/// both traits (ADR-092). Used only to construct a `Connection` for
|
||||
/// tests that exercise `Connection`-level state (alpn, addr, identity)
|
||||
/// without ever reading or writing the stream.
|
||||
pub(crate) struct SinkEmpty;
|
||||
|
||||
impl AsyncRead for SinkEmpty {
|
||||
fn poll_read(
|
||||
self: Pin<&mut Self>,
|
||||
_cx: &mut Context<'_>,
|
||||
_buf: &mut tokio::io::ReadBuf<'_>,
|
||||
) -> Poll<io::Result<()>> {
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
}
|
||||
|
||||
impl AsyncWrite for SinkEmpty {
|
||||
fn poll_write(
|
||||
self: Pin<&mut Self>,
|
||||
_cx: &mut Context<'_>,
|
||||
buf: &[u8],
|
||||
) -> Poll<io::Result<usize>> {
|
||||
Poll::Ready(Ok(buf.len()))
|
||||
}
|
||||
|
||||
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
|
||||
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
}
|
||||
|
||||
fn test_connection() -> Connection {
|
||||
Connection::from_bidi(
|
||||
SinkEmpty,
|
||||
b"alknet/test".to_vec(),
|
||||
Some(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 1234)),
|
||||
)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn capabilities_new_is_empty() {
|
||||
let caps = Capabilities::new();
|
||||
assert!(caps.get("google").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn capabilities_with_api_key_then_get() {
|
||||
let caps = Capabilities::new().with_api_key("google", "sekrit".to_string());
|
||||
let secret = caps.get("google").expect("api key present");
|
||||
assert_eq!(secret.expose_secret(), "sekrit");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn capabilities_with_http_token_then_get() {
|
||||
let caps = Capabilities::new().with_http_token("github", "tok".to_string());
|
||||
let secret = caps.get("github").expect("http token present");
|
||||
assert_eq!(secret.expose_secret(), "tok");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn capabilities_clone_preserves_entries() {
|
||||
let caps = Capabilities::new().with_api_key("google", "k".to_string());
|
||||
let cloned = caps.clone();
|
||||
assert_eq!(
|
||||
cloned.get("google").map(|s| s.expose_secret().clone()),
|
||||
Some("k".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
caps.get("google").map(|s| s.expose_secret().clone()),
|
||||
Some("k".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn capabilities_zeroize_on_drop_clears_secret() {
|
||||
let mut secret = Secret::new("sensitive".to_string());
|
||||
secret.zeroize();
|
||||
assert_eq!(secret.expose_secret(), "");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn capabilities_does_not_derive_serialize() {
|
||||
fn assert_not_serialize<T>() {}
|
||||
assert_not_serialize::<Capabilities>();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn capabilities_debug_redacts_entries() {
|
||||
let caps = Capabilities::new().with_api_key("google", "sekrit".to_string());
|
||||
let s = format!("{:?}", caps);
|
||||
assert!(s.contains("redacted"));
|
||||
assert!(!s.contains("sekrit"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn secret_debug_redacts() {
|
||||
let secret = Secret::new("hidden".to_string());
|
||||
assert_eq!(format!("{:?}", secret), "[REDACTED]");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn set_identity_once_succeeds_twice_errors() {
|
||||
let conn = test_connection();
|
||||
let id = Identity {
|
||||
id: "alk_test".to_string(),
|
||||
scopes: vec!["relay:connect".to_string()],
|
||||
resources: HashMap::new(),
|
||||
};
|
||||
assert!(conn.set_identity(id.clone()).is_ok());
|
||||
assert!(matches!(
|
||||
conn.set_identity(id),
|
||||
Err(IdentityAlreadySet::AlreadySet)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn identity_get_returns_set_value() {
|
||||
let conn = test_connection();
|
||||
assert!(conn.identity().is_none());
|
||||
let id = Identity {
|
||||
id: "alk_test".to_string(),
|
||||
scopes: vec![],
|
||||
resources: HashMap::new(),
|
||||
};
|
||||
conn.set_identity(id.clone()).unwrap();
|
||||
assert_eq!(conn.identity(), Some(&id));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn connection_remote_alpn_and_addr_from_bidi() {
|
||||
let conn = test_connection();
|
||||
assert_eq!(conn.remote_alpn(), b"alknet/test");
|
||||
assert_eq!(
|
||||
conn.remote_addr(),
|
||||
Some(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 1234))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_error_maps_to_handler_error() {
|
||||
assert!(matches!(
|
||||
HandlerError::from(StreamError::ConnectionClosed),
|
||||
HandlerError::ConnectionClosed
|
||||
));
|
||||
match HandlerError::from(StreamError::StreamClosed) {
|
||||
HandlerError::StreamError(e) => assert_eq!(e.kind(), io::ErrorKind::ConnectionReset),
|
||||
other => panic!("expected StreamError, got {other:?}"),
|
||||
}
|
||||
match HandlerError::from(StreamError::Timeout) {
|
||||
HandlerError::StreamError(e) => assert_eq!(e.kind(), io::ErrorKind::TimedOut),
|
||||
other => panic!("expected StreamError, got {other:?}"),
|
||||
}
|
||||
match HandlerError::from(StreamError::Internal(io::Error::other("x"))) {
|
||||
HandlerError::StreamError(e) => assert_eq!(e.kind(), io::ErrorKind::Other),
|
||||
other => panic!("expected StreamError, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn handler_error_auth_required_constructible() {
|
||||
let e = HandlerError::AuthRequired;
|
||||
assert_eq!(format!("{e}"), "authentication required");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn handler_error_debug_covers_all_variants() {
|
||||
assert_eq!(
|
||||
format!("{:?}", HandlerError::ConnectionClosed),
|
||||
"HandlerError::ConnectionClosed"
|
||||
);
|
||||
let io_err = io::Error::new(io::ErrorKind::BrokenPipe, "boom");
|
||||
let dbg = format!("{:?}", HandlerError::StreamError(io_err));
|
||||
assert!(dbg.contains("HandlerError::StreamError"));
|
||||
assert_eq!(
|
||||
format!("{:?}", HandlerError::AuthRequired),
|
||||
"HandlerError::AuthRequired"
|
||||
);
|
||||
let inner: Box<dyn std::error::Error + Send + Sync> = "oops".into();
|
||||
let dbg = format!("{:?}", HandlerError::Internal(inner));
|
||||
assert!(dbg.contains("HandlerError::Internal"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn handler_error_display_covers_all_variants() {
|
||||
assert_eq!(
|
||||
format!("{}", HandlerError::ConnectionClosed),
|
||||
"connection closed"
|
||||
);
|
||||
let io_err = io::Error::new(io::ErrorKind::BrokenPipe, "boom");
|
||||
let s = format!("{}", HandlerError::StreamError(io_err));
|
||||
assert!(s.starts_with("stream error: "));
|
||||
assert_eq!(
|
||||
format!("{}", HandlerError::AuthRequired),
|
||||
"authentication required"
|
||||
);
|
||||
let inner: Box<dyn std::error::Error + Send + Sync> = "oops".into();
|
||||
assert_eq!(
|
||||
format!("{}", HandlerError::Internal(inner)),
|
||||
"internal handler error: oops"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn handler_error_source_covers_all_variants() {
|
||||
use std::error::Error;
|
||||
assert!(HandlerError::ConnectionClosed.source().is_none());
|
||||
assert!(HandlerError::AuthRequired.source().is_none());
|
||||
let stream_err =
|
||||
HandlerError::StreamError(io::Error::new(io::ErrorKind::BrokenPipe, "boom"));
|
||||
assert!(
|
||||
stream_err.source().is_some(),
|
||||
"StreamError must expose its io::Error as source"
|
||||
);
|
||||
let internal_inner: Box<dyn std::error::Error + Send + Sync> = "boom".into();
|
||||
let internal_err = HandlerError::Internal(internal_inner);
|
||||
assert!(
|
||||
internal_err.source().is_some(),
|
||||
"Internal must expose its inner error as source"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_error_debug_covers_all_variants() {
|
||||
assert_eq!(
|
||||
format!("{:?}", StreamError::ConnectionClosed),
|
||||
"StreamError::ConnectionClosed"
|
||||
);
|
||||
assert_eq!(
|
||||
format!("{:?}", StreamError::StreamClosed),
|
||||
"StreamError::StreamClosed"
|
||||
);
|
||||
assert_eq!(
|
||||
format!("{:?}", StreamError::Timeout),
|
||||
"StreamError::Timeout"
|
||||
);
|
||||
let dbg = format!("{:?}", StreamError::Internal(io::Error::other("x")));
|
||||
assert!(dbg.contains("StreamError::Internal"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_error_display_covers_all_variants() {
|
||||
assert_eq!(
|
||||
format!("{}", StreamError::ConnectionClosed),
|
||||
"connection closed"
|
||||
);
|
||||
assert_eq!(format!("{}", StreamError::StreamClosed), "stream closed");
|
||||
assert_eq!(format!("{}", StreamError::Timeout), "stream timed out");
|
||||
assert_eq!(
|
||||
format!("{}", StreamError::Internal(io::Error::other("boom"))),
|
||||
"stream error: boom"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_error_source_covers_all_variants() {
|
||||
use std::error::Error;
|
||||
assert!(StreamError::ConnectionClosed.source().is_none());
|
||||
assert!(StreamError::StreamClosed.source().is_none());
|
||||
assert!(StreamError::Timeout.source().is_none());
|
||||
let internal = StreamError::Internal(io::Error::other("x"));
|
||||
assert!(
|
||||
internal.source().is_some(),
|
||||
"Internal must expose its io::Error as source"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn capabilities_default_is_empty() {
|
||||
let caps = Capabilities::default();
|
||||
assert!(caps.get("anything").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn capabilities_zeroize_clears_entries() {
|
||||
let mut caps = Capabilities::new()
|
||||
.with_api_key("svc-a", "k1".to_string())
|
||||
.with_http_token("svc-b", "t1".to_string());
|
||||
assert!(caps.get("svc-a").is_some());
|
||||
assert!(caps.get("svc-b").is_some());
|
||||
caps.zeroize();
|
||||
assert!(caps.get("svc-a").is_none());
|
||||
assert!(caps.get("svc-b").is_none());
|
||||
}
|
||||
}
|
||||
+25
@@ -0,0 +1,25 @@
|
||||
//! alkcall: Call + channels RPC — operations, streaming, service discovery,
|
||||
//! and N-channel multiplexing over one transport stream.
|
||||
//!
|
||||
//! This crate unifies `alknet-call` (structured JSON RPC: operations,
|
||||
//! streaming subscriptions, service discovery) and `alknet-channels`
|
||||
//! (multiplexing proxy: N logical channels over one transport stream,
|
||||
//! channel 0 pre-negotiated as `alknet/call`).
|
||||
//!
|
||||
//! ## Architecture
|
||||
//!
|
||||
//! - **Vendored core types** ([`core`]): `Connection`, `ProtocolHandler`,
|
||||
//! `BiStream`, `BidiStreamSource`, `AuthContext`, `IdentityProvider`,
|
||||
//! `Capabilities`, `OwnershipProvider` — the home for the former
|
||||
//! `alknet-core` surface.
|
||||
//! - **Registry** ([`registry`]): operation specs, context, dispatch, and
|
||||
//! the operation registry — the call half's dispatch core.
|
||||
//! - **Protocol** ([`protocol`]): wire format, streams, adapter, dispatch
|
||||
//! loop, pending requests, abort cascade — the call half's wire layer.
|
||||
//! - **Client** ([`client`]): `CallClient`, `from_call`, `OperationAdapter`
|
||||
//! — the call half's outbound surface.
|
||||
|
||||
pub mod client;
|
||||
pub mod core;
|
||||
pub mod protocol;
|
||||
pub mod registry;
|
||||
@@ -0,0 +1,393 @@
|
||||
//! Abort cascade logic for nested calls (ADR-016).
|
||||
//!
|
||||
//! When `call.aborted` arrives for a parent request, the protocol cascades
|
||||
//! the abort to all non-terminal descendants in the call tree. The default
|
||||
//! policy is `abort-dependents`; `continue-running` is an opt-in for
|
||||
//! long-running work that should survive a parent's abort.
|
||||
//!
|
||||
//! The call tree is indexed by `parent_request_id` in the
|
||||
//! `PendingRequestMap`. The root request has `parent_request_id: None`;
|
||||
//! each composed call has `parent_request_id: Some(parent.request_id)`.
|
||||
//! Composed child request IDs are internal — they appear in the map for
|
||||
//! abort-cascade indexing but are not sent as `call.requested` to any
|
||||
//! peer. The client only sees `call.aborted` for the root ID it sent; the
|
||||
//! server cascades internally to descendants.
|
||||
|
||||
use super::pending::PendingRequestMap;
|
||||
use crate::registry::context::AbortPolicy;
|
||||
|
||||
pub struct AbortCascade<'a> {
|
||||
pending: &'a mut PendingRequestMap,
|
||||
}
|
||||
|
||||
impl<'a> AbortCascade<'a> {
|
||||
pub fn new(pending: &'a mut PendingRequestMap) -> Self {
|
||||
Self { pending }
|
||||
}
|
||||
|
||||
/// Cascade an abort from the given request ID to all non-terminal
|
||||
/// descendants in the call tree. Returns the list of descendant
|
||||
/// request IDs that were aborted (for logging/auditing), sorted for
|
||||
/// determinism. The root request itself is not touched by this
|
||||
/// method — the caller is responsible for aborting the root (the
|
||||
/// trigger of the cascade).
|
||||
///
|
||||
/// Under `AbortDependents` (default): all descendants are aborted,
|
||||
/// regardless of whether they have started.
|
||||
///
|
||||
/// Under `ContinueRunning`: only descendants that have not started
|
||||
/// are aborted; started descendants continue to completion. No new
|
||||
/// descendants start (the parent is gone). This is the conservative
|
||||
/// approximation noted in ADR-016: a descendant is "started" if
|
||||
/// `PendingEntry::started` is true (the handler has begun
|
||||
/// executing). A `call.aborted` for an unknown request ID is
|
||||
/// silently discarded — `cascade_abort` on an unknown root returns
|
||||
/// an empty list and removes nothing.
|
||||
pub fn cascade_abort(&mut self, root_request_id: &str, policy: AbortPolicy) -> Vec<String> {
|
||||
if !self.pending.contains(root_request_id) {
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
let descendants = self.find_descendants(root_request_id);
|
||||
|
||||
let mut aborted = Vec::new();
|
||||
match policy {
|
||||
AbortPolicy::AbortDependents => {
|
||||
for id in &descendants {
|
||||
if self.pending.handle_aborted(id) {
|
||||
aborted.push(id.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
AbortPolicy::ContinueRunning => {
|
||||
for id in &descendants {
|
||||
let started = self.pending.is_started(id).unwrap_or(false);
|
||||
if !started && self.pending.handle_aborted(id) {
|
||||
aborted.push(id.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
aborted.sort();
|
||||
aborted
|
||||
}
|
||||
|
||||
/// Find all descendants of a request ID in the call tree by walking
|
||||
/// the `parent_request_id` index. Returns descendants in
|
||||
/// breadth-first order with each level's children sorted for
|
||||
/// determinism. The root itself is not included in the result.
|
||||
fn find_descendants(&self, parent_id: &str) -> Vec<String> {
|
||||
let mut descendants = Vec::new();
|
||||
let mut frontier: Vec<String> = vec![parent_id.to_string()];
|
||||
|
||||
while let Some(current) = frontier.pop() {
|
||||
let mut children: Vec<String> = self
|
||||
.pending
|
||||
.request_ids()
|
||||
.into_iter()
|
||||
.filter(|id| {
|
||||
self.pending
|
||||
.parent_of(id)
|
||||
.flatten()
|
||||
.is_some_and(|p| p == current)
|
||||
})
|
||||
.collect();
|
||||
children.sort();
|
||||
for child in children {
|
||||
descendants.push(child.clone());
|
||||
frontier.push(child);
|
||||
}
|
||||
}
|
||||
|
||||
descendants
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::protocol::wire::CallError;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
fn register_call(map: &mut PendingRequestMap, id: &str, parent: Option<&str>) {
|
||||
map.register_call(
|
||||
id.to_string(),
|
||||
Instant::now() + Duration::from_secs(30),
|
||||
parent.map(|p| p.to_string()),
|
||||
);
|
||||
}
|
||||
|
||||
fn register_subscribe(map: &mut PendingRequestMap, id: &str, parent: Option<&str>) {
|
||||
map.register_subscribe(id.to_string(), None, parent.map(|p| p.to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cascade_abort_unknown_root_returns_empty_and_is_noop() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
register_call(&mut map, "r1", None);
|
||||
let mut cascade = AbortCascade::new(&mut map);
|
||||
let aborted = cascade.cascade_abort("does-not-exist", AbortPolicy::AbortDependents);
|
||||
assert!(aborted.is_empty());
|
||||
assert!(cascade.pending.contains("r1"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cascade_abort_abort_dependents_aborts_all_descendants() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
register_call(&mut map, "r1", None);
|
||||
register_call(&mut map, "r1-a", Some("r1"));
|
||||
register_call(&mut map, "r1-b", Some("r1"));
|
||||
register_call(&mut map, "r1-a-1", Some("r1-a"));
|
||||
register_call(&mut map, "r1-a-2", Some("r1-a"));
|
||||
register_call(&mut map, "r1-b-1", Some("r1-b"));
|
||||
|
||||
let mut cascade = AbortCascade::new(&mut map);
|
||||
let aborted = cascade.cascade_abort("r1", AbortPolicy::AbortDependents);
|
||||
|
||||
assert_eq!(
|
||||
aborted,
|
||||
vec![
|
||||
"r1-a".to_string(),
|
||||
"r1-a-1".to_string(),
|
||||
"r1-a-2".to_string(),
|
||||
"r1-b".to_string(),
|
||||
"r1-b-1".to_string(),
|
||||
]
|
||||
);
|
||||
assert!(cascade.pending.contains("r1"));
|
||||
assert!(!cascade.pending.contains("r1-a"));
|
||||
assert!(!cascade.pending.contains("r1-b"));
|
||||
assert!(!cascade.pending.contains("r1-a-1"));
|
||||
assert!(!cascade.pending.contains("r1-a-2"));
|
||||
assert!(!cascade.pending.contains("r1-b-1"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cascade_abort_continue_running_aborts_only_unstarted_descendants() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
register_call(&mut map, "r1", None);
|
||||
register_call(&mut map, "r1-a", Some("r1"));
|
||||
register_call(&mut map, "r1-b", Some("r1"));
|
||||
register_call(&mut map, "r1-a-1", Some("r1-a"));
|
||||
|
||||
map.mark_started("r1-a");
|
||||
// r1-b and r1-a-1 are unstarted
|
||||
|
||||
let mut cascade = AbortCascade::new(&mut map);
|
||||
let aborted = cascade.cascade_abort("r1", AbortPolicy::ContinueRunning);
|
||||
|
||||
assert_eq!(aborted, vec!["r1-a-1".to_string(), "r1-b".to_string()]);
|
||||
assert!(cascade.pending.contains("r1"));
|
||||
assert!(cascade.pending.contains("r1-a"));
|
||||
assert!(!cascade.pending.contains("r1-b"));
|
||||
assert!(!cascade.pending.contains("r1-a-1"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cascade_abort_continue_running_aborts_all_when_none_started() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
register_call(&mut map, "r1", None);
|
||||
register_call(&mut map, "r1-a", Some("r1"));
|
||||
register_call(&mut map, "r1-b", Some("r1"));
|
||||
|
||||
let mut cascade = AbortCascade::new(&mut map);
|
||||
let aborted = cascade.cascade_abort("r1", AbortPolicy::ContinueRunning);
|
||||
|
||||
assert_eq!(aborted, vec!["r1-a".to_string(), "r1-b".to_string()]);
|
||||
assert!(!cascade.pending.contains("r1-a"));
|
||||
assert!(!cascade.pending.contains("r1-b"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cascade_abort_depth_three_aborts_all_descendants() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
register_call(&mut map, "root", None);
|
||||
register_call(&mut map, "root-a", Some("root"));
|
||||
register_call(&mut map, "root-b", Some("root"));
|
||||
register_call(&mut map, "root-a-1", Some("root-a"));
|
||||
register_call(&mut map, "root-a-2", Some("root-a"));
|
||||
register_call(&mut map, "root-a-1-x", Some("root-a-1"));
|
||||
register_call(&mut map, "root-a-1-y", Some("root-a-1"));
|
||||
register_call(&mut map, "root-b-1", Some("root-b"));
|
||||
|
||||
let mut cascade = AbortCascade::new(&mut map);
|
||||
let aborted = cascade.cascade_abort("root", AbortPolicy::AbortDependents);
|
||||
|
||||
assert_eq!(
|
||||
aborted,
|
||||
vec![
|
||||
"root-a".to_string(),
|
||||
"root-a-1".to_string(),
|
||||
"root-a-1-x".to_string(),
|
||||
"root-a-1-y".to_string(),
|
||||
"root-a-2".to_string(),
|
||||
"root-b".to_string(),
|
||||
"root-b-1".to_string(),
|
||||
]
|
||||
);
|
||||
assert!(cascade.pending.contains("root"));
|
||||
assert_eq!(cascade.pending.len(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cascade_abort_root_with_no_descendants_returns_empty() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
register_call(&mut map, "lonely", None);
|
||||
|
||||
let mut cascade = AbortCascade::new(&mut map);
|
||||
let aborted = cascade.cascade_abort("lonely", AbortPolicy::AbortDependents);
|
||||
assert!(aborted.is_empty());
|
||||
assert!(cascade.pending.contains("lonely"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cascade_abort_only_aborts_descendants_not_siblings() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
register_call(&mut map, "r1", None);
|
||||
register_call(&mut map, "r2", None);
|
||||
register_call(&mut map, "r1-a", Some("r1"));
|
||||
register_call(&mut map, "r2-a", Some("r2"));
|
||||
|
||||
let mut cascade = AbortCascade::new(&mut map);
|
||||
let aborted = cascade.cascade_abort("r1", AbortPolicy::AbortDependents);
|
||||
|
||||
assert_eq!(aborted, vec!["r1-a".to_string()]);
|
||||
assert!(cascade.pending.contains("r1"));
|
||||
assert!(cascade.pending.contains("r2"));
|
||||
assert!(cascade.pending.contains("r2-a"));
|
||||
assert!(!cascade.pending.contains("r1-a"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cascade_abort_handles_mixed_call_and_subscribe_entries() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
register_call(&mut map, "r1", None);
|
||||
register_subscribe(&mut map, "r1-sub", Some("r1"));
|
||||
register_call(&mut map, "r1-sub-child", Some("r1-sub"));
|
||||
|
||||
let mut cascade = AbortCascade::new(&mut map);
|
||||
let aborted = cascade.cascade_abort("r1", AbortPolicy::AbortDependents);
|
||||
|
||||
assert_eq!(
|
||||
aborted,
|
||||
vec!["r1-sub".to_string(), "r1-sub-child".to_string(),]
|
||||
);
|
||||
assert!(cascade.pending.contains("r1"));
|
||||
assert_eq!(cascade.pending.len(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cascade_abort_continue_running_with_started_descendant_keeps_its_unstarted_children() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
register_call(&mut map, "r1", None);
|
||||
register_call(&mut map, "r1-a", Some("r1"));
|
||||
register_call(&mut map, "r1-a-1", Some("r1-a"));
|
||||
|
||||
map.mark_started("r1-a");
|
||||
// r1-a is started and continues; r1-a-1 is unstarted.
|
||||
// Under ContinueRunning, r1-a-1 is aborted (conservative: still pending).
|
||||
|
||||
let mut cascade = AbortCascade::new(&mut map);
|
||||
let aborted = cascade.cascade_abort("r1", AbortPolicy::ContinueRunning);
|
||||
|
||||
assert_eq!(aborted, vec!["r1-a-1".to_string()]);
|
||||
assert!(cascade.pending.contains("r1-a"));
|
||||
assert!(!cascade.pending.contains("r1-a-1"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cascade_abort_abort_dependents_aborts_started_descendants_too() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
register_call(&mut map, "r1", None);
|
||||
register_call(&mut map, "r1-a", Some("r1"));
|
||||
register_call(&mut map, "r1-b", Some("r1"));
|
||||
|
||||
map.mark_started("r1-a");
|
||||
map.mark_started("r1-b");
|
||||
|
||||
let mut cascade = AbortCascade::new(&mut map);
|
||||
let aborted = cascade.cascade_abort("r1", AbortPolicy::AbortDependents);
|
||||
|
||||
assert_eq!(aborted, vec!["r1-a".to_string(), "r1-b".to_string()]);
|
||||
assert!(!cascade.pending.contains("r1-a"));
|
||||
assert!(!cascade.pending.contains("r1-b"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn find_descendants_does_not_include_root() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
register_call(&mut map, "r1", None);
|
||||
register_call(&mut map, "r1-a", Some("r1"));
|
||||
|
||||
let cascade = AbortCascade::new(&mut map);
|
||||
let descendants = cascade.find_descendants("r1");
|
||||
assert_eq!(descendants, vec!["r1-a".to_string()]);
|
||||
assert!(!descendants.contains(&"r1".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cascade_abort_default_policy_is_abort_dependents() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
register_call(&mut map, "r1", None);
|
||||
register_call(&mut map, "r1-a", Some("r1"));
|
||||
map.mark_started("r1-a");
|
||||
|
||||
let mut cascade = AbortCascade::new(&mut map);
|
||||
let aborted_default = cascade.cascade_abort("r1", AbortPolicy::default());
|
||||
assert_eq!(aborted_default, vec!["r1-a".to_string()]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cascade_abort_does_not_remove_root() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
register_call(&mut map, "r1", None);
|
||||
register_call(&mut map, "r1-a", Some("r1"));
|
||||
|
||||
let mut cascade = AbortCascade::new(&mut map);
|
||||
let _ = cascade.cascade_abort("r1", AbortPolicy::AbortDependents);
|
||||
assert!(cascade.pending.contains("r1"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cascade_abort_returns_sorted_descendants_for_determinism() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
register_call(&mut map, "r1", None);
|
||||
register_call(&mut map, "r1-z", Some("r1"));
|
||||
register_call(&mut map, "r1-a", Some("r1"));
|
||||
register_call(&mut map, "r1-m", Some("r1"));
|
||||
|
||||
let mut cascade = AbortCascade::new(&mut map);
|
||||
let aborted = cascade.cascade_abort("r1", AbortPolicy::AbortDependents);
|
||||
assert_eq!(
|
||||
aborted,
|
||||
vec!["r1-a".to_string(), "r1-m".to_string(), "r1-z".to_string(),]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unknown_request_id_silently_discarded_no_panic() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
let mut cascade = AbortCascade::new(&mut map);
|
||||
let aborted = cascade.cascade_abort("totally-unknown", AbortPolicy::AbortDependents);
|
||||
assert!(aborted.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cascade_abort_continue_running_started_descendant_survives() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
register_call(&mut map, "r1", None);
|
||||
register_call(&mut map, "r1-a", Some("r1"));
|
||||
map.mark_started("r1-a");
|
||||
|
||||
let mut cascade = AbortCascade::new(&mut map);
|
||||
let aborted = cascade.cascade_abort("r1", AbortPolicy::ContinueRunning);
|
||||
assert!(aborted.is_empty());
|
||||
assert!(cascade.pending.contains("r1-a"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cascade_abort_handles_call_error_unused() {
|
||||
let _ = CallError::internal("unused");
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,18 @@
|
||||
//! Call protocol: wire format, streams, and the call adapter.
|
||||
//!
|
||||
//! Implements `ProtocolHandler` for ALPN `alknet/call` on top of the
|
||||
//! operation registry. See `docs/architecture/crates/call/call-protocol.md`
|
||||
//! for the full specification.
|
||||
|
||||
pub mod abort;
|
||||
pub mod adapter;
|
||||
pub mod connection;
|
||||
pub mod dispatch;
|
||||
pub mod pending;
|
||||
pub mod wire;
|
||||
|
||||
#[cfg(test)]
|
||||
mod test_support;
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) use test_support::sink_empty_connection;
|
||||
@@ -0,0 +1,584 @@
|
||||
use std::collections::HashMap;
|
||||
use std::time::Instant;
|
||||
|
||||
use serde_json::Value;
|
||||
use tokio::sync::{mpsc, oneshot};
|
||||
|
||||
use crate::protocol::wire::CallError;
|
||||
|
||||
const SUBSCRIBE_CHANNEL_CAPACITY: usize = 32;
|
||||
|
||||
pub struct PendingRequestMap {
|
||||
pending: HashMap<String, PendingEntry>,
|
||||
}
|
||||
|
||||
pub(crate) enum PendingEntry {
|
||||
Call {
|
||||
tx: oneshot::Sender<Result<Value, CallError>>,
|
||||
timeout: Instant,
|
||||
parent_request_id: Option<String>,
|
||||
started: bool,
|
||||
},
|
||||
Subscribe {
|
||||
tx: mpsc::Sender<Result<Value, CallError>>,
|
||||
timeout: Option<Instant>,
|
||||
parent_request_id: Option<String>,
|
||||
started: bool,
|
||||
},
|
||||
}
|
||||
|
||||
impl PendingEntry {
|
||||
pub(crate) fn parent_request_id(&self) -> Option<&str> {
|
||||
match self {
|
||||
PendingEntry::Call {
|
||||
parent_request_id, ..
|
||||
} => parent_request_id.as_deref(),
|
||||
PendingEntry::Subscribe {
|
||||
parent_request_id, ..
|
||||
} => parent_request_id.as_deref(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn started(&self) -> bool {
|
||||
match self {
|
||||
PendingEntry::Call { started, .. } => *started,
|
||||
PendingEntry::Subscribe { started, .. } => *started,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl PendingRequestMap {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
pending: HashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn register_call(
|
||||
&mut self,
|
||||
request_id: String,
|
||||
timeout: Instant,
|
||||
parent_request_id: Option<String>,
|
||||
) -> oneshot::Receiver<Result<Value, CallError>> {
|
||||
let (tx, rx) = oneshot::channel();
|
||||
self.pending.insert(
|
||||
request_id,
|
||||
PendingEntry::Call {
|
||||
tx,
|
||||
timeout,
|
||||
parent_request_id,
|
||||
started: false,
|
||||
},
|
||||
);
|
||||
rx
|
||||
}
|
||||
|
||||
pub fn register_subscribe(
|
||||
&mut self,
|
||||
request_id: String,
|
||||
timeout: Option<Instant>,
|
||||
parent_request_id: Option<String>,
|
||||
) -> mpsc::Receiver<Result<Value, CallError>> {
|
||||
let (tx, rx) = mpsc::channel(SUBSCRIBE_CHANNEL_CAPACITY);
|
||||
self.pending.insert(
|
||||
request_id,
|
||||
PendingEntry::Subscribe {
|
||||
tx,
|
||||
timeout,
|
||||
parent_request_id,
|
||||
started: false,
|
||||
},
|
||||
);
|
||||
rx
|
||||
}
|
||||
|
||||
pub fn mark_started(&mut self, request_id: &str) -> bool {
|
||||
let Some(entry) = self.pending.get_mut(request_id) else {
|
||||
return false;
|
||||
};
|
||||
match entry {
|
||||
PendingEntry::Call { started, .. } => *started = true,
|
||||
PendingEntry::Subscribe { started, .. } => *started = true,
|
||||
}
|
||||
true
|
||||
}
|
||||
|
||||
pub fn handle_responded(&mut self, request_id: &str, output: Value) -> bool {
|
||||
let Some(entry) = self.pending.remove(request_id) else {
|
||||
return false;
|
||||
};
|
||||
match entry {
|
||||
PendingEntry::Call { tx, .. } => {
|
||||
let _ = tx.send(Ok(output));
|
||||
true
|
||||
}
|
||||
PendingEntry::Subscribe {
|
||||
tx,
|
||||
timeout,
|
||||
parent_request_id,
|
||||
started,
|
||||
} => {
|
||||
let send_result = tx.try_send(Ok(output));
|
||||
match send_result {
|
||||
Ok(()) => {
|
||||
self.pending.insert(
|
||||
request_id.to_string(),
|
||||
PendingEntry::Subscribe {
|
||||
tx,
|
||||
timeout,
|
||||
parent_request_id,
|
||||
started,
|
||||
},
|
||||
);
|
||||
true
|
||||
}
|
||||
Err(mpsc::error::TrySendError::Full(_)) => {
|
||||
tracing::warn!(
|
||||
request_id,
|
||||
"subscribe channel full; dropping entry and closing subscription"
|
||||
);
|
||||
true
|
||||
}
|
||||
Err(mpsc::error::TrySendError::Closed(_)) => true,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn handle_completed(&mut self, request_id: &str) -> bool {
|
||||
self.pending.remove(request_id).is_some()
|
||||
}
|
||||
|
||||
pub fn handle_aborted(&mut self, request_id: &str) -> bool {
|
||||
self.pending.remove(request_id).is_some()
|
||||
}
|
||||
|
||||
pub fn handle_error(&mut self, request_id: &str, error: CallError) -> bool {
|
||||
let Some(entry) = self.pending.remove(request_id) else {
|
||||
return false;
|
||||
};
|
||||
match entry {
|
||||
PendingEntry::Call { tx, .. } => {
|
||||
let _ = tx.send(Err(error));
|
||||
true
|
||||
}
|
||||
PendingEntry::Subscribe { tx, .. } => {
|
||||
let _ = tx.try_send(Err(error));
|
||||
true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn evict_expired(&mut self) -> Vec<String> {
|
||||
let now = Instant::now();
|
||||
let mut evicted = Vec::new();
|
||||
let mut to_remove: Vec<String> = Vec::new();
|
||||
for (id, entry) in self.pending.iter() {
|
||||
let expired = match entry {
|
||||
PendingEntry::Call { timeout, .. } => *timeout <= now,
|
||||
PendingEntry::Subscribe {
|
||||
timeout: Some(t), ..
|
||||
} => *t <= now,
|
||||
PendingEntry::Subscribe { timeout: None, .. } => false,
|
||||
};
|
||||
if expired {
|
||||
to_remove.push(id.clone());
|
||||
}
|
||||
}
|
||||
for id in to_remove {
|
||||
let Some(entry) = self.pending.remove(&id) else {
|
||||
continue;
|
||||
};
|
||||
let timeout_err = CallError::timeout("request timed out");
|
||||
match entry {
|
||||
PendingEntry::Call { tx, .. } => {
|
||||
let _ = tx.send(Err(timeout_err));
|
||||
}
|
||||
PendingEntry::Subscribe { tx, .. } => {
|
||||
let _ = tx.try_send(Err(timeout_err));
|
||||
}
|
||||
}
|
||||
evicted.push(id);
|
||||
}
|
||||
evicted
|
||||
}
|
||||
|
||||
pub fn fail_all(&mut self, error: CallError) -> Vec<String> {
|
||||
let ids: Vec<String> = self.pending.keys().cloned().collect();
|
||||
for id in &ids {
|
||||
if let Some(entry) = self.pending.remove(id) {
|
||||
match entry {
|
||||
PendingEntry::Call { tx, .. } => {
|
||||
let _ = tx.send(Err(error.clone()));
|
||||
}
|
||||
PendingEntry::Subscribe { tx, .. } => {
|
||||
let _ = tx.try_send(Err(error.clone()));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
ids
|
||||
}
|
||||
|
||||
pub fn contains(&self, request_id: &str) -> bool {
|
||||
self.pending.contains_key(request_id)
|
||||
}
|
||||
|
||||
pub(crate) fn parent_of(&self, request_id: &str) -> Option<Option<String>> {
|
||||
self.pending
|
||||
.get(request_id)
|
||||
.map(|e| e.parent_request_id().map(|s| s.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) fn is_started(&self, request_id: &str) -> Option<bool> {
|
||||
self.pending.get(request_id).map(|e| e.started())
|
||||
}
|
||||
|
||||
pub(crate) fn request_ids(&self) -> Vec<String> {
|
||||
self.pending.keys().cloned().collect()
|
||||
}
|
||||
|
||||
pub fn len(&self) -> usize {
|
||||
self.pending.len()
|
||||
}
|
||||
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.pending.is_empty()
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for PendingRequestMap {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serde_json::json;
|
||||
use std::time::Duration;
|
||||
use tokio::time::timeout;
|
||||
|
||||
fn timeout_error() -> CallError {
|
||||
CallError::timeout("request timed out")
|
||||
}
|
||||
|
||||
fn internal_error(message: &str) -> CallError {
|
||||
CallError::internal(message)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn register_call_then_handle_responded_resolves_oneshot() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
let rx = map.register_call(
|
||||
"req-1".to_string(),
|
||||
Instant::now() + Duration::from_secs(30),
|
||||
None,
|
||||
);
|
||||
|
||||
assert!(map.contains("req-1"));
|
||||
assert_eq!(map.len(), 1);
|
||||
|
||||
assert!(map.handle_responded("req-1", json!(42)));
|
||||
|
||||
let result = timeout(Duration::from_millis(100), rx).await;
|
||||
match result {
|
||||
Ok(Ok(Ok(value))) => assert_eq!(value, json!(42)),
|
||||
other => panic!("expected Ok(42), got {other:?}"),
|
||||
}
|
||||
assert!(!map.contains("req-1"));
|
||||
assert_eq!(map.len(), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn register_subscribe_then_handle_responded_pushes_to_channel() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
let mut rx = map.register_subscribe("sub-1".to_string(), None, None);
|
||||
|
||||
assert!(map.handle_responded("sub-1", json!("first")));
|
||||
assert!(map.handle_responded("sub-1", json!("second")));
|
||||
assert!(map.contains("sub-1"));
|
||||
|
||||
let first = timeout(Duration::from_millis(100), rx.recv()).await;
|
||||
let second = timeout(Duration::from_millis(100), rx.recv()).await;
|
||||
match (first, second) {
|
||||
(Ok(Some(Ok(a))), Ok(Some(Ok(b)))) => {
|
||||
assert_eq!(a, json!("first"));
|
||||
assert_eq!(b, json!("second"));
|
||||
}
|
||||
other => panic!("expected two Ok values, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn subscribe_handle_completed_closes_channel_and_deletes_entry() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
let mut rx = map.register_subscribe("sub-2".to_string(), None, None);
|
||||
|
||||
assert!(map.handle_responded("sub-2", json!("a")));
|
||||
assert!(map.handle_completed("sub-2"));
|
||||
assert!(!map.contains("sub-2"));
|
||||
|
||||
let _ = timeout(Duration::from_millis(100), rx.recv()).await;
|
||||
let after_close = timeout(Duration::from_millis(100), rx.recv()).await;
|
||||
match after_close {
|
||||
Ok(None) => {}
|
||||
other => panic!("expected channel closed (None), got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn expired_call_is_evicted_with_timeout_error() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
let rx = map.register_call(
|
||||
"req-2".to_string(),
|
||||
Instant::now() - Duration::from_millis(1),
|
||||
None,
|
||||
);
|
||||
|
||||
let evicted = map.evict_expired();
|
||||
assert_eq!(evicted, vec!["req-2".to_string()]);
|
||||
assert!(!map.contains("req-2"));
|
||||
|
||||
let result = timeout(Duration::from_millis(100), rx).await;
|
||||
match result {
|
||||
Ok(Ok(Err(e))) => {
|
||||
assert_eq!(e.code, "TIMEOUT");
|
||||
assert!(e.retryable);
|
||||
}
|
||||
other => panic!("expected Err(TIMEOUT), got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn expired_subscribe_is_evicted_with_timeout_error() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
let mut rx = map.register_subscribe(
|
||||
"sub-3".to_string(),
|
||||
Some(Instant::now() - Duration::from_millis(1)),
|
||||
None,
|
||||
);
|
||||
|
||||
let evicted = map.evict_expired();
|
||||
assert_eq!(evicted, vec!["sub-3".to_string()]);
|
||||
|
||||
let result = timeout(Duration::from_millis(100), rx.recv()).await;
|
||||
match result {
|
||||
Ok(Some(Err(e))) => {
|
||||
assert_eq!(e.code, "TIMEOUT");
|
||||
assert!(e.retryable);
|
||||
}
|
||||
other => panic!("expected Err(TIMEOUT), got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn unbounded_subscribe_is_not_evicted() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
let _rx = map.register_subscribe("sub-4".to_string(), None, None);
|
||||
|
||||
let evicted = map.evict_expired();
|
||||
assert!(evicted.is_empty());
|
||||
assert!(map.contains("sub-4"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn fail_all_resolves_all_pending_with_internal_error() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
let rx_call = map.register_call(
|
||||
"c-1".to_string(),
|
||||
Instant::now() + Duration::from_secs(30),
|
||||
None,
|
||||
);
|
||||
let mut rx_sub = map.register_subscribe(
|
||||
"s-1".to_string(),
|
||||
Some(Instant::now() + Duration::from_secs(30)),
|
||||
None,
|
||||
);
|
||||
|
||||
let failed = map.fail_all(internal_error("connection closed"));
|
||||
assert_eq!(failed.len(), 2);
|
||||
assert!(failed.contains(&"c-1".to_string()));
|
||||
assert!(failed.contains(&"s-1".to_string()));
|
||||
assert!(map.is_empty());
|
||||
|
||||
let call_result = timeout(Duration::from_millis(100), rx_call).await;
|
||||
match call_result {
|
||||
Ok(Ok(Err(e))) => {
|
||||
assert_eq!(e.code, "INTERNAL");
|
||||
assert_eq!(e.message, "connection closed");
|
||||
}
|
||||
other => panic!("expected Err(INTERNAL), got {other:?}"),
|
||||
}
|
||||
|
||||
let sub_result = timeout(Duration::from_millis(100), rx_sub.recv()).await;
|
||||
match sub_result {
|
||||
Ok(Some(Err(e))) => {
|
||||
assert_eq!(e.code, "INTERNAL");
|
||||
assert_eq!(e.message, "connection closed");
|
||||
}
|
||||
other => panic!("expected Err(INTERNAL), got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn handle_responded_unknown_request_id_returns_false() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
assert!(!map.handle_responded("nonexistent", json!(1)));
|
||||
assert_eq!(map.len(), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn handle_completed_unknown_request_id_returns_false() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
assert!(!map.handle_completed("nonexistent"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn handle_aborted_unknown_request_id_returns_false() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
assert!(!map.handle_aborted("nonexistent"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn handle_error_unknown_request_id_returns_false() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
assert!(!map.handle_error("nonexistent", internal_error("x")));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn handle_aborted_cancels_pending_call() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
let rx = map.register_call(
|
||||
"req-3".to_string(),
|
||||
Instant::now() + Duration::from_secs(30),
|
||||
None,
|
||||
);
|
||||
|
||||
assert!(map.handle_aborted("req-3"));
|
||||
assert!(!map.contains("req-3"));
|
||||
|
||||
let result = timeout(Duration::from_millis(100), rx).await;
|
||||
match result {
|
||||
Ok(Err(_)) => {}
|
||||
other => panic!("expected sender dropped (Err), got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn handle_error_resolves_call_with_error() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
let rx = map.register_call(
|
||||
"req-4".to_string(),
|
||||
Instant::now() + Duration::from_secs(30),
|
||||
None,
|
||||
);
|
||||
|
||||
let err = CallError::new("FILE_NOT_FOUND", "missing", false);
|
||||
assert!(map.handle_error("req-4", err.clone()));
|
||||
assert!(!map.contains("req-4"));
|
||||
|
||||
let result = timeout(Duration::from_millis(100), rx).await;
|
||||
match result {
|
||||
Ok(Ok(Err(e))) => {
|
||||
assert_eq!(e.code, "FILE_NOT_FOUND");
|
||||
assert!(!e.retryable);
|
||||
}
|
||||
other => panic!("expected Err(FILE_NOT_FOUND), got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn handle_error_pushes_to_subscribe_channel() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
let mut rx = map.register_subscribe("sub-5".to_string(), None, None);
|
||||
|
||||
let err = CallError::new("RATE_LIMITED", "too fast", true);
|
||||
assert!(map.handle_error("sub-5", err.clone()));
|
||||
assert!(!map.contains("sub-5"));
|
||||
|
||||
let result = timeout(Duration::from_millis(100), rx.recv()).await;
|
||||
match result {
|
||||
Ok(Some(Err(e))) => {
|
||||
assert_eq!(e.code, "RATE_LIMITED");
|
||||
assert!(e.retryable);
|
||||
}
|
||||
other => panic!("expected Err(RATE_LIMITED), got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn correlation_by_id_not_by_stream() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
let rx = map.register_call(
|
||||
"req-stream-3".to_string(),
|
||||
Instant::now() + Duration::from_secs(30),
|
||||
None,
|
||||
);
|
||||
|
||||
assert!(map.handle_responded("req-stream-3", json!("response-from-stream-7")));
|
||||
let result = timeout(Duration::from_millis(100), rx).await;
|
||||
match result {
|
||||
Ok(Ok(Ok(value))) => assert_eq!(value, json!("response-from-stream-7")),
|
||||
other => panic!("expected Ok, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn register_call_overwrites_existing_entry() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
let _rx_old = map.register_call(
|
||||
"req-5".to_string(),
|
||||
Instant::now() + Duration::from_secs(30),
|
||||
None,
|
||||
);
|
||||
let rx_new = map.register_call(
|
||||
"req-5".to_string(),
|
||||
Instant::now() + Duration::from_secs(30),
|
||||
None,
|
||||
);
|
||||
assert_eq!(map.len(), 1);
|
||||
|
||||
assert!(map.handle_responded("req-5", json!("new")));
|
||||
let result = timeout(Duration::from_millis(100), rx_new).await;
|
||||
match result {
|
||||
Ok(Ok(Ok(value))) => assert_eq!(value, json!("new")),
|
||||
other => panic!("expected Ok from new receiver, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn evict_expired_skips_non_expired_entries() {
|
||||
let mut map = PendingRequestMap::new();
|
||||
let _rx_expired = map.register_call(
|
||||
"expired".to_string(),
|
||||
Instant::now() - Duration::from_millis(1),
|
||||
None,
|
||||
);
|
||||
let _rx_alive = map.register_call(
|
||||
"alive".to_string(),
|
||||
Instant::now() + Duration::from_secs(60),
|
||||
None,
|
||||
);
|
||||
|
||||
let evicted = map.evict_expired();
|
||||
assert_eq!(evicted, vec!["expired".to_string()]);
|
||||
assert!(map.contains("alive"));
|
||||
assert!(!map.contains("expired"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn default_is_empty_map() {
|
||||
let map = PendingRequestMap::default();
|
||||
assert!(map.is_empty());
|
||||
assert_eq!(map.len(), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn timeout_error_helper() {
|
||||
let err = timeout_error();
|
||||
assert_eq!(err.code, "TIMEOUT");
|
||||
assert!(err.retryable);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
//! Shared test helpers for the call protocol's inline `#[cfg(test)]`
|
||||
//! modules. Kept here (not in each test module) so the `stub_connection()`
|
||||
//! shape is defined once — `Connection::from_stream` was removed (ADR-092)
|
||||
//! and every test stub that previously called it now calls
|
||||
//! `Connection::from_bidi(SinkEmpty, ...)` via `sink_empty_connection()`.
|
||||
|
||||
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
|
||||
use std::pin::Pin;
|
||||
use std::task::{Context, Poll};
|
||||
|
||||
use crate::core::types::Connection;
|
||||
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
|
||||
|
||||
/// A test-only `AsyncRead + AsyncWrite` pair equivalent to
|
||||
/// `tokio::io::sink() + tokio::io::empty()`: reads yield EOF immediately
|
||||
/// (zero bytes), writes discard. Exists because `Connection::from_bidi`
|
||||
/// (ADR-092 — the only public stream constructor, replacing
|
||||
/// `from_stream`) requires a single value that implements both traits.
|
||||
/// Used only to construct a `Connection` for tests that exercise
|
||||
/// `Connection`-level state (alpn, addr, identity, dispatcher run loop
|
||||
/// with an immediately-closed accept stream) without ever reading or
|
||||
/// writing real bytes.
|
||||
pub(crate) struct SinkEmpty;
|
||||
|
||||
impl AsyncRead for SinkEmpty {
|
||||
fn poll_read(
|
||||
self: Pin<&mut Self>,
|
||||
_cx: &mut Context<'_>,
|
||||
_buf: &mut ReadBuf<'_>,
|
||||
) -> Poll<std::io::Result<()>> {
|
||||
// EOF immediately — mirrors `tokio::io::empty()`.
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
}
|
||||
|
||||
impl AsyncWrite for SinkEmpty {
|
||||
fn poll_write(
|
||||
self: Pin<&mut Self>,
|
||||
_cx: &mut Context<'_>,
|
||||
buf: &[u8],
|
||||
) -> Poll<std::io::Result<usize>> {
|
||||
// Discard — mirrors `tokio::io::sink()`.
|
||||
Poll::Ready(Ok(buf.len()))
|
||||
}
|
||||
|
||||
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
|
||||
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
}
|
||||
|
||||
/// Construct a `Connection` whose `accept_bi` yields a `SinkEmpty` once,
|
||||
/// then `ConnectionClosed`. Used by tests that need a `Connection` for
|
||||
/// `CallConnection::new(conn)` or `adapter.handle(conn, &auth)` without
|
||||
/// exercising the wire protocol — `SinkEmpty` reads EOF (so the dispatch
|
||||
/// loop closes immediately) and discards writes.
|
||||
pub(crate) fn sink_empty_connection() -> Connection {
|
||||
Connection::from_bidi(
|
||||
SinkEmpty,
|
||||
b"alknet/call".to_vec(),
|
||||
Some(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 4321)),
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,548 @@
|
||||
//! Wire format: `EventEnvelope`, `ResponseEnvelope`, `CallError`, and
|
||||
//! length-prefixed JSON framing.
|
||||
//!
|
||||
//! See `docs/architecture/crates/call/call-protocol.md` for the full
|
||||
//! specification.
|
||||
|
||||
use std::io;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
|
||||
|
||||
pub const EVENT_REQUESTED: &str = "call.requested";
|
||||
pub const EVENT_RESPONDED: &str = "call.responded";
|
||||
pub const EVENT_COMPLETED: &str = "call.completed";
|
||||
pub const EVENT_ABORTED: &str = "call.aborted";
|
||||
pub const EVENT_ERROR: &str = "call.error";
|
||||
|
||||
const LENGTH_PREFIX_BYTES: usize = 4;
|
||||
const MAX_FRAME_SIZE: u32 = 64 * 1024 * 1024;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct EventEnvelope {
|
||||
#[serde(rename = "type")]
|
||||
pub r#type: String,
|
||||
pub id: String,
|
||||
pub payload: Value,
|
||||
}
|
||||
|
||||
impl EventEnvelope {
|
||||
pub fn new(event_type: impl Into<String>, id: impl Into<String>, payload: Value) -> Self {
|
||||
Self {
|
||||
r#type: event_type.into(),
|
||||
id: id.into(),
|
||||
payload,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn requested(id: impl Into<String>, payload: Value) -> Self {
|
||||
Self::new(EVENT_REQUESTED, id, payload)
|
||||
}
|
||||
|
||||
pub fn responded(id: impl Into<String>, output: Value) -> Self {
|
||||
Self::new(EVENT_RESPONDED, id, serde_json::json!({ "output": output }))
|
||||
}
|
||||
|
||||
pub fn completed(id: impl Into<String>) -> Self {
|
||||
Self::new(EVENT_COMPLETED, id, serde_json::json!({}))
|
||||
}
|
||||
|
||||
pub fn aborted(id: impl Into<String>) -> Self {
|
||||
Self::new(EVENT_ABORTED, id, serde_json::json!({}))
|
||||
}
|
||||
|
||||
pub fn error(id: impl Into<String>, error: &CallError) -> Self {
|
||||
let payload = serde_json::to_value(error).unwrap_or(Value::Null);
|
||||
Self::new(EVENT_ERROR, id, payload)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct CallError {
|
||||
pub code: String,
|
||||
pub message: String,
|
||||
pub retryable: bool,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub details: Option<Value>,
|
||||
}
|
||||
|
||||
impl CallError {
|
||||
pub fn new(code: impl Into<String>, message: impl Into<String>, retryable: bool) -> Self {
|
||||
Self {
|
||||
code: code.into(),
|
||||
message: message.into(),
|
||||
retryable,
|
||||
details: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_details(mut self, details: Value) -> Self {
|
||||
self.details = Some(details);
|
||||
self
|
||||
}
|
||||
|
||||
pub fn not_found(op_name: &str) -> Self {
|
||||
Self::new(
|
||||
"NOT_FOUND",
|
||||
format!("operation not found: {op_name}"),
|
||||
false,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn forbidden(message: impl Into<String>) -> Self {
|
||||
Self::new("FORBIDDEN", message, false)
|
||||
}
|
||||
|
||||
pub fn invalid_input(message: impl Into<String>) -> Self {
|
||||
Self::new("INVALID_INPUT", message, false)
|
||||
}
|
||||
|
||||
pub fn internal(message: impl Into<String>) -> Self {
|
||||
Self::new("INTERNAL", message, false)
|
||||
}
|
||||
|
||||
pub fn timeout(message: impl Into<String>) -> Self {
|
||||
Self::new("TIMEOUT", message, true)
|
||||
}
|
||||
|
||||
pub fn invalid_operation_type(message: impl Into<String>) -> Self {
|
||||
Self::new("INVALID_OPERATION_TYPE", message, false)
|
||||
}
|
||||
}
|
||||
|
||||
impl Eq for CallError {}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct ResponseEnvelope {
|
||||
pub request_id: String,
|
||||
pub result: Result<Value, CallError>,
|
||||
}
|
||||
|
||||
impl ResponseEnvelope {
|
||||
pub fn ok(request_id: impl Into<String>, output: Value) -> Self {
|
||||
Self {
|
||||
request_id: request_id.into(),
|
||||
result: Ok(output),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn error(request_id: impl Into<String>, error: CallError) -> Self {
|
||||
Self {
|
||||
request_id: request_id.into(),
|
||||
result: Err(error),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn not_found(request_id: impl Into<String>, op_name: &str) -> Self {
|
||||
Self::error(request_id, CallError::not_found(op_name))
|
||||
}
|
||||
|
||||
pub fn forbidden(request_id: impl Into<String>, message: impl Into<String>) -> Self {
|
||||
Self::error(request_id, CallError::forbidden(message))
|
||||
}
|
||||
|
||||
pub fn into_event(self) -> EventEnvelope {
|
||||
let id = self.request_id;
|
||||
match self.result {
|
||||
Ok(output) => EventEnvelope::responded(id, output),
|
||||
Err(ref err) => EventEnvelope::error(id, err),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<ResponseEnvelope> for EventEnvelope {
|
||||
fn from(envelope: ResponseEnvelope) -> EventEnvelope {
|
||||
envelope.into_event()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum FrameError {
|
||||
#[error("io error: {0}")]
|
||||
Io(#[from] io::Error),
|
||||
#[error("json error: {0}")]
|
||||
Json(#[from] serde_json::Error),
|
||||
#[error("connection closed")]
|
||||
ConnectionClosed,
|
||||
#[error("invalid frame")]
|
||||
InvalidFrame,
|
||||
}
|
||||
|
||||
pub struct FrameFramedReader<R: AsyncRead + Unpin> {
|
||||
reader: R,
|
||||
len_buf: [u8; LENGTH_PREFIX_BYTES],
|
||||
}
|
||||
|
||||
impl<R: AsyncRead + Unpin> FrameFramedReader<R> {
|
||||
pub fn new(reader: R) -> Self {
|
||||
Self {
|
||||
reader,
|
||||
len_buf: [0u8; LENGTH_PREFIX_BYTES],
|
||||
}
|
||||
}
|
||||
|
||||
pub fn into_inner(self) -> R {
|
||||
self.reader
|
||||
}
|
||||
|
||||
pub async fn read_frame(&mut self) -> Result<EventEnvelope, FrameError> {
|
||||
match self.reader.read_exact(&mut self.len_buf).await {
|
||||
Ok(_) => {}
|
||||
Err(e) if e.kind() == io::ErrorKind::UnexpectedEof => {
|
||||
return Err(FrameError::ConnectionClosed);
|
||||
}
|
||||
Err(e) => return Err(FrameError::Io(e)),
|
||||
}
|
||||
|
||||
let length = u32::from_be_bytes(self.len_buf);
|
||||
if length == 0 {
|
||||
return Err(FrameError::InvalidFrame);
|
||||
}
|
||||
if length > MAX_FRAME_SIZE {
|
||||
return Err(FrameError::InvalidFrame);
|
||||
}
|
||||
|
||||
let mut body = vec![0u8; length as usize];
|
||||
match self.reader.read_exact(&mut body).await {
|
||||
Ok(_) => {}
|
||||
Err(e) if e.kind() == io::ErrorKind::UnexpectedEof => {
|
||||
return Err(FrameError::ConnectionClosed);
|
||||
}
|
||||
Err(e) => return Err(FrameError::Io(e)),
|
||||
}
|
||||
|
||||
let envelope: EventEnvelope = serde_json::from_slice(&body)?;
|
||||
Ok(envelope)
|
||||
}
|
||||
}
|
||||
|
||||
pub struct FrameFramedWriter<W: AsyncWrite + Unpin> {
|
||||
writer: W,
|
||||
}
|
||||
|
||||
impl<W: AsyncWrite + Unpin> FrameFramedWriter<W> {
|
||||
pub fn new(writer: W) -> Self {
|
||||
Self { writer }
|
||||
}
|
||||
|
||||
pub fn into_inner(self) -> W {
|
||||
self.writer
|
||||
}
|
||||
|
||||
pub async fn write_frame(&mut self, envelope: &EventEnvelope) -> Result<(), FrameError> {
|
||||
let body = serde_json::to_vec(envelope)?;
|
||||
let len = body.len();
|
||||
if len > MAX_FRAME_SIZE as usize {
|
||||
return Err(FrameError::InvalidFrame);
|
||||
}
|
||||
let len_bytes = (len as u32).to_be_bytes();
|
||||
self.writer.write_all(&len_bytes).await?;
|
||||
self.writer.write_all(&body).await?;
|
||||
self.writer.flush().await?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use tokio::io::{duplex, AsyncReadExt};
|
||||
|
||||
fn sample_envelope() -> EventEnvelope {
|
||||
EventEnvelope::new(
|
||||
"call.requested",
|
||||
"req-1",
|
||||
serde_json::json!({
|
||||
"operationId": "/fs/readFile",
|
||||
"input": { "path": "/etc/hosts" }
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn round_trip_envelope() {
|
||||
let (client, server) = duplex(8 * 1024);
|
||||
let envelope = sample_envelope();
|
||||
|
||||
let mut writer = FrameFramedWriter::new(client);
|
||||
writer.write_frame(&envelope).await.unwrap();
|
||||
drop(writer);
|
||||
|
||||
let mut reader = FrameFramedReader::new(server);
|
||||
let read = reader.read_frame().await.unwrap();
|
||||
assert_eq!(read, envelope);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn round_trip_multiple_frames() {
|
||||
let (client, server) = duplex(8 * 1024);
|
||||
|
||||
let envelopes = vec![
|
||||
EventEnvelope::responded("a", Value::String("hello".into())),
|
||||
EventEnvelope::completed("a"),
|
||||
EventEnvelope::aborted("b"),
|
||||
];
|
||||
|
||||
{
|
||||
let mut writer = FrameFramedWriter::new(client);
|
||||
for e in &envelopes {
|
||||
writer.write_frame(e).await.unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
let mut reader = FrameFramedReader::new(server);
|
||||
for expected in envelopes {
|
||||
let read = reader.read_frame().await.unwrap();
|
||||
assert_eq!(read, expected);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn read_frame_on_closed_reader_returns_connection_closed() {
|
||||
let (_, server) = duplex(8 * 1024);
|
||||
let mut reader = FrameFramedReader::new(server);
|
||||
match reader.read_frame().await {
|
||||
Err(FrameError::ConnectionClosed) => {}
|
||||
other => panic!("expected ConnectionClosed, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn truncated_body_returns_connection_closed() {
|
||||
let (mut client, server) = duplex(8 * 1024);
|
||||
let envelope = sample_envelope();
|
||||
let body = serde_json::to_vec(&envelope).unwrap();
|
||||
let len_bytes = (body.len() as u32).to_be_bytes();
|
||||
client.write_all(&len_bytes).await.unwrap();
|
||||
client.write_all(&body[..body.len() / 2]).await.unwrap();
|
||||
drop(client);
|
||||
|
||||
let mut reader = FrameFramedReader::new(server);
|
||||
match reader.read_frame().await {
|
||||
Err(FrameError::ConnectionClosed) => {}
|
||||
other => panic!("expected ConnectionClosed, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn zero_length_frame_is_invalid() {
|
||||
let (mut client, server) = duplex(8 * 1024);
|
||||
client.write_all(&[0u8, 0, 0, 0]).await.unwrap();
|
||||
drop(client);
|
||||
|
||||
let mut reader = FrameFramedReader::new(server);
|
||||
match reader.read_frame().await {
|
||||
Err(FrameError::InvalidFrame) => {}
|
||||
other => panic!("expected InvalidFrame, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn oversized_frame_is_invalid() {
|
||||
let (mut client, server) = duplex(8 * 1024);
|
||||
let too_big = (MAX_FRAME_SIZE + 1u32).to_be_bytes();
|
||||
client.write_all(&too_big).await.unwrap();
|
||||
drop(client);
|
||||
|
||||
let mut reader = FrameFramedReader::new(server);
|
||||
match reader.read_frame().await {
|
||||
Err(FrameError::InvalidFrame) => {}
|
||||
other => panic!("expected InvalidFrame, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn framing_handles_large_payload() {
|
||||
let (client, server) = duplex(1024 * 1024);
|
||||
let big = "x".repeat(64 * 1024);
|
||||
let envelope = EventEnvelope::responded("big", Value::String(big.clone()));
|
||||
|
||||
let mut writer = FrameFramedWriter::new(client);
|
||||
writer.write_frame(&envelope).await.unwrap();
|
||||
drop(writer);
|
||||
|
||||
let mut reader = FrameFramedReader::new(server);
|
||||
let read = reader.read_frame().await.unwrap();
|
||||
assert_eq!(read, envelope);
|
||||
match read.payload {
|
||||
Value::Object(map) => match map.get("output") {
|
||||
Some(Value::String(s)) => assert_eq!(s, &big),
|
||||
other => panic!("expected output string, got {other:?}"),
|
||||
},
|
||||
other => panic!("expected object payload, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn response_envelope_ok_produces_call_responded_event() {
|
||||
let response = ResponseEnvelope::ok("req-1", Value::String("hi".into()));
|
||||
let event: EventEnvelope = response.into();
|
||||
assert_eq!(event.r#type, EVENT_RESPONDED);
|
||||
assert_eq!(event.id, "req-1");
|
||||
let map = event.payload.as_object().expect("payload is object");
|
||||
assert_eq!(map.get("output"), Some(&Value::String("hi".into())));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn response_envelope_error_produces_call_error_event() {
|
||||
let err = CallError::new("FILE_NOT_FOUND", "file not found: /etc/x", false)
|
||||
.with_details(serde_json::json!({ "path": "/etc/x" }));
|
||||
let response = ResponseEnvelope::error("req-2", err);
|
||||
let event: EventEnvelope = response.into();
|
||||
assert_eq!(event.r#type, EVENT_ERROR);
|
||||
assert_eq!(event.id, "req-2");
|
||||
assert_eq!(
|
||||
event.payload.get("code"),
|
||||
Some(&Value::String("FILE_NOT_FOUND".into()))
|
||||
);
|
||||
assert_eq!(
|
||||
event.payload.get("message"),
|
||||
Some(&Value::String("file not found: /etc/x".into()))
|
||||
);
|
||||
assert_eq!(event.payload.get("retryable"), Some(&Value::Bool(false)));
|
||||
assert_eq!(
|
||||
event.payload.get("details"),
|
||||
Some(&serde_json::json!({ "path": "/etc/x" }))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn response_envelope_not_found_helper() {
|
||||
let response = ResponseEnvelope::not_found("req-3", "fs/missing");
|
||||
assert_eq!(response.request_id, "req-3");
|
||||
match &response.result {
|
||||
Err(e) => {
|
||||
assert_eq!(e.code, "NOT_FOUND");
|
||||
assert!(!e.retryable);
|
||||
assert!(e.message.contains("fs/missing"));
|
||||
}
|
||||
other => panic!("expected Err, got {other:?}"),
|
||||
}
|
||||
let event: EventEnvelope = response.into();
|
||||
assert_eq!(event.r#type, EVENT_ERROR);
|
||||
assert_eq!(event.id, "req-3");
|
||||
assert_eq!(
|
||||
event.payload.get("code"),
|
||||
Some(&Value::String("NOT_FOUND".into()))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn response_envelope_forbidden_helper() {
|
||||
let response = ResponseEnvelope::forbidden("req-4", "authentication required");
|
||||
match &response.result {
|
||||
Err(e) => {
|
||||
assert_eq!(e.code, "FORBIDDEN");
|
||||
assert_eq!(e.message, "authentication required");
|
||||
}
|
||||
other => panic!("expected Err, got {other:?}"),
|
||||
}
|
||||
let event: EventEnvelope = response.into();
|
||||
assert_eq!(event.r#type, EVENT_ERROR);
|
||||
assert_eq!(event.id, "req-4");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn event_envelope_completed_has_empty_payload() {
|
||||
let event = EventEnvelope::completed("sub-1");
|
||||
assert_eq!(event.r#type, EVENT_COMPLETED);
|
||||
assert_eq!(event.id, "sub-1");
|
||||
assert_eq!(event.payload, serde_json::json!({}));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn event_envelope_aborted_has_empty_payload() {
|
||||
let event = EventEnvelope::aborted("req-9");
|
||||
assert_eq!(event.r#type, EVENT_ABORTED);
|
||||
assert_eq!(event.id, "req-9");
|
||||
assert_eq!(event.payload, serde_json::json!({}));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn event_envelope_responded_wraps_output() {
|
||||
let event = EventEnvelope::responded("req-1", Value::Number(42.into()));
|
||||
assert_eq!(event.r#type, EVENT_RESPONDED);
|
||||
assert_eq!(event.payload.get("output"), Some(&Value::Number(42.into())));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn event_envelope_serializes_type_field() {
|
||||
let event = sample_envelope();
|
||||
let json = serde_json::to_string(&event).unwrap();
|
||||
assert!(json.contains("\"type\":\"call.requested\""));
|
||||
assert!(!json.contains("\"r#type\""));
|
||||
|
||||
let parsed: EventEnvelope = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(parsed, event);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn call_error_skips_missing_details() {
|
||||
let err = CallError::new("INTERNAL", "boom", false);
|
||||
let json = serde_json::to_string(&err).unwrap();
|
||||
assert!(!json.contains("details"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn read_after_eof_then_eof_returns_connection_closed() {
|
||||
let mut data = Vec::new();
|
||||
let envelope = EventEnvelope::responded("one", Value::Null);
|
||||
let body = serde_json::to_vec(&envelope).unwrap();
|
||||
data.extend_from_slice(&(body.len() as u32).to_be_bytes());
|
||||
data.extend_from_slice(&body);
|
||||
let cursor = std::io::Cursor::new(data);
|
||||
let mut reader = FrameFramedReader::new(cursor);
|
||||
let first = reader.read_frame().await.unwrap();
|
||||
assert_eq!(first, envelope);
|
||||
match reader.read_frame().await {
|
||||
Err(FrameError::ConnectionClosed) => {}
|
||||
other => panic!("expected ConnectionClosed, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn writer_into_inner_recovers_stream() {
|
||||
let (client, server) = duplex(8 * 1024);
|
||||
let envelope = sample_envelope();
|
||||
let mut writer = FrameFramedWriter::new(client);
|
||||
writer.write_frame(&envelope).await.unwrap();
|
||||
let mut recovered = writer.into_inner();
|
||||
recovered.shutdown().await.unwrap();
|
||||
drop(recovered);
|
||||
|
||||
let mut reader = FrameFramedReader::new(server);
|
||||
let read = reader.read_frame().await.unwrap();
|
||||
assert_eq!(read, envelope);
|
||||
let _ = reader.into_inner();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn reader_handles_partial_length_prefix() {
|
||||
let (mut client, server) = duplex(8 * 1024);
|
||||
client.write_all(&[0u8, 0]).await.unwrap();
|
||||
drop(client);
|
||||
let mut reader = FrameFramedReader::new(server);
|
||||
match reader.read_frame().await {
|
||||
Err(FrameError::ConnectionClosed) => {}
|
||||
other => panic!("expected ConnectionClosed, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn reader_drains_remaining_after_read() {
|
||||
let mut data = Vec::new();
|
||||
let envelope = sample_envelope();
|
||||
let body = serde_json::to_vec(&envelope).unwrap();
|
||||
data.extend_from_slice(&(body.len() as u32).to_be_bytes());
|
||||
data.extend_from_slice(&body);
|
||||
data.extend_from_slice(&[9u8; 4]);
|
||||
let mut cursor = tokio::io::BufReader::new(std::io::Cursor::new(data));
|
||||
let mut reader = FrameFramedReader::new(&mut cursor);
|
||||
let read = reader.read_frame().await.unwrap();
|
||||
assert_eq!(read, envelope);
|
||||
let mut leftover = Vec::new();
|
||||
let _ = cursor.read_to_end(&mut leftover).await.unwrap();
|
||||
assert_eq!(leftover, vec![9u8; 4]);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,313 @@
|
||||
use std::collections::{HashMap, HashSet};
|
||||
use std::sync::Arc;
|
||||
use std::time::Instant;
|
||||
|
||||
use crate::core::auth::Identity;
|
||||
use crate::core::ownership::OwnershipProvider;
|
||||
use crate::core::types::Capabilities;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::env::{OperationEnv, PeerId, PeerRef};
|
||||
|
||||
pub struct OperationContext {
|
||||
pub request_id: String,
|
||||
pub parent_request_id: Option<String>,
|
||||
pub identity: Option<Identity>,
|
||||
pub handler_identity: Option<CompositionAuthority>,
|
||||
/// The original caller when this call was forwarded by a `from_call`
|
||||
/// handler (ADR-032). **Metadata only** — `AccessControl::check` never
|
||||
/// reads it; the ACL always authorizes `identity` (the direct caller).
|
||||
/// Handlers may read it for logging, auditing, per-user rate limiting,
|
||||
/// or application context. Populated from
|
||||
/// `call.requested.forwarded_for` by the dispatch path; set to `None`
|
||||
/// for composed children (wire-ingress only, not composition-ingress).
|
||||
/// The forwarder's claim, not a verified identity — a malicious hub can
|
||||
/// lie (same property as HTTP `X-Forwarded-For`). See ADR-032.
|
||||
pub forwarded_for: Option<Identity>,
|
||||
pub capabilities: Capabilities,
|
||||
pub metadata: HashMap<String, Value>,
|
||||
pub scoped_env: ScopedPeerEnv,
|
||||
pub env: Arc<dyn OperationEnv + Send + Sync>,
|
||||
pub abort_policy: AbortPolicy,
|
||||
pub deadline: Option<Instant>,
|
||||
pub internal: bool,
|
||||
/// `None` when no ownership provider is wired (backward compat —
|
||||
/// `check` falls back to static `Identity.resources` path). Wired by
|
||||
/// the assembly layer via `CallAdapter`/`Dispatcher` (ADR-050).
|
||||
pub ownership: Option<Arc<dyn OwnershipProvider>>,
|
||||
}
|
||||
|
||||
impl OperationContext {
|
||||
pub fn is_internal(&self) -> bool {
|
||||
self.internal
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
|
||||
pub enum AbortPolicy {
|
||||
#[default]
|
||||
AbortDependents,
|
||||
ContinueRunning,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct CompositionAuthority {
|
||||
pub label: String,
|
||||
pub scopes: Vec<String>,
|
||||
pub resources: HashMap<String, Vec<String>>,
|
||||
}
|
||||
|
||||
impl CompositionAuthority {
|
||||
pub fn none() -> Option<Self> {
|
||||
None
|
||||
}
|
||||
|
||||
pub fn new(label: &str, scopes: impl IntoIterator<Item = String>) -> Self {
|
||||
Self {
|
||||
label: label.to_string(),
|
||||
scopes: scopes.into_iter().collect(),
|
||||
resources: HashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn as_identity(&self) -> Option<Identity> {
|
||||
Some(Identity {
|
||||
id: self.label.clone(),
|
||||
scopes: self.scopes.clone(),
|
||||
resources: self.resources.clone(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ScopedPeerEnv {
|
||||
/// Peer-agnostic reachability — reachable via `PeerRef::Any` or
|
||||
/// `PeerRef::Specific(any)`. The common case (peer-agnostic composition).
|
||||
pub allowed_ops: HashSet<String>,
|
||||
/// Peer-pinned reachability — `"peer-id/op-name"`, reachable only via
|
||||
/// `PeerRef::Specific(that peer)`. Additive to `allowed_ops`; opt-in for
|
||||
/// the disambiguation case (ADR-029 §4).
|
||||
pub peer_pinned: HashSet<String>,
|
||||
}
|
||||
|
||||
impl ScopedPeerEnv {
|
||||
pub fn empty() -> Self {
|
||||
Self {
|
||||
allowed_ops: HashSet::new(),
|
||||
peer_pinned: HashSet::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn new(ops: impl IntoIterator<Item = impl Into<String>>) -> Self {
|
||||
Self {
|
||||
allowed_ops: ops.into_iter().map(|s| s.into()).collect(),
|
||||
peer_pinned: HashSet::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Peer-pinned reachability: `"peer-id/op-name"`. Reachable only via
|
||||
/// `PeerRef::Specific(that peer)`. Additive to `new` — call `new` for the
|
||||
/// peer-agnostic set, then `with_pinned` for the pinned set.
|
||||
pub fn with_pinned(mut self, pinned: impl IntoIterator<Item = impl Into<String>>) -> Self {
|
||||
self.peer_pinned = pinned.into_iter().map(|s| s.into()).collect();
|
||||
self
|
||||
}
|
||||
|
||||
/// Peer-agnostic reachability — unchanged from `ScopedOperationEnv::allows`.
|
||||
/// A name here is reachable via any routing path (`PeerRef::Any` or
|
||||
/// `Specific`).
|
||||
pub fn allows(&self, name: &str) -> bool {
|
||||
self.allowed_ops.contains(name)
|
||||
}
|
||||
|
||||
/// Peer-pinned reachability — reachable only via `PeerRef::Specific(peer)`.
|
||||
/// The entry shape is `"peer-id/op-name"` (ADR-029 §4, OQ-33).
|
||||
pub fn allows_pinned(&self, peer: &PeerId, name: &str) -> bool {
|
||||
self.peer_pinned.contains(&format!("{peer}/{name}"))
|
||||
}
|
||||
|
||||
/// Does this scoped env permit `name` via `peer`? Used by the reachability
|
||||
/// gate in `invoke_peer` / `invoke_with_policy`.
|
||||
/// - `PeerRef::Any` → `allows(name)`
|
||||
/// - `PeerRef::Specific(peer)` → `allows(name) || allows_pinned(peer, name)`
|
||||
pub fn allows_via(&self, peer: &PeerRef, name: &str) -> bool {
|
||||
match peer {
|
||||
PeerRef::Any => self.allows(name),
|
||||
PeerRef::Specific(p) => self.allows(name) || self.allows_pinned(p, name),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for ScopedPeerEnv {
|
||||
fn default() -> Self {
|
||||
Self::empty()
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub(crate) fn generate_request_id() -> String {
|
||||
uuid::Uuid::new_v4().to_string()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn scoped_env_allows_in_set() {
|
||||
let env = ScopedPeerEnv::new(["fs/readFile", "agent/chat"]);
|
||||
assert!(env.allows("fs/readFile"));
|
||||
assert!(env.allows("agent/chat"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn scoped_env_disallows_not_in_set() {
|
||||
let env = ScopedPeerEnv::new(["fs/readFile"]);
|
||||
assert!(!env.allows("agent/chat"));
|
||||
assert!(!env.allows(""));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn scoped_env_empty_allows_nothing() {
|
||||
let env = ScopedPeerEnv::empty();
|
||||
assert!(!env.allows("fs/readFile"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn scoped_peer_env_new_with_pinned_populates_both_fields() {
|
||||
let env = ScopedPeerEnv::new(["fs/readFile"]).with_pinned(["worker-a/container/exec"]);
|
||||
assert!(env.allowed_ops.contains("fs/readFile"));
|
||||
assert!(env.peer_pinned.contains("worker-a/container/exec"));
|
||||
assert!(!env.allowed_ops.contains("worker-a/container/exec"));
|
||||
assert!(!env.peer_pinned.contains("fs/readFile"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn scoped_peer_env_allows_checks_allowed_ops_only() {
|
||||
let env = ScopedPeerEnv::empty().with_pinned(["worker-a/container/exec"]);
|
||||
assert!(
|
||||
!env.allows("container/exec"),
|
||||
"pinned-only op not in allowed_ops"
|
||||
);
|
||||
let env2 = ScopedPeerEnv::new(["container/exec"]).with_pinned(["worker-a/container/exec"]);
|
||||
assert!(
|
||||
env2.allows("container/exec"),
|
||||
"op in allowed_ops is allowed"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn scoped_peer_env_allows_pinned_checks_peer_pinned_shape() {
|
||||
let env = ScopedPeerEnv::empty().with_pinned(["worker-a/container/exec"]);
|
||||
assert!(env.allows_pinned(&"worker-a".to_string(), "container/exec"));
|
||||
assert!(
|
||||
!env.allows_pinned(&"worker-b".to_string(), "container/exec"),
|
||||
"wrong peer"
|
||||
);
|
||||
assert!(
|
||||
!env.allows_pinned(&"worker-a".to_string(), "other/op"),
|
||||
"wrong op"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn scoped_peer_env_allows_via_any_uses_allowed_ops_only() {
|
||||
let env = ScopedPeerEnv::new(["fs/readFile"]).with_pinned(["worker-a/container/exec"]);
|
||||
assert!(
|
||||
env.allows_via(&PeerRef::Any, "fs/readFile"),
|
||||
"allowed op via Any"
|
||||
);
|
||||
assert!(
|
||||
!env.allows_via(&PeerRef::Any, "container/exec"),
|
||||
"pinned-only op NOT reachable via Any"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn scoped_peer_env_allows_via_specific_uses_allowed_ops_or_peer_pinned() {
|
||||
let env = ScopedPeerEnv::new(["fs/readFile"]).with_pinned(["worker-a/container/exec"]);
|
||||
assert!(
|
||||
env.allows_via(&PeerRef::Specific("worker-a".to_string()), "container/exec"),
|
||||
"pinned-only op reachable via Specific(pinned peer)"
|
||||
);
|
||||
assert!(
|
||||
env.allows_via(&PeerRef::Specific("worker-a".to_string()), "fs/readFile"),
|
||||
"allowed op reachable via Specific(any peer)"
|
||||
);
|
||||
assert!(
|
||||
!env.allows_via(&PeerRef::Specific("worker-b".to_string()), "container/exec"),
|
||||
"pinned-only op NOT reachable via Specific(wrong peer)"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn scoped_peer_env_op_in_both_sets_reachable_via_both_any_and_specific() {
|
||||
let env = ScopedPeerEnv::new(["container/exec"]).with_pinned(["worker-a/container/exec"]);
|
||||
assert!(
|
||||
env.allows_via(&PeerRef::Any, "container/exec"),
|
||||
"op in allowed_ops reachable via Any"
|
||||
);
|
||||
assert!(
|
||||
env.allows_via(&PeerRef::Specific("worker-a".to_string()), "container/exec"),
|
||||
"op in both sets reachable via Specific(peer)"
|
||||
);
|
||||
assert!(
|
||||
env.allows_via(&PeerRef::Specific("worker-b".to_string()), "container/exec"),
|
||||
"op in allowed_ops reachable via Specific(other peer) too"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn composition_authority_as_identity_correct() {
|
||||
let mut resources = HashMap::new();
|
||||
resources.insert("service".to_string(), vec!["vastai".to_string()]);
|
||||
let authority = CompositionAuthority {
|
||||
label: "agent-chat".to_string(),
|
||||
scopes: vec!["llm:call".to_string(), "fs:read".to_string()],
|
||||
resources,
|
||||
};
|
||||
let identity = authority.as_identity().expect("as_identity returns Some");
|
||||
assert_eq!(identity.id, "agent-chat");
|
||||
assert_eq!(
|
||||
identity.scopes,
|
||||
vec!["llm:call".to_string(), "fs:read".to_string()]
|
||||
);
|
||||
assert_eq!(
|
||||
identity.resources.get("service"),
|
||||
Some(&vec!["vastai".to_string()])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn composition_authority_new_populates_label_and_scopes() {
|
||||
let authority = CompositionAuthority::new(
|
||||
"agent-chat",
|
||||
["llm:call".to_string(), "fs:read".to_string()],
|
||||
);
|
||||
assert_eq!(authority.label, "agent-chat");
|
||||
assert_eq!(
|
||||
authority.scopes,
|
||||
vec!["llm:call".to_string(), "fs:read".to_string()]
|
||||
);
|
||||
assert!(authority.resources.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn composition_authority_none_is_none() {
|
||||
assert!(CompositionAuthority::none().is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn abort_policy_default_is_abort_dependents() {
|
||||
let policy = AbortPolicy::default();
|
||||
assert!(matches!(policy, AbortPolicy::AbortDependents));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generate_request_id_is_unique_and_non_deterministic() {
|
||||
let a = generate_request_id();
|
||||
let b = generate_request_id();
|
||||
assert_ne!(a, b);
|
||||
assert!(!a.is_empty());
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large.
Load diff
+1325
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,12 @@
|
||||
//! Operation registry: specs, handlers, access control, service discovery.
|
||||
//!
|
||||
//! Maps operation names to specs and handlers, enforces access control, and
|
||||
//! dispatches `call.requested` events to local handlers. The registry is
|
||||
//! layered by trust boundary (ADR-024): a curated layer (immutable after
|
||||
//! startup) plus dynamic session and connection overlays.
|
||||
|
||||
pub mod context;
|
||||
pub mod discovery;
|
||||
pub mod env;
|
||||
pub mod registration;
|
||||
pub mod spec;
|
||||
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,530 @@
|
||||
//! Operation specifications: `OperationSpec`, `OperationType`, `Visibility`,
|
||||
//! `ErrorDefinition`, and `AccessControl`.
|
||||
//!
|
||||
//! See `docs/architecture/crates/call/operation-registry.md` for the full
|
||||
//! specification.
|
||||
|
||||
use crate::core::auth::Identity;
|
||||
use crate::core::ownership::OwnershipProvider;
|
||||
use serde_json::Value;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum OperationType {
|
||||
Query,
|
||||
Mutation,
|
||||
Subscription,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum Visibility {
|
||||
External,
|
||||
Internal,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct ErrorDefinition {
|
||||
pub code: String,
|
||||
pub description: String,
|
||||
pub schema: Value,
|
||||
pub http_status: Option<u16>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default, PartialEq, Eq)]
|
||||
pub struct AccessControl {
|
||||
pub required_scopes: Vec<String>,
|
||||
pub required_scopes_any: Option<Vec<String>>,
|
||||
pub resource_type: Option<String>,
|
||||
pub resource_action: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum AccessResult {
|
||||
Allowed,
|
||||
Forbidden(String),
|
||||
}
|
||||
|
||||
impl AccessResult {
|
||||
pub fn is_allowed(&self) -> bool {
|
||||
matches!(self, AccessResult::Allowed)
|
||||
}
|
||||
}
|
||||
|
||||
impl AccessControl {
|
||||
pub fn has_restrictions(&self) -> bool {
|
||||
!self.required_scopes.is_empty()
|
||||
|| self.required_scopes_any.is_some()
|
||||
|| self.resource_type.is_some()
|
||||
|| self.resource_action.is_some()
|
||||
}
|
||||
|
||||
pub fn check(
|
||||
&self,
|
||||
identity: Option<&Identity>,
|
||||
resource_id: Option<&str>,
|
||||
ownership: Option<&dyn OwnershipProvider>,
|
||||
) -> AccessResult {
|
||||
if !self.has_restrictions() {
|
||||
return AccessResult::Allowed;
|
||||
}
|
||||
let identity = match identity {
|
||||
Some(id) => id,
|
||||
None => return AccessResult::Forbidden("authentication required".to_string()),
|
||||
};
|
||||
|
||||
for scope in &self.required_scopes {
|
||||
if !identity.scopes.iter().any(|s| s == scope) {
|
||||
return AccessResult::Forbidden(format!("missing required scope: {scope}"));
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(any) = &self.required_scopes_any {
|
||||
let has_one = any.iter().any(|s| identity.scopes.iter().any(|i| i == s));
|
||||
if !has_one {
|
||||
return AccessResult::Forbidden(
|
||||
"missing required scope (any of: ".to_string() + &any.join(", ") + ")",
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(p) = ownership {
|
||||
if let Some(rt) = &self.resource_type {
|
||||
match resource_id {
|
||||
Some(rid) => {
|
||||
let action = self.resource_action.as_deref().unwrap_or("");
|
||||
if !p.owns(identity, rt, rid, action) {
|
||||
return AccessResult::Forbidden(format!(
|
||||
"not owner of resource: {rt}/{rid}"
|
||||
));
|
||||
}
|
||||
}
|
||||
None => {
|
||||
if !p.owns_any(identity, rt) {
|
||||
return AccessResult::Forbidden(format!(
|
||||
"no owned resources of type: {rt}"
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
return AccessResult::Allowed;
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(rt) = &self.resource_type {
|
||||
let allowed = identity.resources.get(rt);
|
||||
match &self.resource_action {
|
||||
Some(action) => match allowed {
|
||||
Some(actions) if actions.iter().any(|a| a == action) => {}
|
||||
_ => {
|
||||
return AccessResult::Forbidden(format!("missing resource: {rt}/{action}"))
|
||||
}
|
||||
},
|
||||
None => match allowed {
|
||||
Some(actions) if !actions.is_empty() => {}
|
||||
_ => return AccessResult::Forbidden(format!("missing resource: {rt}")),
|
||||
},
|
||||
}
|
||||
} else if let Some(action) = &self.resource_action {
|
||||
let found = identity
|
||||
.resources
|
||||
.values()
|
||||
.any(|actions| actions.iter().any(|a| a == action));
|
||||
if !found {
|
||||
return AccessResult::Forbidden(format!("missing resource action: {action}"));
|
||||
}
|
||||
}
|
||||
|
||||
AccessResult::Allowed
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct OperationSpec {
|
||||
pub name: String,
|
||||
pub namespace: String,
|
||||
pub op_type: OperationType,
|
||||
pub visibility: Visibility,
|
||||
pub input_schema: Value,
|
||||
pub output_schema: Value,
|
||||
pub error_schemas: Vec<ErrorDefinition>,
|
||||
pub access_control: AccessControl,
|
||||
/// JSON pointer into the input for the resource ID, when
|
||||
/// `access_control.resource_type` is set and the operation targets a
|
||||
/// specific runtime-spawned resource (ADR-050). e.g. `"$.containerId"`
|
||||
/// for `docker/container/exec`. Absent for no-specific-resource
|
||||
/// operations (the `list` case). `None` for operations with no
|
||||
/// `resource_type` or with static resource sets.
|
||||
pub resource_id_path: Option<String>,
|
||||
}
|
||||
|
||||
impl OperationSpec {
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn new(
|
||||
name: impl Into<String>,
|
||||
op_type: OperationType,
|
||||
visibility: Visibility,
|
||||
input_schema: Value,
|
||||
output_schema: Value,
|
||||
error_schemas: Vec<ErrorDefinition>,
|
||||
access_control: AccessControl,
|
||||
resource_id_path: Option<String>,
|
||||
) -> Self {
|
||||
let name = name.into();
|
||||
let namespace = name
|
||||
.split('/')
|
||||
.next()
|
||||
.filter(|s| !s.is_empty())
|
||||
.unwrap_or("")
|
||||
.to_string();
|
||||
Self {
|
||||
name,
|
||||
namespace,
|
||||
op_type,
|
||||
visibility,
|
||||
input_schema,
|
||||
output_schema,
|
||||
error_schemas,
|
||||
access_control,
|
||||
resource_id_path,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn path(&self) -> String {
|
||||
format!("/{}", self.name)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::collections::HashMap;
|
||||
|
||||
fn identity(scopes: &[&str], resources: &[(&str, &[&str])]) -> Identity {
|
||||
let mut res = HashMap::new();
|
||||
for (k, v) in resources {
|
||||
res.insert(
|
||||
(*k).to_string(),
|
||||
v.iter().map(|s| (*s).to_string()).collect(),
|
||||
);
|
||||
}
|
||||
Identity {
|
||||
id: "caller".to_string(),
|
||||
scopes: scopes.iter().map(|s| (*s).to_string()).collect(),
|
||||
resources: res,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn path_has_leading_slash() {
|
||||
let spec = OperationSpec::new(
|
||||
"fs/readFile",
|
||||
OperationType::Query,
|
||||
Visibility::External,
|
||||
serde_json::json!({}),
|
||||
serde_json::json!({}),
|
||||
vec![],
|
||||
AccessControl::default(),
|
||||
None,
|
||||
);
|
||||
assert_eq!(spec.path(), "/fs/readFile");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn namespace_derived_from_name() {
|
||||
let spec = OperationSpec::new(
|
||||
"agent/chat",
|
||||
OperationType::Subscription,
|
||||
Visibility::External,
|
||||
serde_json::json!({}),
|
||||
serde_json::json!({}),
|
||||
vec![],
|
||||
AccessControl::default(),
|
||||
None,
|
||||
);
|
||||
assert_eq!(spec.namespace, "agent");
|
||||
assert_eq!(spec.name, "agent/chat");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn namespace_for_single_segment() {
|
||||
let spec = OperationSpec::new(
|
||||
"list",
|
||||
OperationType::Query,
|
||||
Visibility::Internal,
|
||||
serde_json::json!({}),
|
||||
serde_json::json!({}),
|
||||
vec![],
|
||||
AccessControl::default(),
|
||||
None,
|
||||
);
|
||||
assert_eq!(spec.namespace, "list");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resource_id_path_defaults_to_none() {
|
||||
let spec = OperationSpec::new(
|
||||
"fs/readFile",
|
||||
OperationType::Query,
|
||||
Visibility::External,
|
||||
serde_json::json!({}),
|
||||
serde_json::json!({}),
|
||||
vec![],
|
||||
AccessControl::default(),
|
||||
None,
|
||||
);
|
||||
assert_eq!(spec.resource_id_path, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn empty_access_control_allowed_for_all() {
|
||||
let acl = AccessControl::default();
|
||||
assert_eq!(acl.check(None, None, None), AccessResult::Allowed);
|
||||
let id = identity(&[], &[]);
|
||||
assert_eq!(acl.check(Some(&id), None, None), AccessResult::Allowed);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn none_identity_with_restrictions_forbidden() {
|
||||
let acl = AccessControl {
|
||||
required_scopes: vec!["read".to_string()],
|
||||
..Default::default()
|
||||
};
|
||||
assert_eq!(
|
||||
acl.check(None, None, None),
|
||||
AccessResult::Forbidden("authentication required".to_string())
|
||||
);
|
||||
|
||||
let acl2 = AccessControl {
|
||||
required_scopes_any: Some(vec!["read".to_string()]),
|
||||
..Default::default()
|
||||
};
|
||||
assert_eq!(
|
||||
acl2.check(None, None, None),
|
||||
AccessResult::Forbidden("authentication required".to_string())
|
||||
);
|
||||
|
||||
let acl3 = AccessControl {
|
||||
resource_type: Some("service".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
assert_eq!(
|
||||
acl3.check(None, None, None),
|
||||
AccessResult::Forbidden("authentication required".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn required_scopes_and_checked() {
|
||||
let acl = AccessControl {
|
||||
required_scopes: vec!["a".to_string(), "b".to_string()],
|
||||
..Default::default()
|
||||
};
|
||||
let id_missing = identity(&["a"], &[]);
|
||||
assert!(matches!(
|
||||
acl.check(Some(&id_missing), None, None),
|
||||
AccessResult::Forbidden(_)
|
||||
));
|
||||
let id_ok = identity(&["a", "b", "c"], &[]);
|
||||
assert_eq!(acl.check(Some(&id_ok), None, None), AccessResult::Allowed);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn required_scopes_any_or_checked() {
|
||||
let acl = AccessControl {
|
||||
required_scopes_any: Some(vec!["x".to_string(), "y".to_string()]),
|
||||
..Default::default()
|
||||
};
|
||||
let id_x = identity(&["x"], &[]);
|
||||
assert_eq!(acl.check(Some(&id_x), None, None), AccessResult::Allowed);
|
||||
let id_y = identity(&["y"], &[]);
|
||||
assert_eq!(acl.check(Some(&id_y), None, None), AccessResult::Allowed);
|
||||
let id_none = identity(&["z"], &[]);
|
||||
assert!(matches!(
|
||||
acl.check(Some(&id_none), None, None),
|
||||
AccessResult::Forbidden(_)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resource_check_with_type_and_action() {
|
||||
let acl = AccessControl {
|
||||
resource_type: Some("service".to_string()),
|
||||
resource_action: Some("read".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let id_ok = identity(&[], &[("service", &["read"])]);
|
||||
assert_eq!(acl.check(Some(&id_ok), None, None), AccessResult::Allowed);
|
||||
let id_missing_action = identity(&[], &[("service", &["write"])]);
|
||||
assert!(matches!(
|
||||
acl.check(Some(&id_missing_action), None, None),
|
||||
AccessResult::Forbidden(_)
|
||||
));
|
||||
let id_missing_type = identity(&[], &[("other", &["read"])]);
|
||||
assert!(matches!(
|
||||
acl.check(Some(&id_missing_type), None, None),
|
||||
AccessResult::Forbidden(_)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn combined_scopes_and_resources() {
|
||||
let acl = AccessControl {
|
||||
required_scopes: vec!["admin".to_string()],
|
||||
resource_type: Some("service".to_string()),
|
||||
resource_action: Some("read".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let id_ok = identity(&["admin"], &[("service", &["read"])]);
|
||||
assert_eq!(acl.check(Some(&id_ok), None, None), AccessResult::Allowed);
|
||||
let id_missing_scope = identity(&["user"], &[("service", &["read"])]);
|
||||
assert!(matches!(
|
||||
acl.check(Some(&id_missing_scope), None, None),
|
||||
AccessResult::Forbidden(_)
|
||||
));
|
||||
}
|
||||
|
||||
struct MockOwnership {
|
||||
owned: Vec<(String, String)>,
|
||||
}
|
||||
|
||||
impl OwnershipProvider for MockOwnership {
|
||||
fn owns(
|
||||
&self,
|
||||
_identity: &Identity,
|
||||
resource_type: &str,
|
||||
resource_id: &str,
|
||||
_action: &str,
|
||||
) -> bool {
|
||||
self.owned
|
||||
.iter()
|
||||
.any(|(rt, rid)| rt == resource_type && rid == resource_id)
|
||||
}
|
||||
|
||||
fn owned_resources(&self, _identity: &Identity, resource_type: &str) -> Vec<String> {
|
||||
self.owned
|
||||
.iter()
|
||||
.filter(|(rt, _)| rt == resource_type)
|
||||
.map(|(_, rid)| rid.clone())
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn owns_any(&self, _identity: &Identity, resource_type: &str) -> bool {
|
||||
self.owned.iter().any(|(rt, _)| rt == resource_type)
|
||||
}
|
||||
}
|
||||
|
||||
fn empty_identity(id: &str) -> Identity {
|
||||
Identity {
|
||||
id: id.to_string(),
|
||||
scopes: vec![],
|
||||
resources: HashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ownership_provider_allows_owned_resource() {
|
||||
let acl = AccessControl {
|
||||
resource_type: Some("container".to_string()),
|
||||
resource_action: Some("exec".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let id = empty_identity("alice");
|
||||
let provider = MockOwnership {
|
||||
owned: vec![("container".to_string(), "c1".to_string())],
|
||||
};
|
||||
assert_eq!(
|
||||
acl.check(
|
||||
Some(&id),
|
||||
Some("c1"),
|
||||
Some(&provider as &dyn OwnershipProvider)
|
||||
),
|
||||
AccessResult::Allowed
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ownership_provider_forbids_unowned_resource() {
|
||||
let acl = AccessControl {
|
||||
resource_type: Some("container".to_string()),
|
||||
resource_action: Some("exec".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let id = empty_identity("alice");
|
||||
let provider = MockOwnership {
|
||||
owned: vec![("container".to_string(), "c1".to_string())],
|
||||
};
|
||||
assert!(matches!(
|
||||
acl.check(
|
||||
Some(&id),
|
||||
Some("c2"),
|
||||
Some(&provider as &dyn OwnershipProvider)
|
||||
),
|
||||
AccessResult::Forbidden(_)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ownership_provider_forbids_none_identity() {
|
||||
let acl = AccessControl {
|
||||
resource_type: Some("container".to_string()),
|
||||
resource_action: Some("exec".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let provider = MockOwnership {
|
||||
owned: vec![("container".to_string(), "c1".to_string())],
|
||||
};
|
||||
assert!(matches!(
|
||||
acl.check(None, Some("c1"), Some(&provider as &dyn OwnershipProvider)),
|
||||
AccessResult::Forbidden(_)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ownership_provider_list_allowed_when_owns_any() {
|
||||
let acl = AccessControl {
|
||||
resource_type: Some("container".to_string()),
|
||||
resource_action: Some("exec".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let id = empty_identity("alice");
|
||||
let provider = MockOwnership {
|
||||
owned: vec![("container".to_string(), "c1".to_string())],
|
||||
};
|
||||
assert_eq!(
|
||||
acl.check(Some(&id), None, Some(&provider as &dyn OwnershipProvider)),
|
||||
AccessResult::Allowed
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ownership_provider_list_forbidden_when_not_owns_any() {
|
||||
let acl = AccessControl {
|
||||
resource_type: Some("container".to_string()),
|
||||
resource_action: Some("exec".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let id = empty_identity("alice");
|
||||
let provider = MockOwnership {
|
||||
owned: vec![("volume".to_string(), "v1".to_string())],
|
||||
};
|
||||
assert!(matches!(
|
||||
acl.check(Some(&id), None, Some(&provider as &dyn OwnershipProvider)),
|
||||
AccessResult::Forbidden(_)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ownership_none_falls_back_to_static() {
|
||||
let acl = AccessControl {
|
||||
resource_type: Some("service".to_string()),
|
||||
resource_action: Some("read".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let id_ok = identity(&[], &[("service", &["read"])]);
|
||||
assert_eq!(acl.check(Some(&id_ok), None, None), AccessResult::Allowed);
|
||||
let id_missing = identity(&[], &[("service", &["write"])]);
|
||||
assert!(matches!(
|
||||
acl.check(Some(&id_missing), None, None),
|
||||
AccessResult::Forbidden(_)
|
||||
));
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user