use std::sync::Arc;
use alkcall::channels::client::ChannelClient;
use alkcall::client::{
from_call as import_from_call, AdapterError, FromCallConfig, OperationAdapter,
};
use alkcall::core::types::{Connection, Secret};
use alkcall::protocol::connection::CallConnection;
use alkcall::protocol::wire::CallError;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use crate::websocket::split_tungstenite_to_bytes;
const PENDING_SWEEP_INTERVAL: std::time::Duration = std::time::Duration::from_millis(50);
const PENDING_SWEEP_MAX_POST_EOF: u32 = 8;
fn connection_closed_error() -> CallError {
CallError::connection_closed("from_wss connection dropped")
}
pub struct FromWss {
endpoint: String,
auth_token: Option<Secret<String>>,
namespace: Option<String>,
allow_plaintext: bool,
}
impl FromWss {
pub fn new(endpoint: impl Into<String>) -> Self {
Self {
endpoint: endpoint.into(),
auth_token: None,
namespace: None,
allow_plaintext: false,
}
}
pub fn with_auth_token(mut self, token: impl Into<String>) -> Self {
self.auth_token = Some(Secret::new(token.into()));
self
}
pub fn with_namespace(mut self, namespace: impl Into<String>) -> Self {
self.namespace = Some(namespace.into());
self
}
pub fn allow_plaintext(mut self) -> Self {
self.allow_plaintext = true;
self
}
pub fn endpoint(&self) -> &str {
&self.endpoint
}
pub fn namespace(&self) -> Option<&str> {
self.namespace.as_deref()
}
pub fn auth_token(&self) -> Option<&Secret<String>> {
self.auth_token.as_ref()
}
}
pub struct WssSession {
_client: ChannelClient,
pub call_connection: Arc<CallConnection>,
_monitor: WssDropMonitor,
}
impl WssSession {
#[cfg(all(test, feature = "server"))]
fn monitor_handle(&mut self) -> &mut tokio::task::JoinHandle<()> {
&mut self._monitor.monitor_task
}
}
#[cfg_attr(not(test), allow(dead_code))]
struct WssDropMonitor {
close_tx: Option<tokio::sync::oneshot::Sender<()>>,
#[cfg(all(test, feature = "server"))]
monitor_task: tokio::task::JoinHandle<()>,
}
impl Drop for WssDropMonitor {
fn drop(&mut self) {
if let Some(tx) = self.close_tx.take() {
let _ = tx.send(());
}
}
}
impl WssSession {
pub async fn connect(
endpoint: &str,
auth_token: Option<&str>,
allow_plaintext: bool,
) -> Result<Self, AdapterError> {
if let Some(_token) = auth_token {
if !allow_plaintext {
let scheme = endpoint.split("://").next().unwrap_or_default();
if scheme.eq_ignore_ascii_case("ws") {
return Err(AdapterError::Transport {
message: format!(
"refusing plaintext `ws://` endpoint `{endpoint}` with a Bearer token: \
the credential would ride an unencrypted connection (CON-03); \
use `wss://` or call `FromWss::allow_plaintext` explicitly"
),
});
}
}
}
let mut request = endpoint
.into_client_request()
.map_err(|e| AdapterError::Transport {
message: format!("invalid WSS endpoint `{endpoint}`: {e}"),
})?;
if let Some(token) = auth_token {
let value = format!("Bearer {token}");
request.headers_mut().insert(
"Authorization",
value.parse().map_err(|_| AdapterError::Transport {
message: "bearer token is not a valid header value".to_string(),
})?,
);
}
let (ws, _response) = tokio_tungstenite::connect_async_with_config(
request,
Some(
tokio_tungstenite::tungstenite::protocol::WebSocketConfig::default()
.max_message_size(Some(crate::websocket::INBOUND_WS_MESSAGE_CAP))
.max_frame_size(Some(crate::websocket::INBOUND_WS_FRAME_CAP)),
),
false,
)
.await
.map_err(|e| AdapterError::Transport {
message: format!("WSS connect failed: {e}"),
})?;
let (byte_stream, pumps) = split_tungstenite_to_bytes(ws);
let conn = Connection::from_bidi(byte_stream, b"alk/channels".to_vec(), None);
let client =
ChannelClient::from_connection(conn)
.await
.map_err(|e| AdapterError::Transport {
message: format!("channels session setup failed: {e}"),
})?;
let call_connection =
client
.take_call_connection()
.await
.ok_or_else(|| AdapterError::Transport {
message: "channel client closed before channel 0 was installed".to_string(),
})?;
let (close_tx, mut close_rx) = tokio::sync::oneshot::channel();
let pending = Arc::clone(call_connection.pending());
#[cfg_attr(not(all(test, feature = "server")), allow(unused_variables))]
let monitor_task = tokio::spawn(async move {
let mut eof_rx = pumps.read_eof();
let mut sweep = tokio::time::interval(PENDING_SWEEP_INTERVAL);
sweep.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
let eof_observed = *eof_rx.borrow_and_update();
if eof_observed {
pending.lock().fail_all(connection_closed_error());
}
let mut idle_sweeps_post_eof: u32 = 0;
loop {
tokio::select! {
_ = &mut close_rx => {
pending.lock().fail_all(connection_closed_error());
return;
}
_ = sweep.tick() => {
if *eof_rx.borrow() {
if pending
.lock()
.fail_all(connection_closed_error())
.is_empty()
{
idle_sweeps_post_eof += 1;
} else {
idle_sweeps_post_eof = 0;
}
if idle_sweeps_post_eof >= PENDING_SWEEP_MAX_POST_EOF {
return;
}
}
}
}
}
});
Ok(Self {
_client: client,
call_connection,
_monitor: WssDropMonitor {
close_tx: Some(close_tx),
#[cfg(all(test, feature = "server"))]
monitor_task,
},
})
}
}
fn is_protocol_session_op(name: &str) -> bool {
name == "services/list"
|| name == "services/schema"
|| name == "services/list-peers"
|| name == alkcall::registry::op_register::OP_REGISTER_NAME
|| name == alkcall::channels::operations::OP_CHANNEL_CLOSE
|| name == alkcall::channels::operations::OP_CHANNEL_CONTROL
|| name == alkcall::channels::operations::OP_CHANNEL_RESOURCES_SUBSCRIBE
}
#[async_trait::async_trait]
impl OperationAdapter for FromWss {
async fn import(
&self,
) -> Result<Vec<alkcall::registry::registration::HandlerRegistration>, AdapterError> {
let session = WssSession::connect(
&self.endpoint,
self.auth_token.as_ref().map(|s| s.expose_secret().as_str()),
self.allow_plaintext,
)
.await?;
let config = match &self.namespace {
Some(ns) => FromCallConfig::new().with_namespace_prefix(ns),
None => FromCallConfig::new(),
};
let config = config
.with_operation_filter(protocol_session_ops_from(&session.call_connection).await?);
let bundles = import_from_call(&session.call_connection, config).await;
std::mem::forget(session);
bundles
}
}
async fn protocol_session_ops_from(
connection: &CallConnection,
) -> Result<std::collections::HashSet<String>, AdapterError> {
let response = connection
.call("services/list", serde_json::json!({}))
.await;
let output = response.result.map_err(|e| AdapterError::DiscoveryFailed {
message: format!("services/list failed: {} ({})", e.code, e.message),
})?;
let ops = output
.get("operations")
.and_then(|v| v.as_array())
.ok_or_else(|| AdapterError::SchemaParse {
message: "services/list response missing 'operations' array".to_string(),
})?;
let mut filter = std::collections::HashSet::new();
for op in ops {
if let Some(name) = op.get("name").and_then(|v| v.as_str()) {
if !is_protocol_session_op(name) {
filter.insert(name.to_string());
}
}
}
Ok(filter)
}
#[cfg(all(test, feature = "server"))]
mod tests {
use super::*;
use alkcall::core::auth::{Identity, IdentityProvider};
use alkcall::protocol::wire::ResponseEnvelope;
use alkcall::registry::discovery::{
services_list_handler, services_list_spec, services_schema_handler, services_schema_spec,
};
use alkcall::registry::registration::{
make_handler, HandlerKind, HandlerRegistration, OperationProvenance, OperationRegistry,
};
use alkcall::registry::spec::{AccessControl, OperationSpec, OperationType, Visibility};
use std::collections::HashMap;
use std::sync::Mutex as StdMutex;
type WsPumpsSlot = std::sync::Arc<std::sync::Mutex<Option<crate::websocket::WsPumps>>>;
fn drop_on_signal_producer(registry: Arc<OperationRegistry>) -> (String, WsPumpsSlot) {
let pumps_slot: WsPumpsSlot = std::sync::Arc::new(std::sync::Mutex::new(None));
let provider = provider_with(vec![("tok-1", identity("alice", &[]))]);
async fn killable_upgrade(
axum::extract::State(state): axum::extract::State<(
Arc<OperationRegistry>,
WsPumpsSlot,
)>,
axum::Extension(identity): axum::Extension<Identity>,
ws_upgrade: axum::extract::ws::WebSocketUpgrade,
) -> axum::response::Response {
ws_upgrade.on_upgrade(move |socket| async move {
let (byte_stream, pumps) = crate::websocket::split_ws_to_bytes(socket);
*state.1.lock().unwrap_or_else(|e| e.into_inner()) = Some(pumps);
let conn = alkcall::core::types::Connection::from_bidi(
byte_stream,
b"alk/channels".to_vec(),
None,
);
let _ = conn.set_identity(identity.clone());
let adapter = alkcall::channels::adapter::ChannelsAdapter::new(
crate::websocket::adapter_install_channel_zero(Arc::clone(&state.0)),
std::sync::Arc::new(alkcall::channels::policy::NoCap),
);
let auth = alkcall::core::auth::AuthContext {
identity: Some(identity),
alpn: b"alk/channels".to_vec(),
remote_addr: None,
tls_client_fingerprint: None,
};
if let Err(e) =
alkcall::core::types::ProtocolHandler::handle(&adapter, conn, &auth).await
{
tracing::warn!(error = %e, "kill-test channels session ended");
}
})
}
let app = axum::Router::new()
.route(
"/alk/channels",
axum::routing::get(killable_upgrade).route_layer(
axum::middleware::from_fn_with_state(
Arc::clone(&provider),
crate::websocket::ws_bearer_auth,
),
),
)
.with_state((registry, Arc::clone(&pumps_slot)));
let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("bind");
let addr = listener.local_addr().expect("addr");
let _ = listener.set_nonblocking(true);
std::thread::spawn(move || {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("rt");
rt.block_on(async {
let listener = tokio::net::TcpListener::from_std(listener).expect("tokio listener");
let _ = axum::serve(listener, app).await;
});
});
(format!("ws://{addr}/alk/channels"), pumps_slot)
}
fn identity(id: &str, scopes: &[&str]) -> Identity {
Identity {
id: id.to_string(),
scopes: scopes.iter().map(|s| s.to_string()).collect(),
resources: HashMap::new(),
}
}
struct StaticTokens {
tokens: StdMutex<HashMap<String, Identity>>,
}
impl IdentityProvider for StaticTokens {
fn resolve_from_fingerprint(&self, _: &str) -> Option<Identity> {
None
}
fn resolve_from_token(&self, token: &alkcall::core::auth::AuthToken) -> Option<Identity> {
let s = String::from_utf8_lossy(&token.raw).to_string();
self.tokens
.lock()
.unwrap_or_else(|e| e.into_inner())
.get(&s)
.cloned()
}
}
fn provider_with(tokens: Vec<(&str, Identity)>) -> Arc<dyn IdentityProvider> {
let map: HashMap<String, Identity> = tokens
.into_iter()
.map(|(t, i)| (t.to_string(), i))
.collect();
Arc::new(StaticTokens {
tokens: StdMutex::new(map),
})
}
fn echo_handler() -> alkcall::registry::registration::Handler {
make_handler(|input, ctx| async move { ResponseEnvelope::ok(ctx.request_id, input) })
}
fn noop_context(request_id: &str) -> alkcall::registry::context::OperationContext {
struct NoopEnv;
#[async_trait::async_trait]
impl alkcall::registry::env::OperationEnv for NoopEnv {
async fn invoke_with_policy(
&self,
_ns: &str,
_op: &str,
_input: serde_json::Value,
parent: &alkcall::registry::context::OperationContext,
_policy: alkcall::registry::context::AbortPolicy,
) -> ResponseEnvelope {
ResponseEnvelope::ok(parent.request_id.clone(), serde_json::Value::Null)
}
fn contains(&self, _name: &str) -> bool {
false
}
}
alkcall::registry::context::OperationContext {
request_id: request_id.to_string(),
parent_request_id: None,
identity: None,
handler_identity: None,
forwarded_for: None,
capabilities: alkcall::core::types::Capabilities::new(),
metadata: HashMap::new(),
scoped_env: alkcall::registry::context::ScopedPeerEnv::empty(),
env: Arc::new(NoopEnv),
abort_policy: alkcall::registry::context::AbortPolicy::default(),
deadline: Some(std::time::Instant::now() + std::time::Duration::from_secs(30)),
internal: true,
ownership: None,
}
}
fn producer_registry() -> Arc<OperationRegistry> {
let inner = OperationRegistry::new();
inner
.register(HandlerRegistration::new(
OperationSpec::new(
"echo/run",
OperationType::Query,
Visibility::External,
serde_json::json!({}),
serde_json::json!({}),
vec![],
AccessControl::default(),
None,
),
HandlerKind::Once(echo_handler()),
OperationProvenance::Local,
None,
None,
alkcall::core::types::Capabilities::new(),
))
.unwrap();
inner
.register(HandlerRegistration::new(
OperationSpec::new(
"admin/run",
OperationType::Query,
Visibility::External,
serde_json::json!({}),
serde_json::json!({}),
vec![],
AccessControl {
required_scopes: vec!["admin".to_string()],
..Default::default()
},
None,
),
HandlerKind::Once(echo_handler()),
OperationProvenance::Local,
None,
None,
alkcall::core::types::Capabilities::new(),
))
.unwrap();
let inner = Arc::new(inner);
let registry = OperationRegistry::new();
registry
.register(HandlerRegistration::new(
services_list_spec(),
HandlerKind::Once(services_list_handler(Arc::clone(&inner))),
OperationProvenance::Local,
None,
None,
alkcall::core::types::Capabilities::new(),
))
.unwrap();
registry
.register(HandlerRegistration::new(
services_schema_spec(),
HandlerKind::Once(services_schema_handler(Arc::clone(&inner))),
OperationProvenance::Local,
None,
None,
alkcall::core::types::Capabilities::new(),
))
.unwrap();
for spec in inner.list_operations() {
let name = spec.name.clone();
let reg = inner.registration(&name).unwrap();
registry
.register(HandlerRegistration::new(
reg.spec.clone(),
reg.handler.clone(),
reg.provenance,
reg.composition_authority.clone(),
reg.scoped_env.clone(),
reg.capabilities.clone(),
))
.unwrap();
}
Arc::new(registry)
}
fn slow_producer_registry() -> Arc<OperationRegistry> {
let inner = OperationRegistry::new();
inner
.register(HandlerRegistration::new(
OperationSpec::new(
"slow/op",
OperationType::Query,
Visibility::External,
serde_json::json!({}),
serde_json::json!({}),
vec![],
AccessControl::default(),
None,
),
HandlerKind::Once(make_handler(|_input, _ctx| async move {
tokio::time::sleep(std::time::Duration::from_secs(30)).await;
ResponseEnvelope::ok("never", serde_json::json!({}))
})),
OperationProvenance::Local,
None,
None,
alkcall::core::types::Capabilities::new(),
))
.unwrap();
let inner = Arc::new(inner);
let registry = OperationRegistry::new();
registry
.register(HandlerRegistration::new(
services_list_spec(),
HandlerKind::Once(services_list_handler(Arc::clone(&inner))),
OperationProvenance::Local,
None,
None,
alkcall::core::types::Capabilities::new(),
))
.unwrap();
registry
.register(HandlerRegistration::new(
services_schema_spec(),
HandlerKind::Once(services_schema_handler(Arc::clone(&inner))),
OperationProvenance::Local,
None,
None,
alkcall::core::types::Capabilities::new(),
))
.unwrap();
for spec in inner.list_operations() {
let name = spec.name.clone();
let reg = inner.registration(&name).unwrap();
registry
.register(HandlerRegistration::new(
reg.spec.clone(),
reg.handler.clone(),
reg.provenance,
reg.composition_authority.clone(),
reg.scoped_env.clone(),
reg.capabilities.clone(),
))
.unwrap();
}
Arc::new(registry)
}
async fn spawn_producer(
registry: Arc<OperationRegistry>,
provider: Arc<dyn IdentityProvider>,
) -> String {
let app = axum::Router::new()
.route(
"/alk/channels",
axum::routing::get(crate::websocket::ws_upgrade_handler),
)
.layer(axum::middleware::from_fn_with_state(
provider,
crate::websocket::ws_bearer_auth,
))
.with_state(registry);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = format!("ws://{}", listener.local_addr().unwrap());
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
format!("{addr}/alk/channels")
}
#[test]
fn struct_holds_endpoint_token_namespace() {
let adapter = FromWss::new("ws://localhost:9000/alk/channels");
assert_eq!(adapter.endpoint(), "ws://localhost:9000/alk/channels");
assert_eq!(adapter.namespace(), None);
assert!(adapter.auth_token().is_none());
let with_all = adapter.with_auth_token("tok").with_namespace("remote");
assert_eq!(
with_all.auth_token().map(|s| s.expose_secret().as_str()),
Some("tok")
);
assert_eq!(with_all.namespace(), Some("remote"));
assert!(!with_all.allow_plaintext, "ws:// refused by default");
let opt_in = with_all.allow_plaintext();
assert!(opt_in.allow_plaintext);
}
#[test]
fn token_debug_output_is_redacted() {
let adapter = FromWss::new("wss://localhost/alk/channels").with_auth_token("sekrit");
let debug = format!("{:?}", adapter.auth_token().expect("token present"));
assert_eq!(debug, "[REDACTED]");
assert!(!debug.contains("sekrit"));
}
#[tokio::test]
async fn plaintext_ws_with_token_refused_by_default() {
let endpoint = spawn_producer(
producer_registry(),
provider_with(vec![("tok-1", identity("alice", &[]))]),
)
.await;
let adapter = FromWss::new(&endpoint).with_auth_token("tok-1");
match adapter.import().await {
Ok(_) => panic!("ws:// + token must be refused without an explicit opt-in (CON-03)"),
Err(AdapterError::Transport { message }) => {
assert!(
message.contains("CON-03"),
"error explains the refusal: {message}"
);
assert!(message.contains("ws://"));
}
Err(other) => panic!("expected Transport error, got {other}"),
}
}
#[tokio::test]
async fn plaintext_ws_with_token_allowed_when_explicitly_enabled() {
let endpoint = spawn_producer(
producer_registry(),
provider_with(vec![("tok-1", identity("alice", &["user", "admin"]))]),
)
.await;
let adapter = FromWss::new(&endpoint)
.with_auth_token("tok-1")
.allow_plaintext();
let bundles = adapter.import().await.expect("explicit opt-in dials ws://");
assert!(!bundles.is_empty());
}
#[tokio::test]
async fn plaintext_ws_without_token_is_not_refused_by_the_adapter() {
let endpoint = spawn_producer(producer_registry(), provider_with(vec![])).await;
let adapter = FromWss::new(&endpoint);
match adapter.import().await {
Ok(bundles) => assert!(!bundles.is_empty()),
Err(AdapterError::Transport { message }) => {
assert!(
!message.contains("CON-03"),
"without a token the adapter must not refuse ws://, got: {message}"
);
}
Err(other) => panic!("expected Transport error, got {other}"),
}
}
#[tokio::test]
async fn import_discovers_ops_and_builds_forwarding_handlers() {
let endpoint = spawn_producer(
producer_registry(),
provider_with(vec![("tok-1", identity("alice", &["user", "admin"]))]),
)
.await;
let adapter = FromWss::new(&endpoint)
.with_auth_token("tok-1")
.allow_plaintext();
let bundles = adapter.import().await.expect("import succeeds");
let mut names: Vec<&str> = bundles.iter().map(|b| b.spec.name.as_str()).collect();
names.sort();
assert_eq!(names, vec!["admin/run", "echo/run"]);
for b in &bundles {
assert_eq!(b.provenance, OperationProvenance::FromCall);
assert!(b.composition_authority.is_none());
assert!(b.scoped_env.is_none());
}
}
#[tokio::test]
async fn imported_ops_invoke_end_to_end() {
let endpoint = spawn_producer(
producer_registry(),
provider_with(vec![("tok-1", identity("alice", &["user", "admin"]))]),
)
.await;
let adapter = FromWss::new(&endpoint)
.with_auth_token("tok-1")
.allow_plaintext();
let bundles = adapter.import().await.expect("import succeeds");
let echo = bundles
.into_iter()
.find(|b| b.spec.name == "echo/run")
.expect("echo/run present");
let ctx = noop_context("req-e2e");
let response = match &echo.handler {
HandlerKind::Once(h) => h(serde_json::json!({ "hello": "world" }), ctx).await,
HandlerKind::Stream(_) | HandlerKind::Sink(_) => {
panic!("expected Once handler for query op")
}
};
assert_eq!(response.request_id, "req-e2e");
assert_eq!(response.result, Ok(serde_json::json!({ "hello": "world" })));
}
#[tokio::test]
async fn acl_enforced_end_to_end() {
let endpoint = spawn_producer(
producer_registry(),
provider_with(vec![("tok-alice", identity("alice", &["user"]))]),
)
.await;
let adapter = FromWss::new(&endpoint)
.with_auth_token("tok-alice")
.allow_plaintext();
let bundles = adapter.import().await.expect("import succeeds");
let names: Vec<&str> = bundles.iter().map(|b| b.spec.name.as_str()).collect();
assert_eq!(names, vec!["echo/run"], "ACL-filtered op not discovered");
}
#[tokio::test]
async fn namespace_prefix_applies_to_imported_names() {
let endpoint = spawn_producer(
producer_registry(),
provider_with(vec![("tok-1", identity("alice", &["user", "admin"]))]),
)
.await;
let adapter = FromWss::new(&endpoint)
.with_auth_token("tok-1")
.with_namespace("remote")
.allow_plaintext();
let bundles = adapter.import().await.expect("import succeeds");
let mut names: Vec<&str> = bundles.iter().map(|b| b.spec.name.as_str()).collect();
names.sort();
assert_eq!(names, vec!["remote/admin/run", "remote/echo/run"]);
}
#[tokio::test]
async fn connection_drop_fails_in_flight_calls_retryable_no_hang() {
let endpoint = spawn_producer(
slow_producer_registry(),
provider_with(vec![("tok-1", identity("alice", &[]))]),
)
.await;
let session = WssSession::connect(&endpoint, Some("tok-1"), true)
.await
.expect("connect");
let config = FromCallConfig::new();
let bundles = import_from_call(&session.call_connection, config)
.await
.expect("import");
let slow = bundles
.into_iter()
.find(|b| b.spec.name == "slow/op")
.expect("slow/op present");
let ctx = noop_context("req-drop");
let handler = match &slow.handler {
HandlerKind::Once(h) => h.clone(),
_ => panic!("expected Once handler"),
};
let call_task = tokio::spawn(async move { handler(serde_json::json!({}), ctx).await });
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
drop(session);
let response = tokio::time::timeout(std::time::Duration::from_secs(5), call_task)
.await
.expect("call resolves after drop (no hang until the 30s deadline)")
.expect("join");
match response.result {
Err(e) => {
assert!(e.retryable, "drop error must be retryable, got {e:?}");
}
Ok(_) => panic!("expected Err after connection drop"),
}
}
#[tokio::test]
async fn unreachable_endpoint_returns_transport_error() {
let adapter = FromWss::new("ws://127.0.0.1:1/alk/channels");
match adapter.import().await {
Ok(_) => panic!("expected Err for unreachable endpoint"),
Err(AdapterError::Transport { .. }) => {}
Err(other) => panic!("expected Transport, got {other}"),
}
}
#[tokio::test]
async fn invalid_endpoint_uri_returns_transport_error() {
let adapter = FromWss::new("not a url");
match adapter.import().await {
Ok(_) => panic!("expected Err for invalid endpoint"),
Err(AdapterError::Transport { .. }) => {}
Err(other) => panic!("expected Transport, got {other}"),
}
}
#[test]
fn no_env_vars_used_for_credentials() {
std::env::set_var("WSS_TOKEN", "should-not-be-used");
let adapter = FromWss::new("ws://localhost/alk/channels");
assert!(adapter.auth_token().is_none());
std::env::remove_var("WSS_TOKEN");
}
async fn abort_pumps_and_wait(slot: &WsPumpsSlot) {
let deadline = std::time::Instant::now() + std::time::Duration::from_secs(5);
loop {
let pumps = slot.lock().unwrap_or_else(|e| e.into_inner()).take();
match pumps {
Some(pumps) => {
pumps.abort();
break;
}
None => {
if std::time::Instant::now() > deadline {
panic!("producer pumps never registered");
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
}
}
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
}
async fn race_call_resolves_retryable(drop_before_call: bool) {
let (endpoint, pumps_slot) = drop_on_signal_producer(slow_producer_registry());
let session = WssSession::connect(&endpoint, Some("tok-1"), true)
.await
.expect("connect");
let config = FromCallConfig::new();
let bundles = import_from_call(&session.call_connection, config)
.await
.expect("import");
let slow = bundles
.into_iter()
.find(|b| b.spec.name == "slow/op")
.expect("slow/op present");
let handler = match &slow.handler {
HandlerKind::Once(h) => h.clone(),
_ => panic!("expected Once handler"),
};
if drop_before_call {
std::mem::forget(session);
} else {
drop(session);
}
abort_pumps_and_wait(&pumps_slot).await;
let ctx = noop_context("req-race");
let call_task = tokio::spawn(async move { handler(serde_json::json!({}), ctx).await });
let response = tokio::time::timeout(std::time::Duration::from_secs(5), call_task)
.await
.expect("call resolves after racing drop (no hang)")
.expect("join");
match response.result {
Err(e) => {
assert!(
e.retryable,
"drop error must be retryable regardless of whether the call resolved via \
the drop monitor (CONNECTION_CLOSED) or the CF-001 write-failure mapping \
(also CONNECTION_CLOSED post-CF-001), got {e:?}"
);
assert_eq!(
e.code, "CONNECTION_CLOSED",
"both resolution paths share the retryable wire code post-CF-001: {e:?}"
);
}
Ok(_) => panic!("expected Err after connection drop"),
}
}
#[tokio::test]
async fn forget_session_drop_during_call_registration_resolves_retryable_no_hang() {
race_call_resolves_retryable(true).await;
}
#[tokio::test]
async fn held_session_drop_during_call_registration_resolves_retryable_no_hang() {
race_call_resolves_retryable(false).await;
}
#[tokio::test]
async fn call_registered_after_eof_fails_fast_retryable() {
let (endpoint, pumps_slot) = drop_on_signal_producer(slow_producer_registry());
let session = WssSession::connect(&endpoint, Some("tok-1"), true)
.await
.expect("connect");
let config = FromCallConfig::new();
let bundles = import_from_call(&session.call_connection, config)
.await
.expect("import");
let slow = bundles
.into_iter()
.find(|b| b.spec.name == "slow/op")
.expect("slow/op present");
let handler = match &slow.handler {
HandlerKind::Once(h) => h.clone(),
_ => panic!("expected Once handler"),
};
std::mem::forget(session);
abort_pumps_and_wait(&pumps_slot).await;
let ctx = noop_context("req-late");
let call_task = tokio::spawn(async move { handler(serde_json::json!({}), ctx).await });
let response = tokio::time::timeout(std::time::Duration::from_millis(500), call_task)
.await
.expect("post-EOF registered call fails fast (no 1s sweep wait, no hang)")
.expect("join");
match response.result {
Err(e) => {
assert_eq!(e.code, "CONNECTION_CLOSED", "fast-fail code, got {e:?}");
assert!(e.retryable, "drop error must be retryable, got {e:?}");
}
Ok(_) => panic!("expected Err after connection drop"),
}
}
#[tokio::test]
async fn drop_monitor_ends_after_eof_plus_grace_window() {
let (endpoint, pumps_slot) = drop_on_signal_producer(slow_producer_registry());
let mut session = WssSession::connect(&endpoint, Some("tok-1"), true)
.await
.expect("connect");
let config = FromCallConfig::new();
let bundles = import_from_call(&session.call_connection, config)
.await
.expect("import");
assert!(!bundles.is_empty());
abort_pumps_and_wait(&pumps_slot).await;
tokio::time::timeout(std::time::Duration::from_secs(5), session.monitor_handle())
.await
.expect("monitor ends after EOF + bounded grace window (no permanent task)")
.expect("monitor task was not aborted, no panic payload");
}
}