use super::dispatch::build_hub_ext_for_vctx;
use super::*;
use crate::hub::InMemoryQueue;
use crate::ilink::UpstreamClient;
use crate::store::Store;
use std::sync::atomic::Ordering;
use std::sync::Arc;
pub(crate) struct MockUpstream {
send_ret: i32,
send_msg: &'static str,
send_calls: std::sync::atomic::AtomicU64,
}
impl MockUpstream {
pub(crate) fn returning_ok() -> Arc<dyn crate::ilink::UpstreamSink> {
Arc::new(Self {
send_ret: 0,
send_msg: "",
send_calls: std::sync::atomic::AtomicU64::new(0),
})
}
pub(crate) fn returning_err(
ret: i32,
msg: &'static str,
) -> Arc<dyn crate::ilink::UpstreamSink> {
Arc::new(Self {
send_ret: ret,
send_msg: msg,
send_calls: std::sync::atomic::AtomicU64::new(0),
})
}
}
#[async_trait::async_trait]
impl crate::ilink::UpstreamSink for MockUpstream {
async fn notify_start(&self) -> anyhow::Result<()> {
Ok(())
}
async fn send_message(
&self,
_req: crate::ilink::types::SendMessageRequest,
) -> anyhow::Result<crate::ilink::types::SendMessageResponse> {
self.send_calls.fetch_add(1, Ordering::Relaxed);
if self.send_ret == 0 {
Ok(crate::ilink::types::SendMessageResponse::ok())
} else {
Ok(crate::ilink::types::SendMessageResponse::err(
self.send_ret,
self.send_msg,
))
}
}
async fn send_typing(
&self,
_req: crate::ilink::types::SendTypingRequest,
) -> anyhow::Result<()> {
Ok(())
}
async fn get_config(
&self,
_req: crate::ilink::types::GetConfigRequest,
) -> anyhow::Result<crate::ilink::types::GetConfigResponse> {
Ok(crate::ilink::types::GetConfigResponse::default())
}
async fn get_upload_url(
&self,
_req: crate::ilink::types::GetUploadUrlRequest,
) -> anyhow::Result<crate::ilink::types::GetUploadUrlResponse> {
Ok(crate::ilink::types::GetUploadUrlResponse {
ret: 0,
upload_url: None,
media_id: None,
errmsg: None,
})
}
fn polls_ok(&self) -> u64 {
self.send_calls.load(Ordering::Relaxed)
}
fn polls_err(&self) -> u64 {
0
}
fn relogin_attempts(&self) -> u64 {
0
}
}
fn ok_count(enter: EnterOutcome) -> (usize, PollGuard) {
match enter {
EnterOutcome::Ok { per_vtoken, guard } => (per_vtoken, guard),
other => panic!("expected EnterOutcome::Ok, got {other:?}"),
}
}
#[test]
fn poll_tracker_counts_concurrent_polls_and_releases_on_drop() {
let tracker = Arc::new(PollTracker::default());
tracker.set_hub_cap(MAX_HUB_POLLS_DEFAULT);
let (c1, g1) = ok_count(tracker.enter("vt-a"));
assert_eq!(c1, 1, "first poll is alone");
let (c2, g2) = ok_count(tracker.enter("vt-a"));
assert_eq!(c2, 2, "second concurrent poll on same vtoken detected");
let (c_other, _g_other) = ok_count(tracker.enter("vt-b"));
assert_eq!(c_other, 1);
drop(g2);
let (c3, _g3) = ok_count(tracker.enter("vt-a"));
assert_eq!(
c3, 2,
"count drops when a guard is released, then rises again"
);
drop(g1);
drop(_g3);
let (c4, _g4) = ok_count(tracker.enter("vt-a"));
assert_eq!(c4, 1);
}
#[test]
fn poll_tracker_caps_concurrent() {
let tracker = Arc::new(PollTracker::default());
tracker.set_hub_cap(MAX_HUB_POLLS_DEFAULT);
let mut guards = Vec::with_capacity(MAX_CONCURRENT_POLLS_PER_VTOKEN);
for expected in 1..=MAX_CONCURRENT_POLLS_PER_VTOKEN {
let (c, g) = ok_count(tracker.enter("vt-cap"));
assert_eq!(
c, expected,
"enter #{expected} must report {expected} active polls"
);
guards.push(g);
}
let (over, g_over) = ok_count(tracker.enter("vt-cap"));
assert_eq!(
over,
MAX_CONCURRENT_POLLS_PER_VTOKEN + 1,
"the (MAX+1)th concurrent poll must be observable above the cap"
);
assert!(
over > MAX_CONCURRENT_POLLS_PER_VTOKEN,
"the cap is the 429 boundary; the handler gates on this"
);
drop(g_over);
let (back_to_max, g_back_to_max) = ok_count(tracker.enter("vt-cap"));
assert_eq!(
back_to_max,
MAX_CONCURRENT_POLLS_PER_VTOKEN + 1,
"the freshly entered guard again pushes the count to MAX+1"
);
drop(g_back_to_max);
drop(guards);
}
#[test]
fn poll_tracker_enforces_hub_wide_cap() {
let tracker = Arc::new(PollTracker::default());
tracker.set_hub_cap(2);
let (_c1, g1) = ok_count(tracker.enter("vt-a"));
let (_c2, g2) = ok_count(tracker.enter("vt-b"));
assert_eq!(tracker.total_polls(), 2);
match tracker.enter("vt-c") {
EnterOutcome::HubLimitReached { total, cap } => {
assert_eq!(total, 2);
assert_eq!(cap, 2);
}
other => panic!("expected HubLimitReached, got {other:?}"),
}
assert_eq!(tracker.total_polls(), 2);
drop(g1);
assert_eq!(tracker.total_polls(), 1);
let (_c3, g3) = ok_count(tracker.enter("vt-c"));
assert_eq!(tracker.total_polls(), 2);
drop(g2);
drop(g3);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrent_register_and_route_does_not_deadlock() {
let store = Store::connect("sqlite::memory:")
.await
.expect("in-memory store");
let upstream =
Arc::new(UpstreamClient::new("sk-test".to_string(), None).expect("test upstream client"));
let queue = Arc::new(InMemoryQueue::new());
let (_tx, shutdown_rx) = tokio::sync::watch::channel(false);
let state = HubState::new(
upstream,
Arc::new(store),
queue,
shutdown_rx,
"test-relay-secret".to_string(),
AdminConfig::from_env(),
);
let mut handles = vec![];
for i in 0..8 {
let s = Arc::clone(&state);
handles.push(tokio::spawn(async move {
for j in 0..10 {
crate::server::pairing::register_client_in_hub(
&s,
format!("client-{i}-{j}"),
None,
None,
)
.await;
}
}));
}
for _ in 0..4 {
let s = Arc::clone(&state);
handles.push(tokio::spawn(async move {
for _ in 0..20 {
let _ = s.routing.router.lock().await.get_route("any_user");
tokio::task::yield_now().await;
}
}));
}
let timeout = tokio::time::timeout(
std::time::Duration::from_secs(5),
futures_util::future::join_all(handles),
)
.await;
assert!(
timeout.is_ok(),
"concurrent register+route timed out (possible deadlock)"
);
}
#[tokio::test]
async fn test_build_hub_ext_for_vctx_timeout() {
let store = Store::connect("sqlite::memory:")
.await
.expect("in-memory store");
let _tx = store.pool().begin().await.unwrap();
tokio::time::pause();
let hub_ext = build_hub_ext_for_vctx(&store, "vctx-test", "vtoken-test", None).await;
assert!(hub_ext.is_some());
let ext = hub_ext.unwrap();
assert_eq!(ext.session_name, Some("default".to_string()));
assert_eq!(ext.session_id, None);
}
#[tokio::test]
async fn test_build_hub_ext_for_vctx_timeout_with_session_override() {
let store = Store::connect("sqlite::memory:")
.await
.expect("in-memory store");
let _tx = store.pool().begin().await.unwrap();
tokio::time::pause();
let hub_ext = build_hub_ext_for_vctx(
&store,
"vctx-test",
"vtoken-test",
Some("override".to_string()),
)
.await;
assert!(hub_ext.is_some());
let ext = hub_ext.unwrap();
assert_eq!(ext.session_name, Some("override".to_string()));
assert_eq!(ext.session_id, None);
}
async fn make_state() -> Arc<HubState> {
let upstream =
Arc::new(UpstreamClient::new("sk-test".to_string(), None).expect("test upstream client"));
let store = Arc::new(
Store::connect("sqlite::memory:")
.await
.expect("in-memory store"),
);
let queue = Arc::new(InMemoryQueue::new());
let (_tx, shutdown_rx) = tokio::sync::watch::channel(false);
HubState::new(
upstream,
store,
queue,
shutdown_rx,
"test-relay-secret".to_string(),
AdminConfig::from_env(),
)
}
#[tokio::test]
async fn hub_state_new_populates_all_sub_states() {
let state = make_state().await;
assert!(Arc::strong_count(&state.ilink.upstream) >= 1);
assert_eq!(
state.ilink.ilink_status.load(Ordering::Relaxed),
ilink_status::UNKNOWN,
"iLink status starts at UNKNOWN"
);
let _rx = state.ilink.qr_tx.subscribe();
let _ = state.ilink.relogin_tx.send(());
assert!(
state
.routing
.router
.lock()
.await
.get_route("any_user")
.is_none(),
"fresh Router has no per-user route"
);
assert_eq!(state.clients.registry.read().await.all_clients().len(), 0);
assert!(Arc::strong_count(&state.metrics) >= 1);
assert_eq!(state.metrics.messages_dispatched.load(Ordering::Relaxed), 0);
}
#[tokio::test]
async fn sub_states_are_independently_usable() {
let state = make_state().await;
assert_eq!(state.ilink.upstream.polls_ok(), 0);
let vtoken = "vt-abc".to_string();
state
.routing
.router
.lock()
.await
.set_route("user-x", vtoken.clone());
assert_eq!(
state.routing.router.lock().await.get_route("user-x"),
Some(vtoken.as_str())
);
let weixin_msg = crate::ilink::types::WeixinMessage::default();
let push_result = state.clients.queue.push(&vtoken, weixin_msg).await;
assert!(
push_result.is_ok(),
"in-memory queue accepts the pushed message"
);
}
#[test]
fn sub_state_structs_carry_expected_fields() {
fn assert_ilink_fields(_s: &IlinkConnState) {
let _ = &_s.upstream;
let _ = &_s.shutdown;
let _ = &_s.ilink_status;
let _ = &_s.qr_tx;
let _ = &_s.qr_last_ready;
let _ = &_s.relogin_tx;
}
fn assert_routing_fields(_s: &RoutingState) {
let _ = &_s.router;
}
fn assert_client_fields(_s: &ClientState) {
let _ = &_s.registry;
let _ = &_s.pairing;
let _ = &_s.queue;
let _ = &_s.poll_tracker;
}
let (_tx, _rx) = tokio::sync::watch::channel(false);
let _upstream =
Arc::new(UpstreamClient::new("sk-test".to_string(), None).expect("test upstream client"));
let _queue: Arc<dyn MessageQueue> = Arc::new(InMemoryQueue::new());
let ilink = IlinkConnState::new(
Arc::new(UpstreamClient::new("sk-test".to_string(), None).expect("test upstream client")),
_rx,
);
assert_ilink_fields(&ilink);
let routing = RoutingState::new();
assert_routing_fields(&routing);
let client = ClientState::new(_queue);
assert_client_fields(&client);
}
#[tokio::test]
async fn hub_state_metrics_are_shared_with_sub_state_paths() {
let state = make_state().await;
state
.metrics
.messages_dispatched
.fetch_add(7, Ordering::Relaxed);
let metrics_clone = Arc::clone(&state.metrics);
assert_eq!(metrics_clone.messages_dispatched.load(Ordering::Relaxed), 7);
}
#[test]
fn latency_histogram_submillisecond_observation_increments_sum_us() {
let h = LatencyHistogram::new();
assert_eq!(h.count.load(Ordering::Relaxed), 0);
assert_eq!(h.sum_us.load(Ordering::Relaxed), 0);
h.observe(std::time::Duration::from_micros(500));
assert_eq!(h.count.load(Ordering::Relaxed), 1);
let sum_us = h.sum_us.load(Ordering::Relaxed);
assert!(
sum_us >= 500,
"sub-millisecond observation must contribute to sum_us (got {sum_us})"
);
}
#[test]
fn latency_histogram_submillisecond_observation_lands_in_first_bucket() {
let h = LatencyHistogram::new();
h.observe(std::time::Duration::from_micros(500));
let first_bucket = h.buckets[0].load(Ordering::Relaxed);
assert_eq!(
first_bucket, 1,
"sub-millisecond observation falls into the `le=1` bucket"
);
}
#[test]
fn latency_histogram_millisecond_observation_still_works() {
let h = LatencyHistogram::new();
h.observe(std::time::Duration::from_millis(42));
assert_eq!(h.count.load(Ordering::Relaxed), 1);
let sum_us = h.sum_us.load(Ordering::Relaxed);
assert_eq!(sum_us, 42_000);
let bucket_100 = h.buckets[3].load(Ordering::Relaxed);
assert_eq!(
bucket_100, 1,
"42 ms observation falls into the `le=100` bucket (index 3)"
);
}
#[test]
fn latency_histogram_render_sum_uses_milliseconds() {
let h = LatencyHistogram::new();
for _ in 0..4 {
h.observe(std::time::Duration::from_micros(250));
}
let sum_us = h.sum_us.load(Ordering::Relaxed);
assert_eq!(sum_us, 1_000);
let sum_ms = sum_us / 1000;
assert_eq!(
sum_ms, 1,
"rendered _sum must be sum_us / 1000 (milliseconds on the wire)"
);
}
#[test]
fn latency_guard_records_submillisecond_elapsed() {
let h = LatencyHistogram::new();
{
let _g = LatencyGuard::new(&h);
}
assert_eq!(h.count.load(Ordering::Relaxed), 1);
let _ = h.sum_us.load(Ordering::Relaxed);
}
#[tokio::test]
async fn test_extracted_hub_commands() {
use super::commands::*;
let store = Store::connect("sqlite::memory:")
.await
.expect("in-memory store");
let upstream =
Arc::new(UpstreamClient::new("sk-test".to_string(), None).expect("test upstream client"));
let queue = Arc::new(InMemoryQueue::new());
let (_tx, shutdown_rx) = tokio::sync::watch::channel(false);
let state = Arc::new(HubState::new(
upstream,
Arc::new(store),
queue,
shutdown_rx,
"test-relay-secret".to_string(),
AdminConfig::from_env(),
));
let from_user = "user-123";
let list_res = handle_cmd_list(&state, from_user).await;
assert!(list_res.contains("尚未注册任何后端客户端"));
let use_res = handle_cmd_use(&state, from_user, "non-existent").await;
assert!(use_res.contains("未找到名为 `non-existent` 的后端"));
let out_a =
crate::server::pairing::register_client_in_hub(&state, "client-a".to_string(), None, None)
.await;
let vt_a = out_a.hashed;
let out_b =
crate::server::pairing::register_client_in_hub(&state, "client-b".to_string(), None, None)
.await;
let vt_b = out_b.hashed;
let list_res2 = handle_cmd_list(&state, from_user).await;
assert!(list_res2.contains("🔴 1. `client-a`"));
assert!(list_res2.contains("🔴 2. `client-b`"));
state.clients.registry.write().await.mark_online(&vt_a);
state.clients.registry.write().await.mark_online(&vt_b);
let list_res3 = handle_cmd_list(&state, from_user).await;
assert!(list_res3.contains("🟢 1. `client-a`"));
assert!(list_res3.contains("🟢 2. `client-b`"));
let use_res2 = handle_cmd_use(&state, from_user, "client-a").await;
assert!(use_res2.contains("已切换到 `client-a`"));
{
let router = state.routing.router.lock().await;
assert_eq!(router.get_route(from_user), Some(vt_a.as_str()));
}
let list_res4 = handle_cmd_list(&state, from_user).await;
assert!(list_res4.contains("`client-a` ✅"));
let use_res_alias = handle_cmd_use(&state, from_user, "2").await;
assert!(use_res_alias.contains("已切换到 `client-b`"));
{
let router = state.routing.router.lock().await;
assert_eq!(router.get_route(from_user), Some(vt_b.as_str()));
}
let _ = handle_cmd_use(&state, from_user, "1").await;
let real_ctx = "ctx-789";
let session_res =
handle_cmd_session_new(&state, from_user, real_ctx, None, "session-1", "uuid-1234").await;
assert!(session_res.contains("session `session-1`"));
let vctx =
super::dispatch::resolve_vctx_for_message(&state, real_ctx, from_user, None, None).await;
let active_session = state
.store
.get_active_session_name(&vctx, &vt_a)
.await
.unwrap();
assert_eq!(active_session, "session-1");
let list_sessions = handle_cmd_session_list(&state, from_user, real_ctx, None).await;
assert!(list_sessions.contains("session-1") && list_sessions.contains("✅"));
let session_res_empty = handle_cmd_session_new(&state, from_user, real_ctx, None, "", "").await;
assert!(session_res_empty.contains("session ``"));
let use_session_res =
handle_cmd_session_use(&state, from_user, real_ctx, None, "session-2").await;
assert!(use_session_res.contains("session `session-2`"));
let delete_session_res =
handle_cmd_session_delete(&state, from_user, real_ctx, None, "session-1").await;
assert!(delete_session_res.contains("session `session-1`"));
let delete_active_res =
handle_cmd_session_delete(&state, from_user, real_ctx, None, "session-2").await;
assert!(delete_active_res.contains("无法删除当前活跃的 session"));
}
#[tokio::test]
async fn test_adversarial_hub_commands() {
use super::commands::*;
let store = Store::connect("sqlite::memory:")
.await
.expect("in-memory store");
let upstream =
Arc::new(UpstreamClient::new("sk-test".to_string(), None).expect("test upstream client"));
let queue = Arc::new(InMemoryQueue::new());
let (_tx, shutdown_rx) = tokio::sync::watch::channel(false);
let state = Arc::new(HubState::new(
upstream,
Arc::new(store),
queue,
shutdown_rx,
"test-relay-secret".to_string(),
AdminConfig::from_env(),
));
let from_user = "user-adversarial";
let real_ctx = "ctx-adversarial";
let list_res = handle_cmd_session_list(&state, from_user, real_ctx, None).await;
assert!(list_res.contains("当前未路由到任何后端"));
let new_res = handle_cmd_session_new(&state, from_user, real_ctx, None, "test", "").await;
assert!(new_res.contains("当前未路由到任何后端"));
let use_res = handle_cmd_session_use(&state, from_user, real_ctx, None, "test").await;
assert!(use_res.contains("当前未路由到任何后端"));
let del_res = handle_cmd_session_delete(&state, from_user, real_ctx, None, "test").await;
assert!(del_res.contains("当前未路由到任何后端"));
let out_a =
crate::server::pairing::register_client_in_hub(&state, "client-a".to_string(), None, None)
.await;
let vt_a = out_a.hashed;
state.clients.registry.write().await.mark_online(&vt_a);
let _ = handle_cmd_use(&state, from_user, "client-a").await;
let long_name = "a".repeat(1000);
let new_long = handle_cmd_session_new(&state, from_user, real_ctx, None, &long_name, "").await;
assert!(new_long.contains("session `aaaaaaaa"));
let injection_name = "' OR '1'='1";
let new_inject =
handle_cmd_session_new(&state, from_user, real_ctx, None, injection_name, "").await;
assert!(new_inject.contains("session `' OR '1'='1`"));
let path_name = "../../../etc/passwd";
let new_path = handle_cmd_session_new(&state, from_user, real_ctx, None, path_name, "").await;
assert!(new_path.contains("session `../../../etc/passwd`"));
let unicode_name = "会话🚀🌟";
let new_unicode =
handle_cmd_session_new(&state, from_user, real_ctx, None, unicode_name, "").await;
assert!(new_unicode.contains("session `会话🚀🌟`"));
let use_nonexistent =
handle_cmd_session_use(&state, from_user, real_ctx, None, "nonexistent-session").await;
assert!(use_nonexistent.contains("已切换到 session `nonexistent-session`"));
let delete_nonexistent =
handle_cmd_session_delete(&state, from_user, real_ctx, None, "not-real").await;
assert!(delete_nonexistent.contains("未找到 session `not-real`"));
}