mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 10:57:03 +08:00
Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7b8048c6ae | ||
|
|
ec95f2ca1f | ||
|
|
aa7dbe67d3 |
Generated
+1
-1
@@ -661,7 +661,7 @@ dependencies = [
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "aether-tunnel"
|
name = "aether-tunnel"
|
||||||
version = "0.3.16"
|
version = "0.3.17"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"aether-contracts",
|
"aether-contracts",
|
||||||
"aether-gateway",
|
"aether-gateway",
|
||||||
|
|||||||
@@ -0,0 +1,186 @@
|
|||||||
|
use std::collections::VecDeque;
|
||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
|
use bytes::{Bytes, BytesMut};
|
||||||
|
use parking_lot::Mutex;
|
||||||
|
use tokio::sync::Notify;
|
||||||
|
|
||||||
|
const CHUNK_BYTES: usize = 32 * 1024;
|
||||||
|
|
||||||
|
#[derive(Debug)]
|
||||||
|
pub enum LocalBodyEvent {
|
||||||
|
Chunk(Bytes),
|
||||||
|
End,
|
||||||
|
Error(String),
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Default)]
|
||||||
|
struct BufferState {
|
||||||
|
chunks: VecDeque<BytesMut>,
|
||||||
|
bytes: usize,
|
||||||
|
terminal: Option<Result<(), String>>,
|
||||||
|
receiver_taken: bool,
|
||||||
|
receiver_closed: bool,
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) struct ResponseBuffer {
|
||||||
|
state: Mutex<BufferState>,
|
||||||
|
notify: Notify,
|
||||||
|
capacity: usize,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ResponseBuffer {
|
||||||
|
pub(super) fn new(capacity: usize) -> Arc<Self> {
|
||||||
|
Arc::new(Self {
|
||||||
|
state: Mutex::new(BufferState::default()),
|
||||||
|
notify: Notify::new(),
|
||||||
|
capacity,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) fn take_receiver(self: &Arc<Self>) -> Option<BodyReceiver> {
|
||||||
|
let mut state = self.state.lock();
|
||||||
|
if state.receiver_taken {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
state.receiver_taken = true;
|
||||||
|
Some(BodyReceiver {
|
||||||
|
buffer: Arc::clone(self),
|
||||||
|
finished: false,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) fn push(&self, mut payload: Bytes) -> bool {
|
||||||
|
let mut state = self.state.lock();
|
||||||
|
if state.terminal.is_some()
|
||||||
|
|| state.receiver_closed
|
||||||
|
|| payload.len() > self.capacity.saturating_sub(state.bytes)
|
||||||
|
{
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
state.bytes += payload.len();
|
||||||
|
while !payload.is_empty() {
|
||||||
|
if let Some(tail) = state
|
||||||
|
.chunks
|
||||||
|
.back_mut()
|
||||||
|
.filter(|chunk| chunk.len() < CHUNK_BYTES)
|
||||||
|
{
|
||||||
|
let count = payload.len().min(CHUNK_BYTES - tail.len());
|
||||||
|
tail.extend_from_slice(&payload.split_to(count));
|
||||||
|
} else {
|
||||||
|
let count = payload.len().min(CHUNK_BYTES);
|
||||||
|
let chunk = payload.split_to(count);
|
||||||
|
state.chunks.push_back(
|
||||||
|
chunk
|
||||||
|
.try_into_mut()
|
||||||
|
.unwrap_or_else(|chunk| BytesMut::from(chunk.as_ref())),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
drop(state);
|
||||||
|
self.notify.notify_waiters();
|
||||||
|
true
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) fn finish(&self, result: Result<(), String>) {
|
||||||
|
let mut state = self.state.lock();
|
||||||
|
if state.terminal.is_none() {
|
||||||
|
state.terminal = Some(result);
|
||||||
|
}
|
||||||
|
drop(state);
|
||||||
|
self.notify.notify_waiters();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) struct BodyReceiver {
|
||||||
|
buffer: Arc<ResponseBuffer>,
|
||||||
|
finished: bool,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl BodyReceiver {
|
||||||
|
pub(super) async fn recv(&mut self) -> Option<LocalBodyEvent> {
|
||||||
|
if self.finished {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
loop {
|
||||||
|
let notified = self.buffer.notify.notified();
|
||||||
|
tokio::pin!(notified);
|
||||||
|
notified.as_mut().enable();
|
||||||
|
{
|
||||||
|
let mut state = self.buffer.state.lock();
|
||||||
|
if let Some(chunk) = state.chunks.pop_front() {
|
||||||
|
state.bytes -= chunk.len();
|
||||||
|
return Some(LocalBodyEvent::Chunk(chunk.freeze()));
|
||||||
|
}
|
||||||
|
if let Some(terminal) = state.terminal.take() {
|
||||||
|
self.finished = true;
|
||||||
|
state.receiver_closed = true;
|
||||||
|
return Some(match terminal {
|
||||||
|
Ok(()) => LocalBodyEvent::End,
|
||||||
|
Err(error) => LocalBodyEvent::Error(error),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
}
|
||||||
|
notified.await;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Drop for BodyReceiver {
|
||||||
|
fn drop(&mut self) {
|
||||||
|
let mut state = self.buffer.state.lock();
|
||||||
|
state.receiver_closed = true;
|
||||||
|
state.chunks.clear();
|
||||||
|
state.bytes = 0;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn error_survives_a_full_buffer() {
|
||||||
|
let buffer = ResponseBuffer::new(CHUNK_BYTES);
|
||||||
|
let mut receiver = buffer.take_receiver().unwrap();
|
||||||
|
assert!(buffer.push(Bytes::from(vec![b'x'; CHUNK_BYTES])));
|
||||||
|
buffer.finish(Err("proxy disconnected".into()));
|
||||||
|
assert!(matches!(
|
||||||
|
receiver.recv().await,
|
||||||
|
Some(LocalBodyEvent::Chunk(_))
|
||||||
|
));
|
||||||
|
assert!(
|
||||||
|
matches!(receiver.recv().await, Some(LocalBodyEvent::Error(error)) if error == "proxy disconnected")
|
||||||
|
);
|
||||||
|
assert!(receiver.recv().await.is_none());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn small_frames_are_coalesced_within_the_byte_budget() {
|
||||||
|
let buffer = ResponseBuffer::new(4096);
|
||||||
|
let mut receiver = buffer.take_receiver().unwrap();
|
||||||
|
for _ in 0..4096 {
|
||||||
|
assert!(buffer.push(Bytes::from_static(b"x")));
|
||||||
|
}
|
||||||
|
assert!(!buffer.push(Bytes::from_static(b"x")));
|
||||||
|
assert_eq!(buffer.state.lock().chunks.len(), 1);
|
||||||
|
buffer.finish(Ok(()));
|
||||||
|
assert!(
|
||||||
|
matches!(receiver.recv().await, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 4096)
|
||||||
|
);
|
||||||
|
assert!(matches!(receiver.recv().await, Some(LocalBodyEvent::End)));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn terminal_wakes_an_empty_receiver_and_is_not_overwritten() {
|
||||||
|
let buffer = ResponseBuffer::new(1024);
|
||||||
|
let mut receiver = buffer.take_receiver().unwrap();
|
||||||
|
let task = tokio::spawn(async move { receiver.recv().await });
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
buffer.finish(Err("cancelled".into()));
|
||||||
|
buffer.finish(Ok(()));
|
||||||
|
assert!(
|
||||||
|
matches!(task.await.unwrap(), Some(LocalBodyEvent::Error(error)) if error == "cancelled")
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,250 @@
|
|||||||
|
use super::*;
|
||||||
|
|
||||||
|
async fn fixture(
|
||||||
|
window: u32,
|
||||||
|
capacity: usize,
|
||||||
|
) -> (
|
||||||
|
Arc<HubRouter>,
|
||||||
|
Arc<ProxyConn>,
|
||||||
|
Arc<LocalStream>,
|
||||||
|
aether_runtime::BoundedQueueReceiver<Message>,
|
||||||
|
) {
|
||||||
|
let hub = HubRouter::new(ControlPlaneClient::disabled());
|
||||||
|
let (sender, receiver) = bounded_queue(capacity);
|
||||||
|
let (close_tx, _) = watch::channel(false);
|
||||||
|
let connection = Arc::new(
|
||||||
|
ProxyConn::new(
|
||||||
|
99,
|
||||||
|
"flow-test".into(),
|
||||||
|
"flow-test".into(),
|
||||||
|
sender,
|
||||||
|
close_tx,
|
||||||
|
16,
|
||||||
|
3,
|
||||||
|
)
|
||||||
|
.with_settings(protocol::SettingsPayload {
|
||||||
|
initial_stream_window_bytes: window,
|
||||||
|
min_window_update_bytes: (window / 4).max(1),
|
||||||
|
drain_deadline_ms: 1000,
|
||||||
|
}),
|
||||||
|
);
|
||||||
|
hub.register_proxy(Arc::clone(&connection));
|
||||||
|
let stream = hub.open_local_stream("flow-test", &meta()).await.unwrap();
|
||||||
|
(hub, connection, stream, receiver)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn meta() -> protocol::RequestMeta {
|
||||||
|
protocol::RequestMeta {
|
||||||
|
provider_id: None,
|
||||||
|
endpoint_id: None,
|
||||||
|
key_id: None,
|
||||||
|
method: "GET".into(),
|
||||||
|
url: "https://example.com".into(),
|
||||||
|
headers: HashMap::new(),
|
||||||
|
stream: true,
|
||||||
|
request_timeout_ms: None,
|
||||||
|
stream_first_byte_timeout_ms: None,
|
||||||
|
timeout: 30,
|
||||||
|
follow_redirects: None,
|
||||||
|
http1_only: false,
|
||||||
|
transport_profile: None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn headers(hub: &Arc<HubRouter>, stream: &LocalStream) {
|
||||||
|
let payload = serde_json::to_vec(&protocol::ResponseMeta {
|
||||||
|
status: 200,
|
||||||
|
headers: vec![],
|
||||||
|
})
|
||||||
|
.unwrap();
|
||||||
|
let mut frame = protocol::encode_frame(
|
||||||
|
stream.proxy_stream_id,
|
||||||
|
protocol::RESPONSE_HEADERS,
|
||||||
|
0,
|
||||||
|
&payload,
|
||||||
|
);
|
||||||
|
hub.handle_proxy_frame(stream.proxy_conn_id, &mut frame)
|
||||||
|
.await;
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn window_credit_is_retried_after_queue_pressure_and_cancelled_receive() {
|
||||||
|
let (hub, _, stream, mut outbound) = fixture(128, 1).await;
|
||||||
|
headers(&hub, &stream).await;
|
||||||
|
assert!(stream.push_body_chunk(Bytes::from(vec![b'x'; 64])));
|
||||||
|
let mut receiver = stream.take_body_receiver().unwrap();
|
||||||
|
assert!(
|
||||||
|
tokio::time::timeout(Duration::from_millis(10), receiver.recv())
|
||||||
|
.await
|
||||||
|
.is_err()
|
||||||
|
);
|
||||||
|
assert_eq!(*stream.response_consumed_since_update.lock(), 64);
|
||||||
|
outbound.recv().await.unwrap();
|
||||||
|
let event = tokio::time::timeout(Duration::from_secs(1), receiver.recv())
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert!(matches!(event, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 64));
|
||||||
|
assert_eq!(*stream.response_consumed_since_update.lock(), 0);
|
||||||
|
let Message::Binary(data) = outbound.recv().await.unwrap() else {
|
||||||
|
panic!("expected binary update")
|
||||||
|
};
|
||||||
|
let frame = aether_contracts::tunnel::Frame::decode(data).unwrap();
|
||||||
|
let update: protocol::WindowUpdatePayload = serde_json::from_slice(&frame.payload).unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
frame.msg_type,
|
||||||
|
aether_contracts::tunnel::MsgType::WindowUpdate
|
||||||
|
);
|
||||||
|
assert_eq!(update.delta_bytes, 64);
|
||||||
|
hub.cancel_local_stream(stream.id, "test complete");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn response_credit_is_not_returned_until_consumed() {
|
||||||
|
let (hub, _, stream, mut outbound) = fixture(128, 4).await;
|
||||||
|
outbound.recv().await.unwrap();
|
||||||
|
let mut body = protocol::encode_frame(
|
||||||
|
stream.proxy_stream_id,
|
||||||
|
protocol::RESPONSE_BODY,
|
||||||
|
0,
|
||||||
|
&[b'x'; 128],
|
||||||
|
);
|
||||||
|
hub.handle_proxy_frame(99, &mut body).await;
|
||||||
|
assert!(outbound.try_recv().is_err());
|
||||||
|
let mut receiver = stream.take_body_receiver().unwrap();
|
||||||
|
assert!(
|
||||||
|
matches!(receiver.recv().await, Some(LocalBodyEvent::Chunk(chunk)) if chunk.len() == 128)
|
||||||
|
);
|
||||||
|
assert!(outbound.try_recv().is_ok());
|
||||||
|
hub.cancel_local_stream(stream.id, "test complete");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn cancelled_stream_open_releases_slot_without_resetting_connection() {
|
||||||
|
let (hub, connection, first_stream, mut outbound) = fixture(128, 1).await;
|
||||||
|
let opening_hub = Arc::clone(&hub);
|
||||||
|
let opening =
|
||||||
|
tokio::spawn(async move { opening_hub.open_local_stream("flow-test", &meta()).await });
|
||||||
|
tokio::time::timeout(Duration::from_secs(1), async {
|
||||||
|
while hub.local_streams.len() != 2 {
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
opening.abort();
|
||||||
|
assert!(matches!(opening.await, Err(error) if error.is_cancelled()));
|
||||||
|
assert_eq!(connection.stream_count.load(Ordering::Relaxed), 1);
|
||||||
|
assert_eq!(hub.local_streams.len(), 1);
|
||||||
|
assert_eq!(hub.proxy_to_local.len(), 1);
|
||||||
|
assert!(connection.is_available());
|
||||||
|
outbound.recv().await.unwrap();
|
||||||
|
assert!(outbound.try_recv().is_err());
|
||||||
|
let next_stream = hub.open_local_stream("flow-test", &meta()).await.unwrap();
|
||||||
|
outbound.recv().await.unwrap();
|
||||||
|
hub.cancel_local_stream(first_stream.id, "test complete");
|
||||||
|
outbound.recv().await.unwrap();
|
||||||
|
hub.cancel_local_stream(next_stream.id, "test complete");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn full_response_buffer_preserves_disconnect_error() {
|
||||||
|
let (hub, connection, stream, mut outbound) = fixture(4 * 1024 * 1024, 512).await;
|
||||||
|
outbound.recv().await.unwrap();
|
||||||
|
headers(&hub, &stream).await;
|
||||||
|
let mut receiver = stream.take_body_receiver().unwrap();
|
||||||
|
for _ in 0..128 {
|
||||||
|
let mut frame = protocol::encode_frame(
|
||||||
|
stream.proxy_stream_id,
|
||||||
|
protocol::RESPONSE_BODY,
|
||||||
|
0,
|
||||||
|
&vec![b'x'; 32 * 1024],
|
||||||
|
);
|
||||||
|
hub.handle_proxy_frame(99, &mut frame).await;
|
||||||
|
}
|
||||||
|
hub.unregister_proxy(connection.id, &connection.node_id);
|
||||||
|
let mut bytes = 0;
|
||||||
|
loop {
|
||||||
|
match receiver.recv().await {
|
||||||
|
Some(LocalBodyEvent::Chunk(chunk)) => bytes += chunk.len(),
|
||||||
|
Some(LocalBodyEvent::Error(error)) => {
|
||||||
|
assert!(error.contains("disconnected"));
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
event => panic!("disconnect must not become normal EOF: {event:?}"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
assert_eq!(bytes, 4 * 1024 * 1024);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn slow_stream_does_not_block_another_stream_on_the_same_connection() {
|
||||||
|
let (hub, _, slow, mut outbound) = fixture(128, 512).await;
|
||||||
|
outbound.recv().await.unwrap();
|
||||||
|
let fast = hub.open_local_stream("flow-test", &meta()).await.unwrap();
|
||||||
|
assert!(slow.push_body_chunk(Bytes::from(vec![b'x'; 128])));
|
||||||
|
let mut overflowing = protocol::encode_frame(
|
||||||
|
slow.proxy_stream_id,
|
||||||
|
protocol::RESPONSE_BODY,
|
||||||
|
0,
|
||||||
|
b"overflow",
|
||||||
|
);
|
||||||
|
tokio::time::timeout(Duration::from_secs(1), async {
|
||||||
|
hub.handle_proxy_frame(99, &mut overflowing).await;
|
||||||
|
headers(&hub, &fast).await;
|
||||||
|
assert_eq!(
|
||||||
|
fast.wait_headers(Duration::from_secs(1))
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.status,
|
||||||
|
200
|
||||||
|
);
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.expect("slow stream must not block connection reader");
|
||||||
|
assert!(!hub.local_streams.contains_key(&slow.id));
|
||||||
|
assert!(hub.local_streams.contains_key(&fast.id));
|
||||||
|
hub.cancel_local_stream(fast.id, "test complete");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn cancelling_a_stream_wakes_request_window_waiters() {
|
||||||
|
let (_, _, stream, _) = fixture(128, 512).await;
|
||||||
|
*stream.request_window.available.lock() = 0;
|
||||||
|
let waiter = tokio::spawn({
|
||||||
|
let stream = Arc::clone(&stream);
|
||||||
|
async move {
|
||||||
|
stream
|
||||||
|
.acquire_request_window(1, Duration::from_secs(30))
|
||||||
|
.await
|
||||||
|
}
|
||||||
|
});
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
stream.fail("cancelled");
|
||||||
|
assert!(tokio::time::timeout(Duration::from_secs(1), waiter)
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.unwrap()
|
||||||
|
.is_err());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||||
|
async fn concurrent_headers_and_credit_updates_do_not_lose_notifications() {
|
||||||
|
for index in 0..256 {
|
||||||
|
let stream = Arc::new(LocalStream::new(index, "test".into(), 1, 1, 1));
|
||||||
|
let window = Arc::new(StreamFlowWindow::new(0));
|
||||||
|
let waiter = tokio::spawn({
|
||||||
|
let stream = Arc::clone(&stream);
|
||||||
|
let window = Arc::clone(&window);
|
||||||
|
async move {
|
||||||
|
stream.wait_headers(Duration::from_secs(1)).await.unwrap();
|
||||||
|
window.acquire(1, Duration::from_secs(1)).await.unwrap();
|
||||||
|
}
|
||||||
|
});
|
||||||
|
stream.set_response_headers(protocol::ResponseMeta {
|
||||||
|
status: 200,
|
||||||
|
headers: vec![],
|
||||||
|
});
|
||||||
|
window.add(1);
|
||||||
|
waiter.await.unwrap();
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -12,10 +12,11 @@ use axum::extract::ws::Message;
|
|||||||
use bytes::Bytes;
|
use bytes::Bytes;
|
||||||
use dashmap::DashMap;
|
use dashmap::DashMap;
|
||||||
use parking_lot::{Mutex, RwLock};
|
use parking_lot::{Mutex, RwLock};
|
||||||
use tokio::sync::mpsc;
|
|
||||||
use tokio::sync::{watch, Notify};
|
use tokio::sync::{watch, Notify};
|
||||||
use tracing::{debug, info, warn};
|
use tracing::{debug, info, warn};
|
||||||
|
|
||||||
|
pub use super::body::LocalBodyEvent;
|
||||||
|
use super::body::{BodyReceiver, ResponseBuffer};
|
||||||
use super::control_plane::ControlPlaneClient;
|
use super::control_plane::ControlPlaneClient;
|
||||||
use super::protocol;
|
use super::protocol;
|
||||||
|
|
||||||
@@ -29,6 +30,10 @@ const DEFAULT_DRAIN_DEADLINE_MS: u64 = 30_000;
|
|||||||
const DEFAULT_NODE_STATUS_QUEUE_CAPACITY: usize = 1_024;
|
const DEFAULT_NODE_STATUS_QUEUE_CAPACITY: usize = 1_024;
|
||||||
const CONNECTION_WARMUP: Duration = Duration::from_secs(1);
|
const CONNECTION_WARMUP: Duration = Duration::from_secs(1);
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
#[path = "flow_control_tests.rs"]
|
||||||
|
mod flow_control_tests;
|
||||||
|
|
||||||
static STREAM_INITIAL_WINDOW_BYTES: LazyLock<u32> = LazyLock::new(|| {
|
static STREAM_INITIAL_WINDOW_BYTES: LazyLock<u32> = LazyLock::new(|| {
|
||||||
std::env::var("AETHER_TUNNEL_STREAM_INITIAL_WINDOW_BYTES")
|
std::env::var("AETHER_TUNNEL_STREAM_INITIAL_WINDOW_BYTES")
|
||||||
.ok()
|
.ok()
|
||||||
@@ -53,11 +58,16 @@ static NODE_STATUS_QUEUE_CAPACITY: LazyLock<usize> = LazyLock::new(|| {
|
|||||||
.unwrap_or(DEFAULT_NODE_STATUS_QUEUE_CAPACITY)
|
.unwrap_or(DEFAULT_NODE_STATUS_QUEUE_CAPACITY)
|
||||||
});
|
});
|
||||||
|
|
||||||
static STREAM_MIN_WINDOW_UPDATE_BYTES: LazyLock<u32> = LazyLock::new(|| {
|
pub(super) fn local_settings() -> protocol::SettingsPayload {
|
||||||
STREAM_INITIAL_WINDOW_BYTES
|
protocol::SettingsPayload {
|
||||||
.saturating_div(4)
|
initial_stream_window_bytes: (*STREAM_INITIAL_WINDOW_BYTES)
|
||||||
.clamp(1, 1024 * 1024)
|
.min(aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES as u32),
|
||||||
});
|
min_window_update_bytes: STREAM_INITIAL_WINDOW_BYTES
|
||||||
|
.saturating_div(4)
|
||||||
|
.clamp(1, 1024 * 1024),
|
||||||
|
drain_deadline_ms: *DRAIN_DEADLINE_MS,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||||
pub enum SendStatus {
|
pub enum SendStatus {
|
||||||
@@ -91,6 +101,7 @@ impl ConnHealthState {
|
|||||||
struct StreamFlowWindow {
|
struct StreamFlowWindow {
|
||||||
available: Mutex<u64>,
|
available: Mutex<u64>,
|
||||||
notify: Notify,
|
notify: Notify,
|
||||||
|
closed: AtomicBool,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl StreamFlowWindow {
|
impl StreamFlowWindow {
|
||||||
@@ -98,6 +109,7 @@ impl StreamFlowWindow {
|
|||||||
Self {
|
Self {
|
||||||
available: Mutex::new(u64::from(initial)),
|
available: Mutex::new(u64::from(initial)),
|
||||||
notify: Notify::new(),
|
notify: Notify::new(),
|
||||||
|
closed: AtomicBool::new(false),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -109,6 +121,12 @@ impl StreamFlowWindow {
|
|||||||
let requested = bytes as u64;
|
let requested = bytes as u64;
|
||||||
let started_at = Instant::now();
|
let started_at = Instant::now();
|
||||||
loop {
|
loop {
|
||||||
|
let notified = self.notify.notified();
|
||||||
|
tokio::pin!(notified);
|
||||||
|
notified.as_mut().enable();
|
||||||
|
if self.closed.load(Ordering::Acquire) {
|
||||||
|
return Err(());
|
||||||
|
}
|
||||||
{
|
{
|
||||||
let mut available = self.available.lock();
|
let mut available = self.available.lock();
|
||||||
if *available >= requested {
|
if *available >= requested {
|
||||||
@@ -120,10 +138,7 @@ impl StreamFlowWindow {
|
|||||||
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
|
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
|
||||||
return Err(());
|
return Err(());
|
||||||
};
|
};
|
||||||
if tokio::time::timeout(remaining, self.notify.notified())
|
if tokio::time::timeout(remaining, notified).await.is_err() {
|
||||||
.await
|
|
||||||
.is_err()
|
|
||||||
{
|
|
||||||
return Err(());
|
return Err(());
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -138,6 +153,11 @@ impl StreamFlowWindow {
|
|||||||
drop(available);
|
drop(available);
|
||||||
self.notify.notify_waiters();
|
self.notify.notify_waiters();
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn close(&self) {
|
||||||
|
self.closed.store(true, Ordering::Release);
|
||||||
|
self.notify.notify_waiters();
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, Copy)]
|
#[derive(Debug, Clone, Copy)]
|
||||||
@@ -209,6 +229,10 @@ impl BoundedOutbound {
|
|||||||
pub fn snapshot(&self) -> QueueSnapshot {
|
pub fn snapshot(&self) -> QueueSnapshot {
|
||||||
self.tx.snapshot()
|
self.tx.snapshot()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub(super) fn subscribe_close(&self) -> watch::Receiver<bool> {
|
||||||
|
self.close_tx.subscribe()
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub struct ProxyConn {
|
pub struct ProxyConn {
|
||||||
@@ -231,6 +255,7 @@ pub struct ProxyConn {
|
|||||||
flow_window_blocked_ms: AtomicU64,
|
flow_window_blocked_ms: AtomicU64,
|
||||||
write_latency_last_us: AtomicU64,
|
write_latency_last_us: AtomicU64,
|
||||||
write_latency_ewma_us: AtomicU64,
|
write_latency_ewma_us: AtomicU64,
|
||||||
|
settings: Mutex<protocol::SettingsPayload>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl ProxyConn {
|
impl ProxyConn {
|
||||||
@@ -244,6 +269,7 @@ impl ProxyConn {
|
|||||||
protocol_version: u8,
|
protocol_version: u8,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
Self {
|
Self {
|
||||||
|
settings: Mutex::new(local_settings()),
|
||||||
id,
|
id,
|
||||||
node_id,
|
node_id,
|
||||||
node_name,
|
node_name,
|
||||||
@@ -271,6 +297,11 @@ impl ProxyConn {
|
|||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub(super) fn with_settings(mut self, settings: protocol::SettingsPayload) -> Self {
|
||||||
|
*self.settings.get_mut() = settings;
|
||||||
|
self
|
||||||
|
}
|
||||||
|
|
||||||
pub fn with_tunnel_generation(mut self, tunnel_generation: String) -> Self {
|
pub fn with_tunnel_generation(mut self, tunnel_generation: String) -> Self {
|
||||||
self.node_generation = tunnel_generation;
|
self.node_generation = tunnel_generation;
|
||||||
self
|
self
|
||||||
@@ -565,13 +596,6 @@ pub struct LocalResponseHead {
|
|||||||
pub headers: Vec<(String, String)>,
|
pub headers: Vec<(String, String)>,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug)]
|
|
||||||
pub enum LocalBodyEvent {
|
|
||||||
Chunk(Bytes),
|
|
||||||
End,
|
|
||||||
Error(String),
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Default)]
|
#[derive(Debug, Default)]
|
||||||
struct LocalWaitState {
|
struct LocalWaitState {
|
||||||
response: Option<LocalResponseHead>,
|
response: Option<LocalResponseHead>,
|
||||||
@@ -585,10 +609,11 @@ pub struct LocalStream {
|
|||||||
proxy_stream_id: u32,
|
proxy_stream_id: u32,
|
||||||
request_window: StreamFlowWindow,
|
request_window: StreamFlowWindow,
|
||||||
response_consumed_since_update: Mutex<u64>,
|
response_consumed_since_update: Mutex<u64>,
|
||||||
|
min_window_update_bytes: u32,
|
||||||
|
response_connection: Mutex<Option<std::sync::Weak<ProxyConn>>>,
|
||||||
wait_state: Mutex<LocalWaitState>,
|
wait_state: Mutex<LocalWaitState>,
|
||||||
headers_notify: Notify,
|
headers_notify: Notify,
|
||||||
body_tx: mpsc::Sender<LocalBodyEvent>,
|
body: Arc<ResponseBuffer>,
|
||||||
body_rx: Mutex<Option<mpsc::Receiver<LocalBodyEvent>>>,
|
|
||||||
terminal: AtomicBool,
|
terminal: AtomicBool,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -600,7 +625,6 @@ impl LocalStream {
|
|||||||
proxy_stream_id: u32,
|
proxy_stream_id: u32,
|
||||||
initial_window_bytes: u32,
|
initial_window_bytes: u32,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
let (body_tx, body_rx) = mpsc::channel(128);
|
|
||||||
Self {
|
Self {
|
||||||
id,
|
id,
|
||||||
tunnel_generation,
|
tunnel_generation,
|
||||||
@@ -608,10 +632,11 @@ impl LocalStream {
|
|||||||
proxy_stream_id,
|
proxy_stream_id,
|
||||||
request_window: StreamFlowWindow::new(initial_window_bytes),
|
request_window: StreamFlowWindow::new(initial_window_bytes),
|
||||||
response_consumed_since_update: Mutex::new(0),
|
response_consumed_since_update: Mutex::new(0),
|
||||||
|
min_window_update_bytes: (initial_window_bytes / 4).clamp(1, 1024 * 1024),
|
||||||
|
response_connection: Mutex::new(None),
|
||||||
wait_state: Mutex::new(LocalWaitState::default()),
|
wait_state: Mutex::new(LocalWaitState::default()),
|
||||||
headers_notify: Notify::new(),
|
headers_notify: Notify::new(),
|
||||||
body_tx,
|
body: ResponseBuffer::new(initial_window_bytes as usize),
|
||||||
body_rx: Mutex::new(Some(body_rx)),
|
|
||||||
terminal: AtomicBool::new(false),
|
terminal: AtomicBool::new(false),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -632,26 +657,51 @@ impl LocalStream {
|
|||||||
self.request_window.add(delta);
|
self.request_window.add(delta);
|
||||||
}
|
}
|
||||||
|
|
||||||
fn response_window_update_delta(&self, bytes: usize) -> Option<u32> {
|
async fn flush_response_credit(&self) -> Result<(), String> {
|
||||||
if bytes == 0 {
|
if self.terminal.load(Ordering::Acquire) {
|
||||||
return None;
|
return Ok(());
|
||||||
}
|
}
|
||||||
|
let connection = self
|
||||||
let mut consumed = self.response_consumed_since_update.lock();
|
.response_connection
|
||||||
*consumed = consumed.saturating_add(bytes as u64);
|
.lock()
|
||||||
let threshold = u64::from(*STREAM_MIN_WINDOW_UPDATE_BYTES);
|
.as_ref()
|
||||||
if *consumed < threshold {
|
.and_then(std::sync::Weak::upgrade);
|
||||||
return None;
|
let Some(connection) = connection else {
|
||||||
|
return Ok(());
|
||||||
|
};
|
||||||
|
if connection.protocol_version() < 3 {
|
||||||
|
return Ok(());
|
||||||
}
|
}
|
||||||
|
let delta = {
|
||||||
let delta = (*consumed).min(u64::from(u32::MAX)) as u32;
|
let consumed = self.response_consumed_since_update.lock();
|
||||||
*consumed = consumed.saturating_sub(u64::from(delta));
|
if *consumed < u64::from(self.min_window_update_bytes) {
|
||||||
Some(delta)
|
return Ok(());
|
||||||
|
}
|
||||||
|
(*consumed).min(u64::from(u32::MAX)) as u32
|
||||||
|
};
|
||||||
|
let frame = protocol::encode_window_update(self.proxy_stream_id, delta);
|
||||||
|
if connection
|
||||||
|
.send_wait(Message::Binary(frame.into()), OUTBOUND_BACKPRESSURE_TIMEOUT)
|
||||||
|
.await
|
||||||
|
== SendStatus::Queued
|
||||||
|
{
|
||||||
|
let mut consumed = self.response_consumed_since_update.lock();
|
||||||
|
*consumed = consumed.saturating_sub(u64::from(delta));
|
||||||
|
return Ok(());
|
||||||
|
}
|
||||||
|
if self.terminal.load(Ordering::Acquire) {
|
||||||
|
return Ok(());
|
||||||
|
}
|
||||||
|
connection.request_close();
|
||||||
|
Err("proxy flow-control update failed".to_string())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn wait_headers(&self, timeout: Duration) -> Result<LocalResponseHead, String> {
|
pub async fn wait_headers(&self, timeout: Duration) -> Result<LocalResponseHead, String> {
|
||||||
tokio::time::timeout(timeout, async {
|
tokio::time::timeout(timeout, async {
|
||||||
loop {
|
loop {
|
||||||
|
let notified = self.headers_notify.notified();
|
||||||
|
tokio::pin!(notified);
|
||||||
|
notified.as_mut().enable();
|
||||||
let outcome = {
|
let outcome = {
|
||||||
let state = self.wait_state.lock();
|
let state = self.wait_state.lock();
|
||||||
if let Some(response) = &state.response {
|
if let Some(response) = &state.response {
|
||||||
@@ -662,15 +712,20 @@ impl LocalStream {
|
|||||||
if let Some(error) = outcome {
|
if let Some(error) = outcome {
|
||||||
return Err(error);
|
return Err(error);
|
||||||
}
|
}
|
||||||
self.headers_notify.notified().await;
|
notified.await;
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
.await
|
.await
|
||||||
.map_err(|_| "timed out waiting for response headers".to_string())?
|
.map_err(|_| "timed out waiting for response headers".to_string())?
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn take_body_receiver(&self) -> Option<mpsc::Receiver<LocalBodyEvent>> {
|
pub fn take_body_receiver(self: &Arc<Self>) -> Option<LocalBodyReceiver> {
|
||||||
self.body_rx.lock().take()
|
self.body.take_receiver().map(|receiver| LocalBodyReceiver {
|
||||||
|
receiver,
|
||||||
|
stream: Arc::clone(self),
|
||||||
|
failed: false,
|
||||||
|
pending: None,
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
fn set_response_headers(&self, meta: protocol::ResponseMeta) {
|
fn set_response_headers(&self, meta: protocol::ResponseMeta) {
|
||||||
@@ -690,21 +745,11 @@ impl LocalStream {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn push_body_chunk(&self, payload: Bytes) -> bool {
|
fn push_body_chunk(&self, payload: Bytes) -> bool {
|
||||||
if self.terminal.load(Ordering::Acquire) {
|
if self.terminal.load(Ordering::Acquire) {
|
||||||
return false;
|
return false;
|
||||||
}
|
}
|
||||||
// Use a timeout to prevent a slow consumer from blocking the shared
|
self.body.push(payload)
|
||||||
// proxy-connection reader (head-of-line blocking across streams).
|
|
||||||
match tokio::time::timeout(
|
|
||||||
Duration::from_secs(5),
|
|
||||||
self.body_tx.send(LocalBodyEvent::Chunk(payload)),
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
{
|
|
||||||
Ok(Ok(())) => true,
|
|
||||||
_ => false,
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn finish(&self) {
|
fn finish(&self) {
|
||||||
@@ -722,7 +767,8 @@ impl LocalStream {
|
|||||||
if notify {
|
if notify {
|
||||||
self.headers_notify.notify_waiters();
|
self.headers_notify.notify_waiters();
|
||||||
}
|
}
|
||||||
let _ = self.body_tx.try_send(LocalBodyEvent::End);
|
self.request_window.close();
|
||||||
|
self.body.finish(Ok(()));
|
||||||
}
|
}
|
||||||
|
|
||||||
fn fail(&self, error: impl Into<String>) {
|
fn fail(&self, error: impl Into<String>) {
|
||||||
@@ -742,7 +788,38 @@ impl LocalStream {
|
|||||||
if notify {
|
if notify {
|
||||||
self.headers_notify.notify_waiters();
|
self.headers_notify.notify_waiters();
|
||||||
}
|
}
|
||||||
let _ = self.body_tx.try_send(LocalBodyEvent::Error(error));
|
self.request_window.close();
|
||||||
|
self.body.finish(Err(error));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub struct LocalBodyReceiver {
|
||||||
|
receiver: BodyReceiver,
|
||||||
|
stream: Arc<LocalStream>,
|
||||||
|
failed: bool,
|
||||||
|
pending: Option<LocalBodyEvent>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl LocalBodyReceiver {
|
||||||
|
pub async fn recv(&mut self) -> Option<LocalBodyEvent> {
|
||||||
|
if self.failed {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
if self.pending.is_none() {
|
||||||
|
let event = self.receiver.recv().await?;
|
||||||
|
if let LocalBodyEvent::Chunk(chunk) = &event {
|
||||||
|
let mut consumed = self.stream.response_consumed_since_update.lock();
|
||||||
|
*consumed = consumed.saturating_add(chunk.len() as u64);
|
||||||
|
}
|
||||||
|
self.pending = Some(event);
|
||||||
|
}
|
||||||
|
if matches!(self.pending, Some(LocalBodyEvent::Chunk(_))) {
|
||||||
|
if let Err(error) = self.stream.flush_response_credit().await {
|
||||||
|
self.failed = true;
|
||||||
|
return Some(LocalBodyEvent::Error(error));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
self.pending.take()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -765,6 +842,21 @@ pub struct HubRouter {
|
|||||||
drain_reasons: Mutex<HashMap<String, u64>>,
|
drain_reasons: Mutex<HashMap<String, u64>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
struct PendingStreamGuard<'router> {
|
||||||
|
hub: &'router HubRouter,
|
||||||
|
connection: &'router ProxyConn,
|
||||||
|
stream_id: u64,
|
||||||
|
committed: bool,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Drop for PendingStreamGuard<'_> {
|
||||||
|
fn drop(&mut self) {
|
||||||
|
if !self.committed && self.hub.cleanup_local_stream(self.stream_id) {
|
||||||
|
self.connection.release_stream();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
struct NodeStatusEvent {
|
struct NodeStatusEvent {
|
||||||
node_id: String,
|
node_id: String,
|
||||||
authenticated_key: Option<String>,
|
authenticated_key: Option<String>,
|
||||||
@@ -1166,17 +1258,27 @@ impl HubRouter {
|
|||||||
|
|
||||||
// Frames encoded successfully -- now register the stream.
|
// Frames encoded successfully -- now register the stream.
|
||||||
let local_stream_id = self.next_local_stream_id.fetch_add(1, Ordering::Relaxed);
|
let local_stream_id = self.next_local_stream_id.fetch_add(1, Ordering::Relaxed);
|
||||||
let local_stream = Arc::new(LocalStream::new(
|
let settings = proxy_conn.settings.lock().clone();
|
||||||
|
let mut local_stream = LocalStream::new(
|
||||||
local_stream_id,
|
local_stream_id,
|
||||||
proxy_conn.node_generation.clone(),
|
proxy_conn.node_generation.clone(),
|
||||||
proxy_conn.id,
|
proxy_conn.id,
|
||||||
proxy_stream_id,
|
proxy_stream_id,
|
||||||
*STREAM_INITIAL_WINDOW_BYTES,
|
settings.initial_stream_window_bytes,
|
||||||
));
|
);
|
||||||
|
local_stream.min_window_update_bytes = settings.min_window_update_bytes;
|
||||||
|
*local_stream.response_connection.get_mut() = Some(Arc::downgrade(&proxy_conn));
|
||||||
|
let local_stream = Arc::new(local_stream);
|
||||||
self.local_streams
|
self.local_streams
|
||||||
.insert(local_stream_id, local_stream.clone());
|
.insert(local_stream_id, local_stream.clone());
|
||||||
self.proxy_to_local
|
self.proxy_to_local
|
||||||
.insert((proxy_conn.id, proxy_stream_id), local_stream_id);
|
.insert((proxy_conn.id, proxy_stream_id), local_stream_id);
|
||||||
|
let mut pending_stream = PendingStreamGuard {
|
||||||
|
hub: self,
|
||||||
|
connection: &proxy_conn,
|
||||||
|
stream_id: local_stream_id,
|
||||||
|
committed: false,
|
||||||
|
};
|
||||||
|
|
||||||
let send_status = proxy_conn
|
let send_status = proxy_conn
|
||||||
.send_wait(
|
.send_wait(
|
||||||
@@ -1195,10 +1297,11 @@ impl HubRouter {
|
|||||||
"open_local_stream dispatched"
|
"open_local_stream dispatched"
|
||||||
);
|
);
|
||||||
match send_status {
|
match send_status {
|
||||||
SendStatus::Queued => Ok(local_stream),
|
SendStatus::Queued => {
|
||||||
|
pending_stream.committed = true;
|
||||||
|
Ok(local_stream)
|
||||||
|
}
|
||||||
SendStatus::Closed | SendStatus::Congested => {
|
SendStatus::Closed | SendStatus::Congested => {
|
||||||
self.cleanup_local_stream(local_stream_id);
|
|
||||||
proxy_conn.release_stream();
|
|
||||||
Err("proxy connection congested".to_string())
|
Err("proxy connection congested".to_string())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -1243,7 +1346,9 @@ impl HubRouter {
|
|||||||
.map(|entry| entry.value().clone())
|
.map(|entry| entry.value().clone())
|
||||||
.ok_or_else(|| "proxy connection unavailable".to_string())?;
|
.ok_or_else(|| "proxy connection unavailable".to_string())?;
|
||||||
|
|
||||||
let total_chunks = payload.len().div_ceil(MAX_REQUEST_BODY_FRAME_SIZE);
|
let chunk_size = MAX_REQUEST_BODY_FRAME_SIZE
|
||||||
|
.min(proxy_conn.settings.lock().initial_stream_window_bytes as usize);
|
||||||
|
let total_chunks = payload.len().div_ceil(chunk_size);
|
||||||
let result = if total_chunks == 0 {
|
let result = if total_chunks == 0 {
|
||||||
if end_stream {
|
if end_stream {
|
||||||
self.send_request_body_frame(&proxy_conn, &stream, &[], true)
|
self.send_request_body_frame(&proxy_conn, &stream, &[], true)
|
||||||
@@ -1252,7 +1357,7 @@ impl HubRouter {
|
|||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
for (index, chunk) in payload.chunks(MAX_REQUEST_BODY_FRAME_SIZE).enumerate() {
|
for (index, chunk) in payload.chunks(chunk_size).enumerate() {
|
||||||
let is_last_chunk = index + 1 == total_chunks;
|
let is_last_chunk = index + 1 == total_chunks;
|
||||||
if let Err(error) = self
|
if let Err(error) = self
|
||||||
.send_request_body_frame(
|
.send_request_body_frame(
|
||||||
@@ -1352,17 +1457,20 @@ impl HubRouter {
|
|||||||
} else {
|
} else {
|
||||||
protocol::encode_stream_error(stream.proxy_stream_id, reason)
|
protocol::encode_stream_error(stream.proxy_stream_id, reason)
|
||||||
};
|
};
|
||||||
let _ = pc.send(Message::Binary(frame.into()));
|
if pc.send(Message::Binary(frame.into())) != SendStatus::Queued {
|
||||||
|
pc.request_close();
|
||||||
|
}
|
||||||
}
|
}
|
||||||
stream.fail(reason.to_string());
|
stream.fail(reason.to_string());
|
||||||
}
|
}
|
||||||
|
|
||||||
fn cleanup_local_stream(&self, local_stream_id: u64) {
|
fn cleanup_local_stream(&self, local_stream_id: u64) -> bool {
|
||||||
let Some((_, stream)) = self.local_streams.remove(&local_stream_id) else {
|
let Some((_, stream)) = self.local_streams.remove(&local_stream_id) else {
|
||||||
return;
|
return false;
|
||||||
};
|
};
|
||||||
self.proxy_to_local
|
self.proxy_to_local
|
||||||
.remove(&(stream.proxy_conn_id, stream.proxy_stream_id));
|
.remove(&(stream.proxy_conn_id, stream.proxy_stream_id));
|
||||||
|
true
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn handle_proxy_frame(self: &Arc<Self>, proxy_conn_id: u64, data: &mut [u8]) {
|
pub async fn handle_proxy_frame(self: &Arc<Self>, proxy_conn_id: u64, data: &mut [u8]) {
|
||||||
@@ -1433,9 +1541,7 @@ impl HubRouter {
|
|||||||
.get(&proxy_conn_id)
|
.get(&proxy_conn_id)
|
||||||
.map(|entry| entry.value().clone());
|
.map(|entry| entry.value().clone());
|
||||||
if let Some(pc) = pc {
|
if let Some(pc) = pc {
|
||||||
let _ = pc
|
let _ = pc.send(Message::Binary(pong.into()));
|
||||||
.send_wait(Message::Binary(pong.into()), Duration::from_millis(250))
|
|
||||||
.await;
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
protocol::PONG => {}
|
protocol::PONG => {}
|
||||||
@@ -1509,11 +1615,36 @@ impl HubRouter {
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
protocol::SETTINGS => {
|
protocol::SETTINGS => {
|
||||||
debug!(
|
let settings = protocol::decode_payload_with_limit(
|
||||||
msg_type = header.msg_type,
|
data,
|
||||||
proxy_conn_id = proxy_conn_id,
|
&header,
|
||||||
"received tunnel protocol v3 SETTINGS from proxy"
|
MAX_TUNNEL_CONTROL_PAYLOAD_SIZE,
|
||||||
);
|
)
|
||||||
|
.ok()
|
||||||
|
.and_then(|payload| {
|
||||||
|
serde_json::from_slice::<protocol::SettingsPayload>(&payload).ok()
|
||||||
|
})
|
||||||
|
.filter(|settings| settings.is_valid());
|
||||||
|
if let Some(connection) = self.proxy_conns_by_id.get(&proxy_conn_id) {
|
||||||
|
if header.stream_id != 0 || header.flags != 0 {
|
||||||
|
connection.request_close();
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
let Some(settings) = settings else {
|
||||||
|
connection.request_close();
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
let local = local_settings();
|
||||||
|
let settings = settings
|
||||||
|
.negotiate(local.initial_stream_window_bytes, local.drain_deadline_ms);
|
||||||
|
let mut current = connection.settings.lock();
|
||||||
|
if connection.stream_count.load(Ordering::Acquire) > 0 && *current != settings {
|
||||||
|
drop(current);
|
||||||
|
connection.request_close();
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
*current = settings;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
protocol::WINDOW_UPDATE => {
|
protocol::WINDOW_UPDATE => {
|
||||||
self.handle_window_update(proxy_conn_id, header.stream_id, data, &header);
|
self.handle_window_update(proxy_conn_id, header.stream_id, data, &header);
|
||||||
@@ -1739,19 +1870,8 @@ impl HubRouter {
|
|||||||
None => return,
|
None => return,
|
||||||
};
|
};
|
||||||
|
|
||||||
let payload_len = payload.len();
|
if !stream.push_body_chunk(Bytes::from(payload)) {
|
||||||
if !stream.push_body_chunk(Bytes::from(payload)).await {
|
|
||||||
self.cancel_local_stream(local_id, "local relay response congested");
|
self.cancel_local_stream(local_id, "local relay response congested");
|
||||||
return;
|
|
||||||
}
|
|
||||||
|
|
||||||
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
|
|
||||||
if pc.protocol_version() >= 3 {
|
|
||||||
if let Some(delta) = stream.response_window_update_delta(payload_len) {
|
|
||||||
let frame = protocol::encode_window_update(header.stream_id, delta);
|
|
||||||
let _ = pc.send(Message::Binary(frame.into()));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -9,14 +9,13 @@ use axum::body::{Body, Bytes};
|
|||||||
use axum::extract::{ConnectInfo, Path, Request, State};
|
use axum::extract::{ConnectInfo, Path, Request, State};
|
||||||
use axum::http::{HeaderMap, HeaderName, HeaderValue, Response, StatusCode};
|
use axum::http::{HeaderMap, HeaderName, HeaderValue, Response, StatusCode};
|
||||||
use axum::response::IntoResponse;
|
use axum::response::IntoResponse;
|
||||||
use tokio::sync::mpsc;
|
|
||||||
use tracing::warn;
|
use tracing::warn;
|
||||||
|
|
||||||
use crate::api::response::apply_streaming_response_headers;
|
use crate::api::response::apply_streaming_response_headers;
|
||||||
use crate::headers::should_skip_response_header;
|
use crate::headers::should_skip_response_header;
|
||||||
use crate::maintenance::record_proxy_upgrade_traffic_success_for_generation;
|
use crate::maintenance::record_proxy_upgrade_traffic_success_for_generation;
|
||||||
|
|
||||||
use super::hub::{LocalBodyEvent, LocalStream};
|
use super::hub::{LocalBodyEvent, LocalBodyReceiver, LocalStream};
|
||||||
use super::protocol;
|
use super::protocol;
|
||||||
use super::{AppState, RelayRequestAuthenticated};
|
use super::{AppState, RelayRequestAuthenticated};
|
||||||
|
|
||||||
@@ -40,7 +39,7 @@ impl Drop for StreamGuard {
|
|||||||
pub(crate) struct DirectRelayResponse {
|
pub(crate) struct DirectRelayResponse {
|
||||||
status: u16,
|
status: u16,
|
||||||
headers: Vec<(String, String)>,
|
headers: Vec<(String, String)>,
|
||||||
body_rx: mpsc::Receiver<LocalBodyEvent>,
|
body_rx: LocalBodyReceiver,
|
||||||
request_guard: StreamGuard,
|
request_guard: StreamGuard,
|
||||||
_request_permit: Option<AdmissionPermit>,
|
_request_permit: Option<AdmissionPermit>,
|
||||||
}
|
}
|
||||||
@@ -55,10 +54,13 @@ impl DirectRelayResponse {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) async fn next_chunk(&mut self) -> Result<Option<Bytes>, String> {
|
pub(crate) async fn next_chunk(&mut self) -> Result<Option<Bytes>, String> {
|
||||||
|
if self.request_guard.finished {
|
||||||
|
return Ok(None);
|
||||||
|
}
|
||||||
let event = self.body_rx.recv().await;
|
let event = self.body_rx.recv().await;
|
||||||
match event {
|
match event {
|
||||||
Some(LocalBodyEvent::Chunk(chunk)) => Ok(Some(chunk)),
|
Some(LocalBodyEvent::Chunk(chunk)) => Ok(Some(chunk)),
|
||||||
Some(LocalBodyEvent::End) | None => {
|
Some(LocalBodyEvent::End) => {
|
||||||
self.request_guard.finished = true;
|
self.request_guard.finished = true;
|
||||||
Ok(None)
|
Ok(None)
|
||||||
}
|
}
|
||||||
@@ -66,6 +68,7 @@ impl DirectRelayResponse {
|
|||||||
self.request_guard.finished = true;
|
self.request_guard.finished = true;
|
||||||
Err(error)
|
Err(error)
|
||||||
}
|
}
|
||||||
|
None => Err("tunnel response ended without a terminal frame".to_string()),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -84,6 +87,11 @@ pub(crate) async fn open_direct_relay_stream(
|
|||||||
.open_authorized_local_stream(node_id, &meta)
|
.open_authorized_local_stream(node_id, &meta)
|
||||||
.await
|
.await
|
||||||
.map_err(|error| format!("connect: {error}"))?;
|
.map_err(|error| format!("connect: {error}"))?;
|
||||||
|
let request_guard = StreamGuard {
|
||||||
|
hub: state.hub.clone(),
|
||||||
|
stream_id: stream.id,
|
||||||
|
finished: false,
|
||||||
|
};
|
||||||
if let Err(error) = state
|
if let Err(error) = state
|
||||||
.hub
|
.hub
|
||||||
.push_local_request_body(stream.id, body, true)
|
.push_local_request_body(stream.id, body, true)
|
||||||
@@ -126,11 +134,7 @@ pub(crate) async fn open_direct_relay_stream(
|
|||||||
status: response_head.status,
|
status: response_head.status,
|
||||||
headers: response_head.headers,
|
headers: response_head.headers,
|
||||||
body_rx,
|
body_rx,
|
||||||
request_guard: StreamGuard {
|
request_guard,
|
||||||
hub: state.hub.clone(),
|
|
||||||
stream_id: stream.id,
|
|
||||||
finished: false,
|
|
||||||
},
|
|
||||||
_request_permit: request_permit,
|
_request_permit: request_permit,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -259,6 +263,11 @@ pub async fn relay_request(
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
let request_guard = StreamGuard {
|
||||||
|
hub: state.hub.clone(),
|
||||||
|
stream_id: stream.id,
|
||||||
|
finished: false,
|
||||||
|
};
|
||||||
let body_stream = match spool.body_stream().await {
|
let body_stream = match spool.body_stream().await {
|
||||||
Ok(stream) => stream,
|
Ok(stream) => stream,
|
||||||
Err(error) => {
|
Err(error) => {
|
||||||
@@ -306,12 +315,6 @@ pub async fn relay_request(
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
let request_guard = StreamGuard {
|
|
||||||
hub: state.hub.clone(),
|
|
||||||
stream_id: stream.id,
|
|
||||||
finished: false,
|
|
||||||
};
|
|
||||||
|
|
||||||
let wait_timeout = relay_header_timeout(&meta);
|
let wait_timeout = relay_header_timeout(&meta);
|
||||||
let response_head = match stream.wait_headers(wait_timeout).await {
|
let response_head = match stream.wait_headers(wait_timeout).await {
|
||||||
Ok(response) => response,
|
Ok(response) => response,
|
||||||
@@ -373,6 +376,9 @@ pub async fn relay_request(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
if !guard.finished {
|
||||||
|
yield Err(io::Error::other("tunnel response ended without a terminal frame"));
|
||||||
|
}
|
||||||
guard.finished = true;
|
guard.finished = true;
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -563,6 +569,100 @@ mod tests {
|
|||||||
request
|
request
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn cancelled_relays_reset_streams_during_upload_and_header_wait() {
|
||||||
|
for direct in [true, false] {
|
||||||
|
for during_upload in [true, false] {
|
||||||
|
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![
|
||||||
|
sample_connected_proxy_node("node-123"),
|
||||||
|
]));
|
||||||
|
let data = Arc::new(
|
||||||
|
GatewayDataState::with_proxy_node_repository_for_tests(repository)
|
||||||
|
.with_system_config_values_for_tests(
|
||||||
|
Vec::<(String, serde_json::Value)>::new(),
|
||||||
|
)
|
||||||
|
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||||
|
);
|
||||||
|
let state = test_app_state().with_data(data);
|
||||||
|
let (proxy_tx, mut proxy_rx) = aether_runtime::bounded_queue(8);
|
||||||
|
let (proxy_close_tx, _) = watch::channel(false);
|
||||||
|
let connection = Arc::new(
|
||||||
|
ProxyConn::new(
|
||||||
|
500,
|
||||||
|
"node-123".into(),
|
||||||
|
"Node 123".into(),
|
||||||
|
proxy_tx,
|
||||||
|
proxy_close_tx,
|
||||||
|
16,
|
||||||
|
3,
|
||||||
|
)
|
||||||
|
.with_tunnel_generation(LOCAL_TUNNEL_TEST_GENERATION.to_string())
|
||||||
|
.with_authenticated_key(LOCAL_TUNNEL_TEST_PSK.to_string())
|
||||||
|
.with_settings(protocol::SettingsPayload {
|
||||||
|
initial_stream_window_bytes: 128,
|
||||||
|
min_window_update_bytes: 32,
|
||||||
|
drain_deadline_ms: 1000,
|
||||||
|
}),
|
||||||
|
);
|
||||||
|
state.hub.register_proxy(Arc::clone(&connection));
|
||||||
|
let meta = protocol::RequestMeta {
|
||||||
|
provider_id: None,
|
||||||
|
endpoint_id: None,
|
||||||
|
key_id: None,
|
||||||
|
method: "POST".into(),
|
||||||
|
url: "https://example.com/".into(),
|
||||||
|
headers: HashMap::new(),
|
||||||
|
stream: true,
|
||||||
|
request_timeout_ms: None,
|
||||||
|
stream_first_byte_timeout_ms: None,
|
||||||
|
timeout: 30,
|
||||||
|
follow_redirects: None,
|
||||||
|
http1_only: false,
|
||||||
|
transport_profile: None,
|
||||||
|
};
|
||||||
|
let body = Bytes::from(vec![b'x'; if during_upload { 256 } else { 0 }]);
|
||||||
|
let relay = tokio::spawn(async move {
|
||||||
|
if direct {
|
||||||
|
let _response =
|
||||||
|
super::open_direct_relay_stream(&state, "node-123", meta, body)
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
} else {
|
||||||
|
let request =
|
||||||
|
authenticated_request(encode_relay_envelope(&meta, &body)).await;
|
||||||
|
let _response = relay_request(
|
||||||
|
Path("node-123".into()),
|
||||||
|
State(state),
|
||||||
|
ConnectInfo(SocketAddr::from(([127, 0, 0, 1], 4242))),
|
||||||
|
request,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
}
|
||||||
|
});
|
||||||
|
recv_tunnel_test_frame(&mut proxy_rx, "request headers").await;
|
||||||
|
recv_tunnel_test_frame(&mut proxy_rx, "request body").await;
|
||||||
|
relay.abort();
|
||||||
|
assert!(relay.await.unwrap_err().is_cancelled());
|
||||||
|
let Message::Binary(frame) = recv_tunnel_test_frame(&mut proxy_rx, "reset").await
|
||||||
|
else {
|
||||||
|
panic!("expected binary reset frame")
|
||||||
|
};
|
||||||
|
let frame = aether_contracts::tunnel::Frame::decode(frame).unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
frame.msg_type,
|
||||||
|
aether_contracts::tunnel::MsgType::ResetStream
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
connection
|
||||||
|
.stream_count
|
||||||
|
.load(std::sync::atomic::Ordering::Relaxed),
|
||||||
|
0
|
||||||
|
);
|
||||||
|
assert!(connection.is_available());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn relay_header_timeout_ignores_request_timeout_for_stream_requests() {
|
fn relay_header_timeout_ignores_request_timeout_for_stream_requests() {
|
||||||
let meta = protocol::RequestMeta {
|
let meta = protocol::RequestMeta {
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
mod body;
|
||||||
mod control_plane;
|
mod control_plane;
|
||||||
mod hub;
|
mod hub;
|
||||||
mod local_relay;
|
mod local_relay;
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ use aether_runtime::bounded_queue;
|
|||||||
use axum::extract::ws::{Message, WebSocket};
|
use axum::extract::ws::{Message, WebSocket};
|
||||||
use futures_util::{SinkExt, StreamExt};
|
use futures_util::{SinkExt, StreamExt};
|
||||||
use tokio::sync::watch;
|
use tokio::sync::watch;
|
||||||
|
use tokio::task::JoinSet;
|
||||||
use tracing::{debug, info, warn};
|
use tracing::{debug, info, warn};
|
||||||
|
|
||||||
use super::hub::{ConnConfig, HubRouter, ProxyConn, ProxyManagementTokenCredential, SendStatus};
|
use super::hub::{ConnConfig, HubRouter, ProxyConn, ProxyManagementTokenCredential, SendStatus};
|
||||||
@@ -84,6 +85,35 @@ pub async fn handle_proxy_connection(
|
|||||||
|
|
||||||
let (tx, mut rx) = bounded_queue::<Message>(cfg.outbound_queue_capacity);
|
let (tx, mut rx) = bounded_queue::<Message>(cfg.outbound_queue_capacity);
|
||||||
let (close_tx, mut close_rx) = watch::channel(false);
|
let (close_tx, mut close_rx) = watch::channel(false);
|
||||||
|
let settings = if protocol_version >= 3 {
|
||||||
|
let Some(settings) = read_proxy_settings(
|
||||||
|
&mut ws_tx,
|
||||||
|
&mut ws_rx,
|
||||||
|
security.as_deref(),
|
||||||
|
protocol_version,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
else {
|
||||||
|
warn!(conn_id, "proxy SETTINGS negotiation failed");
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
let local = super::hub::local_settings();
|
||||||
|
let negotiated =
|
||||||
|
settings.negotiate(local.initial_stream_window_bytes, local.drain_deadline_ms);
|
||||||
|
let message = Message::Binary(protocol::encode_settings(&negotiated).into());
|
||||||
|
let Ok(message) = encrypt_message(message, security.as_deref()) else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
if !matches!(
|
||||||
|
tokio::time::timeout(PROXY_HELLO_TIMEOUT, ws_tx.send(message)).await,
|
||||||
|
Ok(Ok(()))
|
||||||
|
) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
negotiated
|
||||||
|
} else {
|
||||||
|
super::hub::local_settings()
|
||||||
|
};
|
||||||
let conn = ProxyConn::new(
|
let conn = ProxyConn::new(
|
||||||
conn_id,
|
conn_id,
|
||||||
node_id.clone(),
|
node_id.clone(),
|
||||||
@@ -93,7 +123,8 @@ pub async fn handle_proxy_connection(
|
|||||||
max_streams,
|
max_streams,
|
||||||
protocol_version,
|
protocol_version,
|
||||||
)
|
)
|
||||||
.with_tunnel_generation(node_generation);
|
.with_tunnel_generation(node_generation)
|
||||||
|
.with_settings(settings);
|
||||||
let conn = match (security_key.clone(), management_token_credential) {
|
let conn = match (security_key.clone(), management_token_credential) {
|
||||||
(Some(key), None) => Arc::new(conn.with_authenticated_key(key)),
|
(Some(key), None) => Arc::new(conn.with_authenticated_key(key)),
|
||||||
(None, Some(credential)) => Arc::new(conn.with_management_token_credential(credential)),
|
(None, Some(credential)) => Arc::new(conn.with_management_token_credential(credential)),
|
||||||
@@ -416,19 +447,22 @@ async fn run_proxy_reader(
|
|||||||
let idle_enabled = !idle_timeout.is_zero();
|
let idle_enabled = !idle_timeout.is_zero();
|
||||||
let mut oversized_count = 0u32;
|
let mut oversized_count = 0u32;
|
||||||
let mut frames_received: u64 = 0;
|
let mut frames_received: u64 = 0;
|
||||||
|
let mut close_rx = conn.outbound.subscribe_close();
|
||||||
|
let mut heartbeats = JoinSet::new();
|
||||||
loop {
|
loop {
|
||||||
let msg = if idle_enabled {
|
if conn.outbound.is_closing() {
|
||||||
tokio::select! {
|
break;
|
||||||
msg = ws_rx.next() => msg,
|
}
|
||||||
_ = tokio::time::sleep(idle_timeout) => {
|
while heartbeats.try_join_next().is_some() {}
|
||||||
warn!(conn_id = conn.id, node_id = %conn.node_id, "proxy idle timeout");
|
let msg = tokio::select! {
|
||||||
let _ = conn.send(Message::Binary(protocol::encode_goaway().into()));
|
biased;
|
||||||
conn.request_close();
|
_ = close_rx.changed() => break,
|
||||||
break;
|
msg = ws_rx.next() => msg,
|
||||||
}
|
_ = tokio::time::sleep(idle_timeout), if idle_enabled => {
|
||||||
|
warn!(conn_id = conn.id, node_id = %conn.node_id, "proxy idle timeout");
|
||||||
|
conn.request_close();
|
||||||
|
break;
|
||||||
}
|
}
|
||||||
} else {
|
|
||||||
ws_rx.next().await
|
|
||||||
};
|
};
|
||||||
|
|
||||||
match msg {
|
match msg {
|
||||||
@@ -463,7 +497,27 @@ async fn run_proxy_reader(
|
|||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
||||||
hub.handle_proxy_frame(conn.id, &mut data).await;
|
let is_heartbeat = protocol::FrameHeader::parse(&data)
|
||||||
|
.is_some_and(|header| header.msg_type == protocol::HEARTBEAT_DATA);
|
||||||
|
if is_heartbeat {
|
||||||
|
if heartbeats.is_empty() {
|
||||||
|
let heartbeat_hub = Arc::clone(&hub);
|
||||||
|
let conn_id = conn.id;
|
||||||
|
heartbeats.spawn(async move {
|
||||||
|
if tokio::time::timeout(
|
||||||
|
Duration::from_secs(10),
|
||||||
|
heartbeat_hub.handle_proxy_frame(conn_id, &mut data),
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.is_err()
|
||||||
|
{
|
||||||
|
warn!(conn_id, "proxy heartbeat processing timed out");
|
||||||
|
}
|
||||||
|
});
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
hub.handle_proxy_frame(conn.id, &mut data).await;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
Some(Ok(Message::Close(_))) | None => {
|
Some(Ok(Message::Close(_))) | None => {
|
||||||
info!(
|
info!(
|
||||||
@@ -489,6 +543,56 @@ async fn run_proxy_reader(
|
|||||||
_ => {}
|
_ => {}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
heartbeats.shutdown().await;
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn read_proxy_settings(
|
||||||
|
ws_tx: &mut futures_util::stream::SplitSink<WebSocket, Message>,
|
||||||
|
ws_rx: &mut futures_util::stream::SplitStream<WebSocket>,
|
||||||
|
security: Option<&SecureFrameCodec>,
|
||||||
|
protocol_version: u8,
|
||||||
|
) -> Option<protocol::SettingsPayload> {
|
||||||
|
tokio::time::timeout(PROXY_HELLO_TIMEOUT, async {
|
||||||
|
let mut hello_received = security.is_some();
|
||||||
|
for _ in 0..MAX_PREAUTH_PINGS {
|
||||||
|
match ws_rx.next().await? {
|
||||||
|
Ok(Message::Binary(data)) => {
|
||||||
|
if data.len() > 256 * 1024 {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
let data = decrypt_message(data, security).ok()?;
|
||||||
|
let frame = Frame::decode(data.into()).ok()?;
|
||||||
|
if frame.stream_id != 0 || frame.flags != 0 {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
match frame.msg_type {
|
||||||
|
MsgType::Hello if !hello_received => {
|
||||||
|
let hello =
|
||||||
|
serde_json::from_slice::<HelloPayload>(&frame.payload).ok()?;
|
||||||
|
if hello.protocol_version != protocol_version {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
hello_received = true;
|
||||||
|
}
|
||||||
|
MsgType::Settings if hello_received => {
|
||||||
|
let settings =
|
||||||
|
serde_json::from_slice::<protocol::SettingsPayload>(&frame.payload)
|
||||||
|
.ok()?;
|
||||||
|
return settings.is_valid().then_some(settings);
|
||||||
|
}
|
||||||
|
_ => return None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Ok(Message::Ping(payload)) => ws_tx.send(Message::Pong(payload)).await.ok()?,
|
||||||
|
Ok(Message::Pong(_)) => {}
|
||||||
|
_ => return None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
None
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.ok()
|
||||||
|
.flatten()
|
||||||
}
|
}
|
||||||
|
|
||||||
fn encrypt_message(
|
fn encrypt_message(
|
||||||
@@ -523,6 +627,158 @@ fn decrypt_message(
|
|||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
|
#[cfg(feature = "testkit")]
|
||||||
|
#[tokio::test]
|
||||||
|
async fn slow_heartbeat_does_not_block_response_frames() {
|
||||||
|
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||||
|
use tokio_tungstenite::tungstenite::{client::IntoClientRequest, Message as ClientMessage};
|
||||||
|
|
||||||
|
let called = Arc::new(AtomicUsize::new(0));
|
||||||
|
let callback_called = Arc::clone(&called);
|
||||||
|
let control_plane = super::super::control_plane::ControlPlaneClient::local(
|
||||||
|
move |_, _| {
|
||||||
|
callback_called.fetch_add(1, Ordering::SeqCst);
|
||||||
|
Box::pin(std::future::pending())
|
||||||
|
},
|
||||||
|
|_, _, _, _| Box::pin(async { Ok(()) }),
|
||||||
|
);
|
||||||
|
let data = crate::data::GatewayDataState::with_tunnel_management_auth_for_testkit(
|
||||||
|
"heartbeat-test",
|
||||||
|
"heartbeat-generation",
|
||||||
|
"ae-tunnel-harness-management-token",
|
||||||
|
aether_crypto::DEVELOPMENT_ENCRYPTION_KEY,
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
let state = super::super::AppState::new(
|
||||||
|
control_plane,
|
||||||
|
ConnConfig {
|
||||||
|
ping_interval: Duration::from_secs(60),
|
||||||
|
idle_timeout: Duration::ZERO,
|
||||||
|
outbound_queue_capacity: 128,
|
||||||
|
},
|
||||||
|
16,
|
||||||
|
)
|
||||||
|
.with_data(Arc::new(data));
|
||||||
|
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||||
|
let address = listener.local_addr().unwrap();
|
||||||
|
let router = super::super::build_router_with_state(state.clone());
|
||||||
|
let server = tokio::spawn(async move {
|
||||||
|
axum::serve(
|
||||||
|
listener,
|
||||||
|
router.into_make_service_with_connect_info::<std::net::SocketAddr>(),
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
});
|
||||||
|
let mut request = format!("ws://{address}/api/internal/proxy-tunnel")
|
||||||
|
.into_client_request()
|
||||||
|
.unwrap();
|
||||||
|
let headers = request.headers_mut();
|
||||||
|
headers.insert("x-node-id", "heartbeat-test".parse().unwrap());
|
||||||
|
headers.insert(
|
||||||
|
aether_contracts::tunnel_security::TUNNEL_GENERATION_HEADER,
|
||||||
|
"heartbeat-generation".parse().unwrap(),
|
||||||
|
);
|
||||||
|
headers.insert(
|
||||||
|
aether_contracts::tunnel::TUNNEL_PROTOCOL_VERSION_HEADER,
|
||||||
|
"3".parse().unwrap(),
|
||||||
|
);
|
||||||
|
headers.insert(
|
||||||
|
"authorization",
|
||||||
|
"Bearer ae-tunnel-harness-management-token".parse().unwrap(),
|
||||||
|
);
|
||||||
|
let (mut websocket, _) = tokio_tungstenite::connect_async(request).await.unwrap();
|
||||||
|
let hello = HelloPayload {
|
||||||
|
protocol_version: 3,
|
||||||
|
capabilities: vec![],
|
||||||
|
session_id: None,
|
||||||
|
replica_id: None,
|
||||||
|
};
|
||||||
|
websocket
|
||||||
|
.send(ClientMessage::Binary(protocol::encode_hello(&hello).into()))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
websocket
|
||||||
|
.send(ClientMessage::Binary(
|
||||||
|
protocol::encode_settings(&super::super::hub::local_settings()).into(),
|
||||||
|
))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
let ClientMessage::Binary(settings) = websocket.next().await.unwrap().unwrap() else {
|
||||||
|
panic!("expected SETTINGS")
|
||||||
|
};
|
||||||
|
assert_eq!(Frame::decode(settings).unwrap().msg_type, MsgType::Settings);
|
||||||
|
tokio::time::timeout(Duration::from_secs(1), async {
|
||||||
|
while !state.hub.has_local_proxy("heartbeat-test") {
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
let meta: protocol::RequestMeta = serde_json::from_value(serde_json::json!({
|
||||||
|
"method": "GET", "url": "https://example.com", "headers": {}, "stream": true, "timeout": 10
|
||||||
|
})).unwrap();
|
||||||
|
let stream = state
|
||||||
|
.hub
|
||||||
|
.open_local_stream("heartbeat-test", &meta)
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
let ClientMessage::Binary(request) = websocket.next().await.unwrap().unwrap() else {
|
||||||
|
panic!("expected request headers")
|
||||||
|
};
|
||||||
|
let stream_id = Frame::decode(request).unwrap().stream_id;
|
||||||
|
let heartbeat = Frame::control(
|
||||||
|
MsgType::HeartbeatData,
|
||||||
|
serde_json::to_vec(&serde_json::json!({"node_id": "heartbeat-test"})).unwrap(),
|
||||||
|
);
|
||||||
|
websocket
|
||||||
|
.send(ClientMessage::Binary(heartbeat.encode()))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
tokio::time::timeout(Duration::from_secs(1), async {
|
||||||
|
while called.load(Ordering::SeqCst) == 0 {
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
for _ in 0..8 {
|
||||||
|
websocket
|
||||||
|
.send(ClientMessage::Binary(heartbeat.encode()))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
}
|
||||||
|
let response = Frame::new(
|
||||||
|
stream_id,
|
||||||
|
MsgType::ResponseHeaders,
|
||||||
|
0,
|
||||||
|
serde_json::to_vec(&serde_json::json!({"status": 200, "headers": []})).unwrap(),
|
||||||
|
);
|
||||||
|
websocket
|
||||||
|
.send(ClientMessage::Binary(response.encode()))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
stream
|
||||||
|
.wait_headers(Duration::from_secs(1))
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.status,
|
||||||
|
200
|
||||||
|
);
|
||||||
|
assert_eq!(called.load(Ordering::SeqCst), 1);
|
||||||
|
state.hub.request_close_all_proxies();
|
||||||
|
tokio::time::timeout(Duration::from_secs(1), async {
|
||||||
|
while state.hub.has_local_proxy("heartbeat-test") {
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
server.abort();
|
||||||
|
let _ = server.await;
|
||||||
|
}
|
||||||
|
|
||||||
use super::*;
|
use super::*;
|
||||||
|
|
||||||
const KEY: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=";
|
const KEY: &str = "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=";
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
[package]
|
[package]
|
||||||
name = "aether-tunnel"
|
name = "aether-tunnel"
|
||||||
version = "0.3.16"
|
version = "0.3.17"
|
||||||
edition = "2021"
|
edition = "2021"
|
||||||
description = "Tunnel agent for Aether"
|
description = "Tunnel agent for Aether"
|
||||||
|
|
||||||
@@ -47,3 +47,4 @@ uuid.workspace = true
|
|||||||
|
|
||||||
[dev-dependencies]
|
[dev-dependencies]
|
||||||
aether-gateway = { workspace = true, features = ["testkit"] }
|
aether-gateway = { workspace = true, features = ["testkit"] }
|
||||||
|
tokio = { version = "1", features = ["test-util"] }
|
||||||
|
|||||||
@@ -4,6 +4,15 @@ Aether Tunnel 代理节点,部署在海外 VPS 上,通过 WebSocket 隧道
|
|||||||
|
|
||||||
Tunnel 模式下代理节点**无需对外监听端口**,仅需出站连接到 Aether 服务器。
|
Tunnel 模式下代理节点**无需对外监听端口**,仅需出站连接到 Aether 服务器。
|
||||||
|
|
||||||
|
## 流式传输与升级注意事项
|
||||||
|
|
||||||
|
- 协议 v3 连接在 `HELLO` / `SETTINGS` 协商后才接收业务请求。实际双向流窗口取 gateway 与 agent 配置的较小值,信用更新阈值不超过该窗口的四分之一;单帧也不会超过协商窗口。
|
||||||
|
- 响应缓冲按字节限额并合并小帧,结束和错误状态独立保存。慢消费者不会阻塞同一隧道其他流的读取;超出窗口或缓冲预算的流会被明确终止,不会静默截断。
|
||||||
|
- 信用更新在消费数据后可靠入队;启用重定向重放时,进入有界重放缓存也视为请求体消费。持续无法投递关键控制帧时会关闭连接并向在途请求报告错误。
|
||||||
|
- 客户端取消会终止对应上游请求,断连会回收 session 的 writer、heartbeat 和请求任务。正常 drain 在配置期限内继续处理已有流,期限到达后终止残留任务。
|
||||||
|
- 建议先升级 gateway,再升级 agent。既有 v3 agent 已发送 `HELLO` / `SETTINGS`,可连接新 gateway;自定义 v3 节点必须完成这两步握手。协议 v1/v2 保留旧握手。与旧 gateway 混用时应保持默认窗口配置,不能依赖旧 gateway 应用新的窗口协商。
|
||||||
|
- 自动重连恢复后续请求,不会自动续传已经输出的 SSE,也不会无条件重放已经发送的请求。
|
||||||
|
|
||||||
## 安装
|
## 安装
|
||||||
|
|
||||||
`aether-tunnel` 会根据宿主机自动选择服务管理器:
|
`aether-tunnel` 会根据宿主机自动选择服务管理器:
|
||||||
|
|||||||
@@ -764,6 +764,13 @@ impl Config {
|
|||||||
if self.tunnel_stream_initial_window_bytes == 0 {
|
if self.tunnel_stream_initial_window_bytes == 0 {
|
||||||
anyhow::bail!("tunnel_stream_initial_window_bytes must be > 0");
|
anyhow::bail!("tunnel_stream_initial_window_bytes must be > 0");
|
||||||
}
|
}
|
||||||
|
if u64::from(self.tunnel_stream_initial_window_bytes)
|
||||||
|
> aether_contracts::tunnel::MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES as u64
|
||||||
|
{
|
||||||
|
anyhow::bail!(
|
||||||
|
"tunnel_stream_initial_window_bytes exceeds the maximum tunnel payload size"
|
||||||
|
);
|
||||||
|
}
|
||||||
if self.tunnel_drain_deadline_ms == 0 {
|
if self.tunnel_drain_deadline_ms == 0 {
|
||||||
anyhow::bail!("tunnel_drain_deadline_ms must be > 0");
|
anyhow::bail!("tunnel_drain_deadline_ms must be > 0");
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -203,19 +203,34 @@ pub async fn connect_and_run(
|
|||||||
let (ws_sink, ws_read) = futures_util::StreamExt::split(ws_stream);
|
let (ws_sink, ws_read) = futures_util::StreamExt::split(ws_stream);
|
||||||
|
|
||||||
// Spawn writer task (with WebSocket ping keepalive)
|
// Spawn writer task (with WebSocket ping keepalive)
|
||||||
let (frame_tx, mut writer_handle) = writer::spawn_writer_with_metrics_and_security(
|
let (frame_tx, writer_handle) = writer::spawn_writer_with_metrics_and_security(
|
||||||
ws_sink,
|
ws_sink,
|
||||||
ping_interval,
|
ping_interval,
|
||||||
Some(Arc::clone(&server.tunnel_metrics)),
|
Some(Arc::clone(&server.tunnel_metrics)),
|
||||||
security.clone(),
|
security.clone(),
|
||||||
);
|
);
|
||||||
|
let mut writer_handle = super::task::SessionTask::new(writer_handle);
|
||||||
send_protocol_v3_hello(&frame_tx, &security_session, state).await;
|
send_protocol_v3_hello(&frame_tx, &security_session, state).await;
|
||||||
let drain_signal = spawn_drain_signal(
|
let (session_drain_tx, session_drain_rx) = watch::channel(*drain.borrow());
|
||||||
|
let forward_drain_tx = session_drain_tx.clone();
|
||||||
|
let mut external_drain = drain;
|
||||||
|
let forward_drain = super::task::SessionTask::new(tokio::spawn(async move {
|
||||||
|
loop {
|
||||||
|
if *external_drain.borrow() {
|
||||||
|
let _ = forward_drain_tx.send(true);
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
if external_drain.changed().await.is_err() {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
let drain_signal = super::task::SessionTask::new(spawn_drain_signal(
|
||||||
conn_idx,
|
conn_idx,
|
||||||
frame_tx.clone(),
|
frame_tx.clone(),
|
||||||
drain.clone(),
|
session_drain_rx.clone(),
|
||||||
state.config.tunnel_drain_deadline_ms,
|
state.config.tunnel_drain_deadline_ms,
|
||||||
);
|
));
|
||||||
|
|
||||||
// Spawn heartbeat task (only for primary connection to avoid
|
// Spawn heartbeat task (only for primary connection to avoid
|
||||||
// resetting shared atomic metrics via swap(0))
|
// resetting shared atomic metrics via swap(0))
|
||||||
@@ -237,16 +252,19 @@ pub async fn connect_and_run(
|
|||||||
// ensures we detect this and trigger a reconnect promptly.
|
// ensures we detect this and trigger a reconnect promptly.
|
||||||
let state_clone = Arc::clone(state);
|
let state_clone = Arc::clone(state);
|
||||||
let server_clone = Arc::clone(server);
|
let server_clone = Arc::clone(server);
|
||||||
let outcome = tokio::select! {
|
let outcome = {
|
||||||
result = dispatcher::run_with_security(
|
let dispatch = dispatcher::run_with_security(
|
||||||
state_clone,
|
state_clone,
|
||||||
server_clone,
|
server_clone,
|
||||||
ws_read,
|
ws_read,
|
||||||
frame_tx.clone(),
|
frame_tx.clone(),
|
||||||
hb_handle,
|
hb_handle,
|
||||||
drain.clone(),
|
session_drain_rx,
|
||||||
security.clone(),
|
security.clone(),
|
||||||
) => {
|
);
|
||||||
|
tokio::pin!(dispatch);
|
||||||
|
tokio::select! {
|
||||||
|
result = &mut dispatch => {
|
||||||
match result {
|
match result {
|
||||||
Ok(()) => Ok(TunnelOutcome::Disconnected),
|
Ok(()) => Ok(TunnelOutcome::Disconnected),
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
@@ -258,6 +276,8 @@ pub async fn connect_and_run(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
writer_result = &mut writer_handle => {
|
writer_result = &mut writer_handle => {
|
||||||
|
frame_tx.close();
|
||||||
|
let _ = tokio::time::timeout(Duration::from_secs(1), &mut dispatch).await;
|
||||||
match writer_result {
|
match writer_result {
|
||||||
Ok(()) => warn!("writer task exited normally, triggering reconnect"),
|
Ok(()) => warn!("writer task exited normally, triggering reconnect"),
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
@@ -278,24 +298,37 @@ pub async fn connect_and_run(
|
|||||||
}
|
}
|
||||||
_ = shutdown.changed() => {
|
_ = shutdown.changed() => {
|
||||||
debug!("shutdown during tunnel dispatch");
|
debug!("shutdown during tunnel dispatch");
|
||||||
|
let _ = session_drain_tx.send(true);
|
||||||
|
let deadline = Duration::from_millis(state.config.tunnel_drain_deadline_ms).saturating_add(Duration::from_secs(1));
|
||||||
|
let _ = tokio::time::timeout(deadline, &mut dispatch).await;
|
||||||
Ok(TunnelOutcome::Shutdown)
|
Ok(TunnelOutcome::Shutdown)
|
||||||
}
|
}
|
||||||
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
// Drop our sender; the writer will exit once all stream handler clones
|
// Drop our sender; the writer will exit once all stream handler clones
|
||||||
// are also dropped (i.e. after they finish their in-flight work).
|
// are also dropped (i.e. after they finish their in-flight work).
|
||||||
drop(frame_tx);
|
drop(frame_tx);
|
||||||
|
forward_drain.abort();
|
||||||
|
let _ = forward_drain.await;
|
||||||
if !drain_signal.is_finished() {
|
if !drain_signal.is_finished() {
|
||||||
drain_signal.abort();
|
drain_signal.abort();
|
||||||
let _ = drain_signal.await;
|
let _ = drain_signal.await;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Wait for the writer task to finish with a generous timeout — the
|
|
||||||
// dispatcher already waits up to 30s for stream handlers, so 35s here
|
|
||||||
// covers that plus a small margin.
|
|
||||||
// Skip if the writer already exited (the select branch that fired).
|
|
||||||
if !writer_handle.is_finished() {
|
if !writer_handle.is_finished() {
|
||||||
let _ = tokio::time::timeout(Duration::from_secs(35), writer_handle).await;
|
let flush_timeout = if *session_drain_tx.borrow() {
|
||||||
|
Duration::from_millis(state.config.tunnel_drain_deadline_ms)
|
||||||
|
} else {
|
||||||
|
Duration::from_secs(1)
|
||||||
|
};
|
||||||
|
if tokio::time::timeout(flush_timeout, &mut writer_handle)
|
||||||
|
.await
|
||||||
|
.is_err()
|
||||||
|
{
|
||||||
|
writer_handle.abort();
|
||||||
|
let _ = writer_handle.await;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
let connected_for = connected_at.elapsed();
|
let connected_for = connected_at.elapsed();
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ use std::time::Duration;
|
|||||||
use bytes::Bytes;
|
use bytes::Bytes;
|
||||||
use futures_util::StreamExt;
|
use futures_util::StreamExt;
|
||||||
use tokio::sync::{mpsc, watch, OwnedSemaphorePermit, Semaphore};
|
use tokio::sync::{mpsc, watch, OwnedSemaphorePermit, Semaphore};
|
||||||
use tokio::task::JoinHandle;
|
use tokio::task::{AbortHandle, JoinSet};
|
||||||
use tokio_tungstenite::tungstenite::Message;
|
use tokio_tungstenite::tungstenite::Message;
|
||||||
use tracing::{debug, error, info, warn};
|
use tracing::{debug, error, info, warn};
|
||||||
|
|
||||||
@@ -41,13 +41,25 @@ impl AsRef<[u8]> for BudgetedFramePayload {
|
|||||||
enum StreamDispatchStatus {
|
enum StreamDispatchStatus {
|
||||||
Delivered,
|
Delivered,
|
||||||
Closed,
|
Closed,
|
||||||
TimedOut,
|
Congested,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Clone)]
|
#[derive(Clone)]
|
||||||
struct StreamDispatchTarget {
|
struct StreamDispatchTarget {
|
||||||
body_tx: mpsc::Sender<Frame>,
|
body_tx: mpsc::Sender<Frame>,
|
||||||
response_window: Arc<StreamSendWindow>,
|
response_window: Arc<StreamSendWindow>,
|
||||||
|
handler: Option<AbortHandle>,
|
||||||
|
}
|
||||||
|
|
||||||
|
struct StreamCompletion {
|
||||||
|
stream_id: u32,
|
||||||
|
finished_tx: mpsc::UnboundedSender<u32>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Drop for StreamCompletion {
|
||||||
|
fn drop(&mut self) {
|
||||||
|
let _ = self.finished_tx.send(self.stream_id);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// A request stream is identified by a non-zero id and may only be opened
|
/// A request stream is identified by a non-zero id and may only be opened
|
||||||
@@ -109,7 +121,7 @@ where
|
|||||||
// reopen the same id and bypass the stream admission limit.
|
// reopen the same id and bypass the stream admission limit.
|
||||||
let mut active_handler_ids: HashSet<u32> = HashSet::new();
|
let mut active_handler_ids: HashSet<u32> = HashSet::new();
|
||||||
// Track spawned stream handlers so we can wait for them on shutdown
|
// Track spawned stream handlers so we can wait for them on shutdown
|
||||||
let mut handler_handles: Vec<JoinHandle<()>> = Vec::new();
|
let mut handler_handles = JoinSet::new();
|
||||||
let (handler_finished_tx, mut handler_finished_rx) = mpsc::unbounded_channel::<u32>();
|
let (handler_finished_tx, mut handler_finished_rx) = mpsc::unbounded_channel::<u32>();
|
||||||
let max_streams = state.config.tunnel_max_streams.unwrap_or(128) as usize;
|
let max_streams = state.config.tunnel_max_streams.unwrap_or(128) as usize;
|
||||||
let mut frames_since_cleanup: u32 = 0;
|
let mut frames_since_cleanup: u32 = 0;
|
||||||
@@ -121,30 +133,48 @@ where
|
|||||||
// Track last time we received any data to detect stale connections
|
// Track last time we received any data to detect stale connections
|
||||||
let mut last_data_at = tokio::time::Instant::now();
|
let mut last_data_at = tokio::time::Instant::now();
|
||||||
let mut draining = *drain.borrow();
|
let mut draining = *drain.borrow();
|
||||||
|
let mut drain_open = true;
|
||||||
|
let mut drain_deadline = draining.then(|| {
|
||||||
|
tokio::time::Instant::now() + Duration::from_millis(state.config.tunnel_drain_deadline_ms)
|
||||||
|
});
|
||||||
|
let mut initial_window_bytes = state.config.tunnel_stream_initial_window_bytes;
|
||||||
|
let mut close_rx = frame_tx.subscribe_close();
|
||||||
|
|
||||||
let read_err = loop {
|
let read_err = loop {
|
||||||
|
if *close_rx.borrow() {
|
||||||
|
break None;
|
||||||
|
}
|
||||||
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
||||||
info!("tunnel drained after in-flight streams completed");
|
info!("tunnel drained after in-flight streams completed");
|
||||||
break None;
|
break None;
|
||||||
}
|
}
|
||||||
|
|
||||||
let msg_result = tokio::select! {
|
let msg_result = tokio::select! {
|
||||||
|
_ = close_rx.changed() => break None,
|
||||||
msg = ws_stream.next() => {
|
msg = ws_stream.next() => {
|
||||||
match msg {
|
match msg {
|
||||||
Some(r) => r,
|
Some(r) => r,
|
||||||
None => break None,
|
None => break None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
changed = drain.changed() => {
|
changed = drain.changed(), if drain_open => {
|
||||||
if changed.is_err() {
|
if changed.is_err() {
|
||||||
|
drain_open = false;
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
if *drain.borrow() {
|
if *drain.borrow() {
|
||||||
info!("tunnel drain requested, waiting for in-flight streams");
|
info!("tunnel drain requested, waiting for in-flight streams");
|
||||||
draining = true;
|
draining = true;
|
||||||
|
drain_deadline.get_or_insert_with(|| tokio::time::Instant::now() + Duration::from_millis(state.config.tunnel_drain_deadline_ms));
|
||||||
}
|
}
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
_ = async {
|
||||||
|
match drain_deadline {
|
||||||
|
Some(deadline) => tokio::time::sleep_until(deadline).await,
|
||||||
|
None => std::future::pending().await,
|
||||||
|
}
|
||||||
|
} => break None,
|
||||||
finished = handler_finished_rx.recv() => {
|
finished = handler_finished_rx.recv() => {
|
||||||
if let Some(stream_id) = finished {
|
if let Some(stream_id) = finished {
|
||||||
active_handler_ids.remove(&stream_id);
|
active_handler_ids.remove(&stream_id);
|
||||||
@@ -238,20 +268,7 @@ where
|
|||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
if draining {
|
if draining {
|
||||||
if frame_tx
|
try_send_stream_error(&frame_tx, frame.stream_id, "tunnel draining");
|
||||||
.try_send(Frame::new(
|
|
||||||
frame.stream_id,
|
|
||||||
MsgType::StreamError,
|
|
||||||
0,
|
|
||||||
Bytes::from("tunnel draining"),
|
|
||||||
))
|
|
||||||
.is_err()
|
|
||||||
{
|
|
||||||
warn!(
|
|
||||||
stream_id = frame.stream_id,
|
|
||||||
"writer channel full, StreamError dropped during drain"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -263,6 +280,11 @@ where
|
|||||||
Ok(p) => p,
|
Ok(p) => p,
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
warn!(stream_id = frame.stream_id, error = %e, "frame decompress failed");
|
warn!(stream_id = frame.stream_id, error = %e, "frame decompress failed");
|
||||||
|
try_send_stream_error(
|
||||||
|
&frame_tx,
|
||||||
|
frame.stream_id,
|
||||||
|
"invalid request metadata",
|
||||||
|
);
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
@@ -270,21 +292,11 @@ where
|
|||||||
Ok(m) => m,
|
Ok(m) => m,
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
warn!(stream_id = frame.stream_id, error = %e, "invalid request metadata");
|
warn!(stream_id = frame.stream_id, error = %e, "invalid request metadata");
|
||||||
// Use try_send to avoid blocking the read loop
|
try_send_stream_error(
|
||||||
if frame_tx
|
&frame_tx,
|
||||||
.try_send(Frame::new(
|
frame.stream_id,
|
||||||
frame.stream_id,
|
"invalid request metadata",
|
||||||
MsgType::StreamError,
|
);
|
||||||
0,
|
|
||||||
Bytes::from(format!("invalid request metadata: {e}")),
|
|
||||||
))
|
|
||||||
.is_err()
|
|
||||||
{
|
|
||||||
warn!(
|
|
||||||
stream_id = frame.stream_id,
|
|
||||||
"writer channel full, StreamError dropped"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
@@ -294,33 +306,27 @@ where
|
|||||||
stream_id = frame.stream_id,
|
stream_id = frame.stream_id,
|
||||||
"max concurrent streams reached"
|
"max concurrent streams reached"
|
||||||
);
|
);
|
||||||
if frame_tx
|
try_send_stream_error(
|
||||||
.try_send(Frame::new(
|
&frame_tx,
|
||||||
frame.stream_id,
|
frame.stream_id,
|
||||||
MsgType::StreamError,
|
"max concurrent streams reached",
|
||||||
0,
|
);
|
||||||
Bytes::from("max concurrent streams reached"),
|
|
||||||
))
|
|
||||||
.is_err()
|
|
||||||
{
|
|
||||||
warn!(
|
|
||||||
stream_id = frame.stream_id,
|
|
||||||
"writer channel full, StreamError dropped"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Create body channel and spawn handler
|
// Create body channel and spawn handler
|
||||||
let (body_tx, body_rx) = mpsc::channel::<Frame>(64);
|
let body_capacity = (initial_window_bytes as usize)
|
||||||
let response_window = Arc::new(StreamSendWindow::new(
|
.div_ceil(32 * 1024)
|
||||||
state.config.tunnel_stream_initial_window_bytes,
|
.saturating_add(1)
|
||||||
));
|
.max(64);
|
||||||
|
let (body_tx, body_rx) = mpsc::channel::<Frame>(body_capacity);
|
||||||
|
let response_window = Arc::new(StreamSendWindow::new(initial_window_bytes));
|
||||||
streams.insert(
|
streams.insert(
|
||||||
frame.stream_id,
|
frame.stream_id,
|
||||||
StreamDispatchTarget {
|
StreamDispatchTarget {
|
||||||
body_tx,
|
body_tx,
|
||||||
response_window: Arc::clone(&response_window),
|
response_window: Arc::clone(&response_window),
|
||||||
|
handler: None,
|
||||||
},
|
},
|
||||||
);
|
);
|
||||||
active_handler_ids.insert(frame.stream_id);
|
active_handler_ids.insert(frame.stream_id);
|
||||||
@@ -329,9 +335,13 @@ where
|
|||||||
let state_clone = Arc::clone(&state);
|
let state_clone = Arc::clone(&state);
|
||||||
let server_clone = Arc::clone(&server);
|
let server_clone = Arc::clone(&server);
|
||||||
let tx_clone = frame_tx.clone();
|
let tx_clone = frame_tx.clone();
|
||||||
let finished_tx = handler_finished_tx.clone();
|
|
||||||
let sid = frame.stream_id;
|
let sid = frame.stream_id;
|
||||||
let handle = tokio::spawn(async move {
|
let completion = StreamCompletion {
|
||||||
|
stream_id: sid,
|
||||||
|
finished_tx: handler_finished_tx.clone(),
|
||||||
|
};
|
||||||
|
let handle = handler_handles.spawn(async move {
|
||||||
|
let _completion = completion;
|
||||||
stream_handler::handle_stream(
|
stream_handler::handle_stream(
|
||||||
state_clone,
|
state_clone,
|
||||||
server_clone,
|
server_clone,
|
||||||
@@ -342,9 +352,8 @@ where
|
|||||||
response_window,
|
response_window,
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
let _ = finished_tx.send(sid);
|
|
||||||
});
|
});
|
||||||
handler_handles.push(handle);
|
streams.get_mut(&sid).expect("new stream exists").handler = Some(handle);
|
||||||
|
|
||||||
if request_headers_end_stream {
|
if request_headers_end_stream {
|
||||||
if let Some(target) = streams.get(&sid) {
|
if let Some(target) = streams.get(&sid) {
|
||||||
@@ -365,19 +374,21 @@ where
|
|||||||
let is_end = frame.is_end_stream();
|
let is_end = frame.is_end_stream();
|
||||||
let sid = frame.stream_id;
|
let sid = frame.stream_id;
|
||||||
let dispatch = dispatch_stream_frame(&target.body_tx, frame).await;
|
let dispatch = dispatch_stream_frame(&target.body_tx, frame).await;
|
||||||
if dispatch != StreamDispatchStatus::Delivered {
|
if dispatch == StreamDispatchStatus::Congested {
|
||||||
streams.remove(&sid);
|
if let Some(target) = streams.remove(&sid) {
|
||||||
if dispatch == StreamDispatchStatus::TimedOut {
|
if let Some(handler) = target.handler {
|
||||||
server.tunnel_metrics.record_error(
|
handler.abort();
|
||||||
"stream_dispatch_timeout",
|
}
|
||||||
&format!("request body dispatch timed out for stream {}", sid),
|
|
||||||
);
|
|
||||||
try_send_stream_error(
|
|
||||||
&frame_tx,
|
|
||||||
sid,
|
|
||||||
"tunnel request body dispatch stalled",
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
|
server.tunnel_metrics.record_error(
|
||||||
|
"stream_dispatch_timeout",
|
||||||
|
&format!("request body dispatch congested for stream {}", sid),
|
||||||
|
);
|
||||||
|
try_send_stream_error(
|
||||||
|
&frame_tx,
|
||||||
|
sid,
|
||||||
|
"tunnel request body dispatch stalled",
|
||||||
|
);
|
||||||
if is_end && draining && streams.is_empty() && active_handler_ids.is_empty()
|
if is_end && draining && streams.is_empty() && active_handler_ids.is_empty()
|
||||||
{
|
{
|
||||||
info!("tunnel drained after request body completion");
|
info!("tunnel drained after request body completion");
|
||||||
@@ -387,10 +398,29 @@ where
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
MsgType::StreamEnd | MsgType::StreamError | MsgType::ResetStream => {
|
MsgType::StreamEnd => {
|
||||||
|
if let Some(target) = streams.get(&frame.stream_id) {
|
||||||
|
if dispatch_stream_frame(&target.body_tx, frame.clone()).await
|
||||||
|
== StreamDispatchStatus::Congested
|
||||||
|
{
|
||||||
|
if let Some(handler) = &target.handler {
|
||||||
|
handler.abort();
|
||||||
|
}
|
||||||
|
try_send_stream_error(
|
||||||
|
&frame_tx,
|
||||||
|
frame.stream_id,
|
||||||
|
"tunnel request body dispatch stalled",
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
MsgType::StreamError | MsgType::ResetStream => {
|
||||||
// Client-side cancellation or end
|
// Client-side cancellation or end
|
||||||
if let Some(target) = streams.remove(&frame.stream_id) {
|
if let Some(target) = streams.remove(&frame.stream_id) {
|
||||||
let _ = dispatch_stream_frame(&target.body_tx, frame).await;
|
if let Some(handler) = target.handler {
|
||||||
|
handler.abort();
|
||||||
|
}
|
||||||
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
||||||
info!("tunnel drained after stream termination");
|
info!("tunnel drained after stream termination");
|
||||||
break None;
|
break None;
|
||||||
@@ -409,12 +439,22 @@ where
|
|||||||
}
|
}
|
||||||
|
|
||||||
MsgType::HeartbeatAck => {
|
MsgType::HeartbeatAck => {
|
||||||
heartbeat.on_ack(frame.payload).await;
|
heartbeat.on_ack(frame.payload);
|
||||||
}
|
}
|
||||||
|
|
||||||
MsgType::GoAway => {
|
MsgType::GoAway => {
|
||||||
info!("received GOAWAY");
|
info!("received GOAWAY");
|
||||||
break None;
|
draining = true;
|
||||||
|
let deadline_ms =
|
||||||
|
serde_json::from_slice::<aether_contracts::tunnel::GoAwayPayload>(
|
||||||
|
&frame.payload,
|
||||||
|
)
|
||||||
|
.map(|payload| payload.drain_deadline_ms)
|
||||||
|
.unwrap_or(state.config.tunnel_drain_deadline_ms)
|
||||||
|
.min(state.config.tunnel_drain_deadline_ms);
|
||||||
|
drain_deadline.get_or_insert_with(|| {
|
||||||
|
tokio::time::Instant::now() + Duration::from_millis(deadline_ms)
|
||||||
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
MsgType::WindowUpdate => {
|
MsgType::WindowUpdate => {
|
||||||
@@ -433,7 +473,31 @@ where
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
MsgType::Hello | MsgType::Settings | MsgType::LoadReport => {
|
MsgType::Settings => {
|
||||||
|
if frame.stream_id != 0 || frame.flags != 0 {
|
||||||
|
break None;
|
||||||
|
}
|
||||||
|
let settings = serde_json::from_slice::<aether_contracts::tunnel::SettingsPayload>(
|
||||||
|
&frame.payload,
|
||||||
|
)
|
||||||
|
.ok()
|
||||||
|
.filter(|settings| settings.is_valid());
|
||||||
|
let Some(settings) = settings else {
|
||||||
|
warn!("invalid tunnel SETTINGS");
|
||||||
|
break None;
|
||||||
|
};
|
||||||
|
if !streams.is_empty()
|
||||||
|
&& settings.initial_stream_window_bytes != initial_window_bytes
|
||||||
|
{
|
||||||
|
warn!("tunnel SETTINGS changed with active streams");
|
||||||
|
break None;
|
||||||
|
}
|
||||||
|
initial_window_bytes = settings
|
||||||
|
.initial_stream_window_bytes
|
||||||
|
.min(state.config.tunnel_stream_initial_window_bytes);
|
||||||
|
}
|
||||||
|
|
||||||
|
MsgType::Hello | MsgType::LoadReport => {
|
||||||
debug!(
|
debug!(
|
||||||
msg_type = ?frame.msg_type,
|
msg_type = ?frame.msg_type,
|
||||||
stream_id = frame.stream_id,
|
stream_id = frame.stream_id,
|
||||||
@@ -455,7 +519,7 @@ where
|
|||||||
// Trigger every 64 frames OR when the count exceeds max_streams.
|
// Trigger every 64 frames OR when the count exceeds max_streams.
|
||||||
frames_since_cleanup += 1;
|
frames_since_cleanup += 1;
|
||||||
if frames_since_cleanup >= 64 || handler_handles.len() > max_streams {
|
if frames_since_cleanup >= 64 || handler_handles.len() > max_streams {
|
||||||
handler_handles.retain(|h| !h.is_finished());
|
while handler_handles.try_join_next().is_some() {}
|
||||||
frames_since_cleanup = 0;
|
frames_since_cleanup = 0;
|
||||||
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
if draining && streams.is_empty() && active_handler_ids.is_empty() {
|
||||||
info!("tunnel drained after cleanup");
|
info!("tunnel drained after cleanup");
|
||||||
@@ -467,9 +531,7 @@ where
|
|||||||
// Drop body senders so stream handlers waiting on body_rx will unblock
|
// Drop body senders so stream handlers waiting on body_rx will unblock
|
||||||
streams.clear();
|
streams.clear();
|
||||||
|
|
||||||
// Wait for active stream handlers to finish so their frame_tx clones
|
handler_handles.shutdown().await;
|
||||||
// are dropped before the writer closes the sink.
|
|
||||||
drain_handlers(handler_handles).await;
|
|
||||||
|
|
||||||
match read_err {
|
match read_err {
|
||||||
Some(e) => Err(e.into()),
|
Some(e) => Err(e.into()),
|
||||||
@@ -478,30 +540,13 @@ where
|
|||||||
}
|
}
|
||||||
|
|
||||||
async fn dispatch_stream_frame(tx: &mpsc::Sender<Frame>, frame: Frame) -> StreamDispatchStatus {
|
async fn dispatch_stream_frame(tx: &mpsc::Sender<Frame>, frame: Frame) -> StreamDispatchStatus {
|
||||||
let stream_id = frame.stream_id;
|
let Some(frame) = attach_request_body_queue_budget(frame).await else {
|
||||||
let dispatched = tokio::time::timeout(stream_frame_dispatch_timeout(), async {
|
return StreamDispatchStatus::Congested;
|
||||||
let frame = attach_request_body_queue_budget(frame).await?;
|
};
|
||||||
tx.send(frame).await.ok()?;
|
match tx.try_send(frame) {
|
||||||
Some(())
|
Ok(()) => StreamDispatchStatus::Delivered,
|
||||||
})
|
Err(mpsc::error::TrySendError::Closed(_)) => StreamDispatchStatus::Closed,
|
||||||
.await;
|
Err(mpsc::error::TrySendError::Full(_)) => StreamDispatchStatus::Congested,
|
||||||
match dispatched {
|
|
||||||
Ok(Some(())) => StreamDispatchStatus::Delivered,
|
|
||||||
Ok(None) => {
|
|
||||||
warn!(
|
|
||||||
stream_id,
|
|
||||||
"stream handler channel or request body budget closed while dispatching tunnel frame"
|
|
||||||
);
|
|
||||||
StreamDispatchStatus::Closed
|
|
||||||
}
|
|
||||||
Err(_) => {
|
|
||||||
warn!(
|
|
||||||
stream_id,
|
|
||||||
timeout_ms = stream_frame_dispatch_timeout().as_millis(),
|
|
||||||
"stream handler channel blocked while dispatching tunnel frame"
|
|
||||||
);
|
|
||||||
StreamDispatchStatus::TimedOut
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -523,7 +568,7 @@ async fn attach_request_body_queue_budget_with(
|
|||||||
return Some(frame);
|
return Some(frame);
|
||||||
}
|
}
|
||||||
let permits = request_body_queue_permits(&frame, budget_bytes)?;
|
let permits = request_body_queue_permits(&frame, budget_bytes)?;
|
||||||
let permit = budget.acquire_many_owned(permits).await.ok()?;
|
let permit = budget.try_acquire_many_owned(permits).ok()?;
|
||||||
frame.payload = Bytes::from_owner(BudgetedFramePayload {
|
frame.payload = Bytes::from_owner(BudgetedFramePayload {
|
||||||
bytes: frame.payload,
|
bytes: frame.payload,
|
||||||
_permit: permit,
|
_permit: permit,
|
||||||
@@ -549,20 +594,6 @@ fn request_body_queue_permits(frame: &Frame, budget_bytes: usize) -> Option<u32>
|
|||||||
u32::try_from(retained_bytes).ok()
|
u32::try_from(retained_bytes).ok()
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Bound how long a single stream handler is allowed to block the shared
|
|
||||||
/// WebSocket read loop while receiving request-body frames.
|
|
||||||
fn stream_frame_dispatch_timeout() -> Duration {
|
|
||||||
#[cfg(test)]
|
|
||||||
{
|
|
||||||
Duration::from_millis(25)
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(not(test))]
|
|
||||||
{
|
|
||||||
Duration::from_millis(500)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn try_send_stream_error(frame_tx: &FrameSender, stream_id: u32, message: &'static str) {
|
fn try_send_stream_error(frame_tx: &FrameSender, stream_id: u32, message: &'static str) {
|
||||||
if frame_tx
|
if frame_tx
|
||||||
.try_send(Frame::new(
|
.try_send(Frame::new(
|
||||||
@@ -573,6 +604,7 @@ fn try_send_stream_error(frame_tx: &FrameSender, stream_id: u32, message: &'stat
|
|||||||
))
|
))
|
||||||
.is_err()
|
.is_err()
|
||||||
{
|
{
|
||||||
|
frame_tx.close();
|
||||||
warn!(
|
warn!(
|
||||||
stream_id,
|
stream_id,
|
||||||
"writer channel full, StreamError dropped while aborting stalled stream"
|
"writer channel full, StreamError dropped while aborting stalled stream"
|
||||||
@@ -587,21 +619,6 @@ fn prune_closed_stream_senders(streams: &mut HashMap<u32, StreamDispatchTarget>)
|
|||||||
before.saturating_sub(streams.len())
|
before.saturating_sub(streams.len())
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Wait for all active stream handlers to finish (with a timeout).
|
|
||||||
async fn drain_handlers(handles: Vec<JoinHandle<()>>) {
|
|
||||||
if handles.is_empty() {
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
let count = handles.len();
|
|
||||||
debug!(count, "waiting for active stream handlers to finish");
|
|
||||||
let _ = tokio::time::timeout(Duration::from_secs(30), async {
|
|
||||||
for h in handles {
|
|
||||||
let _ = h.await;
|
|
||||||
}
|
|
||||||
})
|
|
||||||
.await;
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
@@ -633,7 +650,7 @@ mod tests {
|
|||||||
|
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
stalled_send.await.expect("dispatch task should join"),
|
stalled_send.await.expect("dispatch task should join"),
|
||||||
StreamDispatchStatus::TimedOut
|
StreamDispatchStatus::Congested
|
||||||
);
|
);
|
||||||
|
|
||||||
let retained = rx
|
let retained = rx
|
||||||
@@ -737,6 +754,7 @@ mod tests {
|
|||||||
StreamDispatchTarget {
|
StreamDispatchTarget {
|
||||||
body_tx: closed_tx,
|
body_tx: closed_tx,
|
||||||
response_window: Arc::new(StreamSendWindow::new(1024)),
|
response_window: Arc::new(StreamSendWindow::new(1024)),
|
||||||
|
handler: None,
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
@@ -744,6 +762,7 @@ mod tests {
|
|||||||
StreamDispatchTarget {
|
StreamDispatchTarget {
|
||||||
body_tx: open_tx,
|
body_tx: open_tx,
|
||||||
response_window: Arc::new(StreamSendWindow::new(1024)),
|
response_window: Arc::new(StreamSendWindow::new(1024)),
|
||||||
|
handler: None,
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
]);
|
]);
|
||||||
@@ -763,6 +782,7 @@ mod tests {
|
|||||||
StreamDispatchTarget {
|
StreamDispatchTarget {
|
||||||
body_tx: tx,
|
body_tx: tx,
|
||||||
response_window: Arc::new(StreamSendWindow::new(1024)),
|
response_window: Arc::new(StreamSendWindow::new(1024)),
|
||||||
|
handler: None,
|
||||||
},
|
},
|
||||||
)]);
|
)]);
|
||||||
let mut active_handler_ids = HashSet::from([7]);
|
let mut active_handler_ids = HashSet::from([7]);
|
||||||
|
|||||||
@@ -31,14 +31,22 @@ enum AckDecision {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Handle for the dispatcher to forward HeartbeatAck frames.
|
/// Handle for the dispatcher to forward HeartbeatAck frames.
|
||||||
#[derive(Clone)]
|
|
||||||
pub struct HeartbeatHandle {
|
pub struct HeartbeatHandle {
|
||||||
ack_tx: tokio::sync::mpsc::Sender<Bytes>,
|
ack_tx: tokio::sync::mpsc::Sender<Bytes>,
|
||||||
|
task: Option<tokio::task::JoinHandle<()>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl HeartbeatHandle {
|
impl HeartbeatHandle {
|
||||||
pub async fn on_ack(&self, payload: Bytes) {
|
pub fn on_ack(&self, payload: Bytes) {
|
||||||
let _ = self.ack_tx.send(payload).await;
|
let _ = self.ack_tx.try_send(payload);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Drop for HeartbeatHandle {
|
||||||
|
fn drop(&mut self) {
|
||||||
|
if let Some(task) = self.task.take() {
|
||||||
|
task.abort();
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -48,7 +56,7 @@ impl HeartbeatHandle {
|
|||||||
pub fn spawn_noop() -> HeartbeatHandle {
|
pub fn spawn_noop() -> HeartbeatHandle {
|
||||||
let (ack_tx, _) = tokio::sync::mpsc::channel::<Bytes>(1);
|
let (ack_tx, _) = tokio::sync::mpsc::channel::<Bytes>(1);
|
||||||
// receiver is immediately dropped; on_ack() calls will silently fail
|
// receiver is immediately dropped; on_ack() calls will silently fail
|
||||||
HeartbeatHandle { ack_tx }
|
HeartbeatHandle { ack_tx, task: None }
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, Copy, Default)]
|
#[derive(Debug, Clone, Copy, Default)]
|
||||||
@@ -74,7 +82,7 @@ pub fn spawn(
|
|||||||
) -> HeartbeatHandle {
|
) -> HeartbeatHandle {
|
||||||
let (ack_tx, mut ack_rx) = tokio::sync::mpsc::channel::<Bytes>(4);
|
let (ack_tx, mut ack_rx) = tokio::sync::mpsc::channel::<Bytes>(4);
|
||||||
|
|
||||||
tokio::spawn(async move {
|
let task = tokio::spawn(async move {
|
||||||
// Read initial interval from dynamic config (may be updated by remote config).
|
// Read initial interval from dynamic config (may be updated by remote config).
|
||||||
let initial_interval = Duration::from_secs(server.dynamic.load().heartbeat_interval);
|
let initial_interval = Duration::from_secs(server.dynamic.load().heartbeat_interval);
|
||||||
let mut current_interval = initial_interval;
|
let mut current_interval = initial_interval;
|
||||||
@@ -151,7 +159,8 @@ pub fn spawn(
|
|||||||
current_interval = new_interval;
|
current_interval = new_interval;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Some(ack_payload) = ack_rx.recv() => {
|
ack_payload = ack_rx.recv() => {
|
||||||
|
let Some(ack_payload) = ack_payload else { break; };
|
||||||
match handle_ack(&server, &ack_payload) {
|
match handle_ack(&server, &ack_payload) {
|
||||||
AckDecision::Accept {
|
AckDecision::Accept {
|
||||||
heartbeat_id: ack_id,
|
heartbeat_id: ack_id,
|
||||||
@@ -179,7 +188,10 @@ pub fn spawn(
|
|||||||
}
|
}
|
||||||
});
|
});
|
||||||
|
|
||||||
HeartbeatHandle { ack_tx }
|
HeartbeatHandle {
|
||||||
|
ack_tx,
|
||||||
|
task: Some(task),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn build_heartbeat_payload(
|
async fn build_heartbeat_payload(
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ pub mod dispatcher;
|
|||||||
pub mod heartbeat;
|
pub mod heartbeat;
|
||||||
pub mod protocol;
|
pub mod protocol;
|
||||||
pub mod stream_handler;
|
pub mod stream_handler;
|
||||||
|
mod task;
|
||||||
pub mod writer;
|
pub mod writer;
|
||||||
|
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
@@ -332,9 +333,9 @@ mod tests {
|
|||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
assert_eq!(gateway_state.force_close_all_tunnel_proxies(), 1);
|
|
||||||
tokio::time::sleep(Duration::from_millis(200)).await;
|
|
||||||
gateway_handle.abort();
|
gateway_handle.abort();
|
||||||
|
let _ = (&mut gateway_handle).await;
|
||||||
|
assert_eq!(gateway_state.force_close_all_tunnel_proxies(), 1);
|
||||||
|
|
||||||
let (_restarted_gateway_state, restarted_gateway_handle) =
|
let (_restarted_gateway_state, restarted_gateway_handle) =
|
||||||
start_gateway_on_port_retry(gateway_port)
|
start_gateway_on_port_retry(gateway_port)
|
||||||
@@ -349,6 +350,7 @@ mod tests {
|
|||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
|
assert!(server.tunnel_metrics.snapshot().connect_successes >= 2);
|
||||||
let _ = shutdown_tx.send(true);
|
let _ = shutdown_tx.send(true);
|
||||||
tokio::time::timeout(Duration::from_secs(5), tunnel_task)
|
tokio::time::timeout(Duration::from_secs(5), tunnel_task)
|
||||||
.await
|
.await
|
||||||
@@ -380,7 +382,17 @@ mod tests {
|
|||||||
gateway_base_url: &str,
|
gateway_base_url: &str,
|
||||||
node_id: &str,
|
node_id: &str,
|
||||||
) -> Option<(StatusCode, String)> {
|
) -> Option<(StatusCode, String)> {
|
||||||
let payload = relay_probe_envelope();
|
let response = relay_response(gateway_base_url, node_id, relay_probe_envelope()).await?;
|
||||||
|
let status = response.status();
|
||||||
|
let body = response.text().await.unwrap_or_default();
|
||||||
|
Some((status, body))
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn relay_response(
|
||||||
|
gateway_base_url: &str,
|
||||||
|
node_id: &str,
|
||||||
|
payload: Vec<u8>,
|
||||||
|
) -> Option<reqwest::Response> {
|
||||||
let timestamp = SystemTime::now()
|
let timestamp = SystemTime::now()
|
||||||
.duration_since(UNIX_EPOCH)
|
.duration_since(UNIX_EPOCH)
|
||||||
.expect("test clock should be after epoch")
|
.expect("test clock should be after epoch")
|
||||||
@@ -398,7 +410,7 @@ mod tests {
|
|||||||
&nonce,
|
&nonce,
|
||||||
&digest,
|
&digest,
|
||||||
);
|
);
|
||||||
let response = reqwest::Client::new()
|
reqwest::Client::new()
|
||||||
.post(format!(
|
.post(format!(
|
||||||
"{gateway_base_url}/api/internal/tunnel/relay/{node_id}"
|
"{gateway_base_url}/api/internal/tunnel/relay/{node_id}"
|
||||||
))
|
))
|
||||||
@@ -421,10 +433,7 @@ mod tests {
|
|||||||
.body(payload)
|
.body(payload)
|
||||||
.send()
|
.send()
|
||||||
.await
|
.await
|
||||||
.ok()?;
|
.ok()
|
||||||
let status = response.status();
|
|
||||||
let body = response.text().await.unwrap_or_default();
|
|
||||||
Some((status, body))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn relay_probe_envelope() -> Vec<u8> {
|
fn relay_probe_envelope() -> Vec<u8> {
|
||||||
@@ -456,30 +465,150 @@ mod tests {
|
|||||||
) -> Result<(GatewayAppState, tokio::task::JoinHandle<()>), std::io::Error> {
|
) -> Result<(GatewayAppState, tokio::task::JoinHandle<()>), std::io::Error> {
|
||||||
// The embedded gateway now fails closed when relay authentication is
|
// The embedded gateway now fails closed when relay authentication is
|
||||||
// not configured. Keep this integration fixture explicitly authenticated.
|
// not configured. Keep this integration fixture explicitly authenticated.
|
||||||
let previous_secret = std::env::var_os("AETHER_TUNNEL_RELAY_AUTH_SECRET");
|
static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
|
||||||
let previous_instance = std::env::var_os("AETHER_GATEWAY_INSTANCE_ID");
|
let state = {
|
||||||
std::env::set_var(
|
let _guard = ENV_LOCK.lock().unwrap();
|
||||||
"AETHER_TUNNEL_RELAY_AUTH_SECRET",
|
let previous_secret = std::env::var_os("AETHER_TUNNEL_RELAY_AUTH_SECRET");
|
||||||
"tunnel-reconnect-test-secret-at-least-32-bytes",
|
let previous_instance = std::env::var_os("AETHER_GATEWAY_INSTANCE_ID");
|
||||||
);
|
std::env::set_var(
|
||||||
std::env::set_var(
|
"AETHER_TUNNEL_RELAY_AUTH_SECRET",
|
||||||
"AETHER_GATEWAY_INSTANCE_ID",
|
"tunnel-reconnect-test-secret-at-least-32-bytes",
|
||||||
"tunnel-reconnect-test-gateway",
|
);
|
||||||
);
|
std::env::set_var(
|
||||||
let mut state = GatewayAppState::new().expect("gateway test state should build");
|
"AETHER_GATEWAY_INSTANCE_ID",
|
||||||
aether_gateway::configure_test_tunnel_security(
|
"tunnel-reconnect-test-gateway",
|
||||||
&mut state,
|
);
|
||||||
"node-recovery",
|
let mut state = GatewayAppState::new().expect("gateway test state should build");
|
||||||
"test-generation-1",
|
aether_gateway::configure_test_tunnel_security(
|
||||||
"BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=",
|
&mut state,
|
||||||
);
|
"node-recovery",
|
||||||
restore_test_env("AETHER_TUNNEL_RELAY_AUTH_SECRET", previous_secret);
|
"test-generation-1",
|
||||||
restore_test_env("AETHER_GATEWAY_INSTANCE_ID", previous_instance);
|
"BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=",
|
||||||
|
);
|
||||||
|
restore_test_env("AETHER_TUNNEL_RELAY_AUTH_SECRET", previous_secret);
|
||||||
|
restore_test_env("AETHER_GATEWAY_INSTANCE_ID", previous_instance);
|
||||||
|
state
|
||||||
|
};
|
||||||
let router = build_router_with_state(state.clone());
|
let router = build_router_with_state(state.clone());
|
||||||
let handle = spawn_router_on_port(port, router).await?;
|
let handle = spawn_router_on_port(port, router).await?;
|
||||||
Ok((state, handle))
|
Ok((state, handle))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn negotiated_small_window_streams_large_responses_and_cancels_idle_upstream() {
|
||||||
|
use axum::body::{Body, Bytes};
|
||||||
|
use axum::routing::get;
|
||||||
|
use futures_util::StreamExt;
|
||||||
|
|
||||||
|
ensure_rustls_provider();
|
||||||
|
let upstream_port = reserve_local_port().unwrap();
|
||||||
|
let upstream = Router::new()
|
||||||
|
.route(
|
||||||
|
"/large",
|
||||||
|
get(|| async { Body::from(vec![b'x'; 2 * 1024 * 1024]) }),
|
||||||
|
)
|
||||||
|
.route(
|
||||||
|
"/idle",
|
||||||
|
get(|| async {
|
||||||
|
let first = futures_util::stream::once(async {
|
||||||
|
Ok::<_, std::io::Error>(Bytes::from_static(b"data: started\n\n"))
|
||||||
|
});
|
||||||
|
(
|
||||||
|
[("content-type", "text/event-stream")],
|
||||||
|
Body::from_stream(first.chain(futures_util::stream::pending())),
|
||||||
|
)
|
||||||
|
}),
|
||||||
|
);
|
||||||
|
let upstream_task = super::task::SessionTask::new(
|
||||||
|
spawn_router_on_port(upstream_port, upstream).await.unwrap(),
|
||||||
|
);
|
||||||
|
let gateway_port = reserve_local_port().unwrap();
|
||||||
|
let gateway_url = format!("http://127.0.0.1:{gateway_port}");
|
||||||
|
let (_, gateway_task) = start_gateway_on_port(gateway_port).await.unwrap();
|
||||||
|
let gateway_task = super::task::SessionTask::new(gateway_task);
|
||||||
|
let mut config = sample_config(&gateway_url);
|
||||||
|
config.tunnel_security = crate::config::TunnelSecurity::NonTlsRequired;
|
||||||
|
config.tunnel_encryption_key = Some("BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc=".into());
|
||||||
|
config.tunnel_stream_initial_window_bytes = 512 * 1024;
|
||||||
|
config.tunnel_drain_deadline_ms = 100;
|
||||||
|
config.allow_private_targets = true;
|
||||||
|
config.allowed_ports.push(upstream_port);
|
||||||
|
let state = sample_state(config);
|
||||||
|
let server = sample_server(&state, "node-recovery");
|
||||||
|
let (shutdown_tx, shutdown_rx) = watch::channel(false);
|
||||||
|
let (_drain_tx, drain_rx) = watch::channel(false);
|
||||||
|
let tunnel_task = super::task::SessionTask::new(tokio::spawn({
|
||||||
|
let state = Arc::clone(&state);
|
||||||
|
let server = Arc::clone(&server);
|
||||||
|
async move {
|
||||||
|
run(&state, &server, 0, shutdown_rx, drain_rx).await;
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
wait_until_relay_status(&gateway_url, "node-recovery", StatusCode::GATEWAY_TIMEOUT).await;
|
||||||
|
|
||||||
|
let envelope = |path: &str| {
|
||||||
|
let mut meta: protocol::RequestMeta =
|
||||||
|
serde_json::from_slice(&relay_probe_envelope()[4..]).unwrap();
|
||||||
|
meta.url = format!("http://127.0.0.1:{upstream_port}/{path}");
|
||||||
|
meta.stream = true;
|
||||||
|
meta.timeout = 10;
|
||||||
|
meta.stream_first_byte_timeout_ms = Some(10_000);
|
||||||
|
let encoded = serde_json::to_vec(&meta).unwrap();
|
||||||
|
let mut result = (encoded.len() as u32).to_be_bytes().to_vec();
|
||||||
|
result.extend(encoded);
|
||||||
|
result
|
||||||
|
};
|
||||||
|
let response = relay_response(&gateway_url, "node-recovery", envelope("large"))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(response.status(), StatusCode::OK);
|
||||||
|
let body = tokio::time::timeout(Duration::from_secs(10), response.bytes())
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(body.len(), 2 * 1024 * 1024);
|
||||||
|
assert!(body.iter().all(|byte| *byte == b'x'));
|
||||||
|
|
||||||
|
let mut response = relay_response(&gateway_url, "node-recovery", envelope("idle"))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
response.chunk().await.unwrap().unwrap(),
|
||||||
|
"data: started\n\n"
|
||||||
|
);
|
||||||
|
drop(response);
|
||||||
|
tokio::time::timeout(Duration::from_secs(3), async {
|
||||||
|
while server
|
||||||
|
.active_connections
|
||||||
|
.load(std::sync::atomic::Ordering::Acquire)
|
||||||
|
!= 0
|
||||||
|
{
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.expect("cancelled SSE must release the upstream handler");
|
||||||
|
|
||||||
|
let mut response = relay_response(&gateway_url, "node-recovery", envelope("idle"))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert!(response.chunk().await.unwrap().is_some());
|
||||||
|
shutdown_tx.send(true).unwrap();
|
||||||
|
tokio::time::timeout(Duration::from_secs(3), tunnel_task)
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(
|
||||||
|
server
|
||||||
|
.active_connections
|
||||||
|
.load(std::sync::atomic::Ordering::Acquire),
|
||||||
|
0
|
||||||
|
);
|
||||||
|
drop(response);
|
||||||
|
drop(gateway_task);
|
||||||
|
drop(upstream_task);
|
||||||
|
}
|
||||||
|
|
||||||
fn restore_test_env(key: &str, value: Option<std::ffi::OsString>) {
|
fn restore_test_env(key: &str, value: Option<std::ffi::OsString>) {
|
||||||
if let Some(value) = value {
|
if let Some(value) = value {
|
||||||
std::env::set_var(key, value);
|
std::env::set_var(key, value);
|
||||||
|
|||||||
@@ -52,6 +52,7 @@ static REDIRECT_REPLAY_BUFFERED_BYTES: AtomicUsize = AtomicUsize::new(0);
|
|||||||
|
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
pub(crate) struct StreamSendWindow {
|
pub(crate) struct StreamSendWindow {
|
||||||
|
initial_window_bytes: u32,
|
||||||
available: Mutex<u64>,
|
available: Mutex<u64>,
|
||||||
notify: Notify,
|
notify: Notify,
|
||||||
}
|
}
|
||||||
@@ -59,6 +60,7 @@ pub(crate) struct StreamSendWindow {
|
|||||||
impl StreamSendWindow {
|
impl StreamSendWindow {
|
||||||
pub(crate) fn new(initial_window_bytes: u32) -> Self {
|
pub(crate) fn new(initial_window_bytes: u32) -> Self {
|
||||||
Self {
|
Self {
|
||||||
|
initial_window_bytes: initial_window_bytes.max(1),
|
||||||
available: Mutex::new(u64::from(initial_window_bytes.max(1))),
|
available: Mutex::new(u64::from(initial_window_bytes.max(1))),
|
||||||
notify: Notify::new(),
|
notify: Notify::new(),
|
||||||
}
|
}
|
||||||
@@ -82,6 +84,9 @@ impl StreamSendWindow {
|
|||||||
let requested = bytes as u64;
|
let requested = bytes as u64;
|
||||||
let started_at = Instant::now();
|
let started_at = Instant::now();
|
||||||
loop {
|
loop {
|
||||||
|
let notified = self.notify.notified();
|
||||||
|
tokio::pin!(notified);
|
||||||
|
notified.as_mut().enable();
|
||||||
{
|
{
|
||||||
let mut available = self.available.lock().expect("stream window lock poisoned");
|
let mut available = self.available.lock().expect("stream window lock poisoned");
|
||||||
if *available >= requested {
|
if *available >= requested {
|
||||||
@@ -93,10 +98,7 @@ impl StreamSendWindow {
|
|||||||
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
|
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
|
||||||
return Err(());
|
return Err(());
|
||||||
};
|
};
|
||||||
if tokio::time::timeout(remaining, self.notify.notified())
|
if tokio::time::timeout(remaining, notified).await.is_err() {
|
||||||
.await
|
|
||||||
.is_err()
|
|
||||||
{
|
|
||||||
return Err(());
|
return Err(());
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -173,31 +175,33 @@ fn safe_stream_error_message(message: &str) -> &'static str {
|
|||||||
"upstream request failed"
|
"upstream request failed"
|
||||||
}
|
}
|
||||||
|
|
||||||
fn try_send_window_update(frame_tx: &FrameSender, stream_id: u32, bytes: usize) {
|
async fn send_window_update(frame_tx: &FrameSender, stream_id: u32, bytes: usize) -> bool {
|
||||||
if bytes == 0 {
|
if bytes == 0 {
|
||||||
return;
|
return true;
|
||||||
}
|
}
|
||||||
let delta = bytes.min(u32::MAX as usize) as u32;
|
let delta = bytes.min(u32::MAX as usize) as u32;
|
||||||
if frame_tx
|
if matches!(
|
||||||
.try_send(TunnelFrame::new(
|
tokio::time::timeout(
|
||||||
stream_id,
|
FLOW_CONTROL_WAIT_TIMEOUT,
|
||||||
MsgType::WindowUpdate,
|
frame_tx.send(TunnelFrame::new(
|
||||||
0,
|
stream_id,
|
||||||
Bytes::from(
|
MsgType::WindowUpdate,
|
||||||
serde_json::to_vec(&aether_contracts::tunnel::WindowUpdatePayload {
|
0,
|
||||||
delta_bytes: delta,
|
Bytes::from(
|
||||||
})
|
serde_json::to_vec(&aether_contracts::tunnel::WindowUpdatePayload {
|
||||||
.expect("window update payload should serialize"),
|
delta_bytes: delta,
|
||||||
),
|
})
|
||||||
))
|
.expect("window update payload should serialize"),
|
||||||
.is_err()
|
),
|
||||||
{
|
))
|
||||||
warn!(
|
)
|
||||||
stream_id,
|
.await,
|
||||||
delta_bytes = delta,
|
Ok(Ok(()))
|
||||||
"writer channel full, WINDOW_UPDATE dropped"
|
) {
|
||||||
);
|
return true;
|
||||||
}
|
}
|
||||||
|
frame_tx.close();
|
||||||
|
false
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Match reqwest's default redirect budget so direct execution and tunnel relay
|
/// Match reqwest's default redirect budget so direct execution and tunnel relay
|
||||||
@@ -242,6 +246,23 @@ enum ReplayableRequestBody {
|
|||||||
struct PreparedRequestBody {
|
struct PreparedRequestBody {
|
||||||
first_request_body: Option<upstream_client::UpstreamRequestBody>,
|
first_request_body: Option<upstream_client::UpstreamRequestBody>,
|
||||||
replay_body: ReplayableRequestBody,
|
replay_body: ReplayableRequestBody,
|
||||||
|
spool_task: Option<tokio::task::JoinHandle<()>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Drop for PreparedRequestBody {
|
||||||
|
fn drop(&mut self) {
|
||||||
|
if let Some(task) = self.spool_task.take() {
|
||||||
|
task.abort();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
struct ActiveStreamGuard(Arc<ServerContext>);
|
||||||
|
|
||||||
|
impl Drop for ActiveStreamGuard {
|
||||||
|
fn drop(&mut self) {
|
||||||
|
self.0.active_connections.fetch_sub(1, Ordering::Release);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, Copy)]
|
#[derive(Debug, Clone, Copy)]
|
||||||
@@ -331,7 +352,10 @@ impl hyper::body::Body for ReplayRequestBody {
|
|||||||
|
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
enum SpoolBodyEvent {
|
enum SpoolBodyEvent {
|
||||||
Data(Bytes),
|
Data {
|
||||||
|
payload: Bytes,
|
||||||
|
credit_returned: bool,
|
||||||
|
},
|
||||||
Error(String),
|
Error(String),
|
||||||
End,
|
End,
|
||||||
}
|
}
|
||||||
@@ -563,8 +587,9 @@ impl RequestBodyReplayState {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn push_chunk(&self, payload: Bytes) {
|
fn push_chunk(&self, payload: Bytes) -> bool {
|
||||||
let mut disable_replay = false;
|
let mut disable_replay = false;
|
||||||
|
let mut retained = false;
|
||||||
let mut state = self.state.lock().expect("request body replay state lock");
|
let mut state = self.state.lock().expect("request body replay state lock");
|
||||||
if let RequestBodyReplayStatus::Collecting {
|
if let RequestBodyReplayStatus::Collecting {
|
||||||
chunks,
|
chunks,
|
||||||
@@ -577,7 +602,7 @@ impl RequestBodyReplayState {
|
|||||||
drop(state);
|
drop(state);
|
||||||
self.release_reserved_bytes();
|
self.release_reserved_bytes();
|
||||||
self.ready.notify_waiters();
|
self.ready.notify_waiters();
|
||||||
return;
|
return false;
|
||||||
};
|
};
|
||||||
let accounted_bytes = payload.len().checked_add(std::mem::size_of::<Bytes>());
|
let accounted_bytes = payload.len().checked_add(std::mem::size_of::<Bytes>());
|
||||||
if next_len > self.budget_bytes
|
if next_len > self.budget_bytes
|
||||||
@@ -590,6 +615,7 @@ impl RequestBodyReplayState {
|
|||||||
} else {
|
} else {
|
||||||
*buffered_len = next_len;
|
*buffered_len = next_len;
|
||||||
chunks.push(payload);
|
chunks.push(payload);
|
||||||
|
retained = true;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
drop(state);
|
drop(state);
|
||||||
@@ -597,6 +623,7 @@ impl RequestBodyReplayState {
|
|||||||
self.release_reserved_bytes();
|
self.release_reserved_bytes();
|
||||||
self.ready.notify_waiters();
|
self.ready.notify_waiters();
|
||||||
}
|
}
|
||||||
|
retained
|
||||||
}
|
}
|
||||||
|
|
||||||
fn try_reserve_bytes(&self, bytes: usize) -> bool {
|
fn try_reserve_bytes(&self, bytes: usize) -> bool {
|
||||||
@@ -951,9 +978,6 @@ pub(super) fn decode_request_body_frame(frame: TunnelFrame) -> Result<Bytes, std
|
|||||||
Ok(frame.payload)
|
Ok(frame.payload)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Drain tunnel body frames on a detached task so the shared dispatcher is no
|
|
||||||
// longer coupled to upstream body polling. Redirect replay retains a bounded
|
|
||||||
// copy; crossing either replay budget only disables replay for this request.
|
|
||||||
fn prepare_request_body(
|
fn prepare_request_body(
|
||||||
stream_id: u32,
|
stream_id: u32,
|
||||||
body_rx: mpsc::Receiver<TunnelFrame>,
|
body_rx: mpsc::Receiver<TunnelFrame>,
|
||||||
@@ -973,19 +997,20 @@ fn prepare_request_body(
|
|||||||
None => ReplayableRequestBody::NonReplayable,
|
None => ReplayableRequestBody::NonReplayable,
|
||||||
};
|
};
|
||||||
|
|
||||||
tokio::spawn(spool_request_body(
|
let spool_task = tokio::spawn(spool_request_body(
|
||||||
stream_id,
|
stream_id,
|
||||||
body_rx,
|
body_rx,
|
||||||
spool_tx,
|
spool_tx,
|
||||||
replay_state,
|
replay_state,
|
||||||
body_size,
|
body_size,
|
||||||
deadline,
|
deadline,
|
||||||
frame_tx,
|
frame_tx.clone(),
|
||||||
));
|
));
|
||||||
|
|
||||||
PreparedRequestBody {
|
PreparedRequestBody {
|
||||||
first_request_body: Some(build_spooled_request_body(spool_rx)),
|
first_request_body: Some(build_spooled_request_body(spool_rx, stream_id, frame_tx)),
|
||||||
replay_body,
|
replay_body,
|
||||||
|
spool_task: Some(spool_task),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1001,6 +1026,7 @@ fn prepare_bodyless_request_body(
|
|||||||
} else {
|
} else {
|
||||||
ReplayableRequestBody::NonReplayable
|
ReplayableRequestBody::NonReplayable
|
||||||
},
|
},
|
||||||
|
spool_task: None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1056,11 +1082,16 @@ async fn spool_request_body(
|
|||||||
};
|
};
|
||||||
|
|
||||||
let Some(frame) = frame else {
|
let Some(frame) = frame else {
|
||||||
|
let message = "tunnel request body closed before stream end".to_string();
|
||||||
if let Some(state) = &replay_state {
|
if let Some(state) = &replay_state {
|
||||||
state.finish();
|
state.fail(message.clone());
|
||||||
}
|
}
|
||||||
let _ =
|
let _ = send_spool_event(
|
||||||
send_spool_event(&mut spool_tx, SpoolBodyEvent::End, replay_state.as_ref()).await;
|
&mut spool_tx,
|
||||||
|
SpoolBodyEvent::Error(message),
|
||||||
|
replay_state.as_ref(),
|
||||||
|
)
|
||||||
|
.await;
|
||||||
return;
|
return;
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -1086,13 +1117,23 @@ async fn spool_request_body(
|
|||||||
|
|
||||||
if !payload.is_empty() {
|
if !payload.is_empty() {
|
||||||
body_size.fetch_add(payload.len(), Ordering::Relaxed);
|
body_size.fetch_add(payload.len(), Ordering::Relaxed);
|
||||||
try_send_window_update(&frame_tx, stream_id, payload.len());
|
let credit_returned = replay_state
|
||||||
if let Some(state) = &replay_state {
|
.as_ref()
|
||||||
state.push_chunk(payload.clone());
|
.is_some_and(|state| state.push_chunk(payload.clone()));
|
||||||
|
if credit_returned
|
||||||
|
&& !send_window_update(&frame_tx, stream_id, payload.len()).await
|
||||||
|
{
|
||||||
|
if let Some(state) = &replay_state {
|
||||||
|
state.fail("tunnel flow-control update failed".to_string());
|
||||||
|
}
|
||||||
|
return;
|
||||||
}
|
}
|
||||||
if send_spool_event(
|
if send_spool_event(
|
||||||
&mut spool_tx,
|
&mut spool_tx,
|
||||||
SpoolBodyEvent::Data(payload),
|
SpoolBodyEvent::Data {
|
||||||
|
payload,
|
||||||
|
credit_returned,
|
||||||
|
},
|
||||||
replay_state.as_ref(),
|
replay_state.as_ref(),
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
@@ -1479,6 +1520,7 @@ where
|
|||||||
}
|
}
|
||||||
|
|
||||||
let mut stream = response.into_body().into_data_stream();
|
let mut stream = response.into_body().into_data_stream();
|
||||||
|
let chunk_size = MAX_CHUNK_SIZE.min(response_window.initial_window_bytes as usize);
|
||||||
loop {
|
loop {
|
||||||
let chunk_result = if let Some(deadline) = response_body_deadline {
|
let chunk_result = if let Some(deadline) = response_body_deadline {
|
||||||
let Some(remaining) = remaining_timeout(deadline) else {
|
let Some(remaining) = remaining_timeout(deadline) else {
|
||||||
@@ -1531,7 +1573,7 @@ where
|
|||||||
|
|
||||||
match chunk_result {
|
match chunk_result {
|
||||||
Ok(chunk) => {
|
Ok(chunk) => {
|
||||||
if chunk.len() <= MAX_CHUNK_SIZE {
|
if chunk.len() <= chunk_size {
|
||||||
let (payload, extra_flags) = raw_payload(chunk);
|
let (payload, extra_flags) = raw_payload(chunk);
|
||||||
if !acquire_response_credit(response_window, frame_tx, stream_id, payload.len())
|
if !acquire_response_credit(response_window, frame_tx, stream_id, payload.len())
|
||||||
.await
|
.await
|
||||||
@@ -1561,7 +1603,7 @@ where
|
|||||||
} else {
|
} else {
|
||||||
let mut offset = 0;
|
let mut offset = 0;
|
||||||
while offset < chunk.len() {
|
while offset < chunk.len() {
|
||||||
let end = (offset + MAX_CHUNK_SIZE).min(chunk.len());
|
let end = (offset + chunk_size).min(chunk.len());
|
||||||
let slice = chunk.slice(offset..end);
|
let slice = chunk.slice(offset..end);
|
||||||
let (payload, extra_flags) = raw_payload(slice);
|
let (payload, extra_flags) = raw_payload(slice);
|
||||||
if !acquire_response_credit(
|
if !acquire_response_credit(
|
||||||
@@ -1735,6 +1777,7 @@ pub async fn handle_stream(
|
|||||||
};
|
};
|
||||||
|
|
||||||
server.active_connections.fetch_add(1, Ordering::Release);
|
server.active_connections.fetch_add(1, Ordering::Release);
|
||||||
|
let _active_stream = ActiveStreamGuard(Arc::clone(&server));
|
||||||
|
|
||||||
let stream_io = StreamIo {
|
let stream_io = StreamIo {
|
||||||
body_rx,
|
body_rx,
|
||||||
@@ -1745,7 +1788,6 @@ pub async fn handle_stream(
|
|||||||
|
|
||||||
let connect_elapsed = handle_stream_inner(&state, &server, stream_id, meta, stream_io).await;
|
let connect_elapsed = handle_stream_inner(&state, &server, stream_id, meta, stream_io).await;
|
||||||
|
|
||||||
server.active_connections.fetch_sub(1, Ordering::Release);
|
|
||||||
if let Some(d) = connect_elapsed {
|
if let Some(d) = connect_elapsed {
|
||||||
server.metrics.record_request(d);
|
server.metrics.record_request(d);
|
||||||
}
|
}
|
||||||
@@ -1772,6 +1814,18 @@ async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool {
|
|||||||
timeout_ms = FLOW_CONTROL_WAIT_TIMEOUT.as_millis() as u64,
|
timeout_ms = FLOW_CONTROL_WAIT_TIMEOUT.as_millis() as u64,
|
||||||
"writer channel stalled for body frame, abandoning stream"
|
"writer channel stalled for body frame, abandoning stream"
|
||||||
);
|
);
|
||||||
|
let reset = TunnelFrame::new(
|
||||||
|
stream_id,
|
||||||
|
MsgType::ResetStream,
|
||||||
|
0,
|
||||||
|
Bytes::from_static(b"{\"reason\":\"tunnel writer stalled\"}"),
|
||||||
|
);
|
||||||
|
if !matches!(
|
||||||
|
tokio::time::timeout(CONTROL_FRAME_SEND_TIMEOUT, tx.send(reset)).await,
|
||||||
|
Ok(Ok(()))
|
||||||
|
) {
|
||||||
|
tx.close();
|
||||||
|
}
|
||||||
false
|
false
|
||||||
}
|
}
|
||||||
Ok(Err(QueueSendError::Full(_))) => {
|
Ok(Err(QueueSendError::Full(_))) => {
|
||||||
@@ -1781,7 +1835,10 @@ async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool {
|
|||||||
} else {
|
} else {
|
||||||
match tokio::time::timeout(CONTROL_FRAME_SEND_TIMEOUT, tx.send(frame)).await {
|
match tokio::time::timeout(CONTROL_FRAME_SEND_TIMEOUT, tx.send(frame)).await {
|
||||||
Ok(Ok(())) => true,
|
Ok(Ok(())) => true,
|
||||||
Ok(Err(_)) => false,
|
Ok(Err(_)) => {
|
||||||
|
tx.close();
|
||||||
|
false
|
||||||
|
}
|
||||||
Err(_) => {
|
Err(_) => {
|
||||||
warn!(
|
warn!(
|
||||||
stream_id,
|
stream_id,
|
||||||
@@ -1789,6 +1846,7 @@ async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool {
|
|||||||
flags = flags,
|
flags = flags,
|
||||||
"control frame send timeout (writer congested), abandoning stream"
|
"control frame send timeout (writer congested), abandoning stream"
|
||||||
);
|
);
|
||||||
|
tx.close();
|
||||||
false
|
false
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -2126,7 +2184,6 @@ async fn handle_stream_inner(
|
|||||||
}
|
}
|
||||||
|
|
||||||
async fn send_error(tx: &FrameSender, stream_id: u32, msg: &str) {
|
async fn send_error(tx: &FrameSender, stream_id: u32, msg: &str) {
|
||||||
// Error frames use best-effort delivery — don't block if writer is congested
|
|
||||||
let safe_message = safe_stream_error_message(msg);
|
let safe_message = safe_stream_error_message(msg);
|
||||||
let _ = send_frame(
|
let _ = send_frame(
|
||||||
tx,
|
tx,
|
||||||
@@ -2163,22 +2220,42 @@ fn build_streaming_request_body(
|
|||||||
|
|
||||||
fn build_spooled_request_body(
|
fn build_spooled_request_body(
|
||||||
spool_rx: mpsc::Receiver<SpoolBodyEvent>,
|
spool_rx: mpsc::Receiver<SpoolBodyEvent>,
|
||||||
|
stream_id: u32,
|
||||||
|
frame_tx: FrameSender,
|
||||||
) -> upstream_client::UpstreamRequestBody {
|
) -> upstream_client::UpstreamRequestBody {
|
||||||
let body_stream = stream::unfold((spool_rx, false), |(mut spool_rx, finished)| async move {
|
let body_stream = stream::unfold(
|
||||||
if finished {
|
(spool_rx, frame_tx, false),
|
||||||
return None;
|
move |(mut spool_rx, frame_tx, finished)| async move {
|
||||||
}
|
if finished {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
|
||||||
match spool_rx.recv().await {
|
match spool_rx.recv().await {
|
||||||
Some(SpoolBodyEvent::Data(payload)) => {
|
Some(SpoolBodyEvent::Data {
|
||||||
Some((Ok(BodyFrame::data(payload)), (spool_rx, false)))
|
payload,
|
||||||
|
credit_returned,
|
||||||
|
}) => {
|
||||||
|
if !credit_returned
|
||||||
|
&& !send_window_update(&frame_tx, stream_id, payload.len()).await
|
||||||
|
{
|
||||||
|
return Some((
|
||||||
|
Err(io::Error::other("tunnel flow-control update failed")),
|
||||||
|
(spool_rx, frame_tx, true),
|
||||||
|
));
|
||||||
|
}
|
||||||
|
Some((Ok(BodyFrame::data(payload)), (spool_rx, frame_tx, false)))
|
||||||
|
}
|
||||||
|
Some(SpoolBodyEvent::Error(message)) => {
|
||||||
|
Some((Err(io::Error::other(message)), (spool_rx, frame_tx, true)))
|
||||||
|
}
|
||||||
|
Some(SpoolBodyEvent::End) => None,
|
||||||
|
None => Some((
|
||||||
|
Err(io::Error::other("tunnel request body ended unexpectedly")),
|
||||||
|
(spool_rx, frame_tx, true),
|
||||||
|
)),
|
||||||
}
|
}
|
||||||
Some(SpoolBodyEvent::Error(message)) => {
|
},
|
||||||
Some((Err(io::Error::other(message)), (spool_rx, true)))
|
);
|
||||||
}
|
|
||||||
Some(SpoolBodyEvent::End) | None => None,
|
|
||||||
}
|
|
||||||
});
|
|
||||||
|
|
||||||
upstream_client::stream_request_body(body_stream)
|
upstream_client::stream_request_body(body_stream)
|
||||||
}
|
}
|
||||||
@@ -2249,6 +2326,105 @@ fn build_prefixed_request_body(
|
|||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
|
#[tokio::test(start_paused = true)]
|
||||||
|
async fn window_updates_wait_for_capacity_instead_of_disappearing() {
|
||||||
|
let (high_tx, mut high_rx) = aether_runtime::bounded_queue(1);
|
||||||
|
let (normal_tx, _normal_rx) = aether_runtime::bounded_queue(1);
|
||||||
|
let sender = FrameSender::from_test_queues(high_tx, normal_tx);
|
||||||
|
sender
|
||||||
|
.try_send(TunnelFrame::control(MsgType::Ping, Bytes::new()))
|
||||||
|
.unwrap();
|
||||||
|
let task = tokio::spawn(async move { send_window_update(&sender, 7, 1024).await });
|
||||||
|
tokio::time::sleep(Duration::from_secs(1)).await;
|
||||||
|
assert!(!task.is_finished());
|
||||||
|
high_rx.recv().await.unwrap();
|
||||||
|
assert!(task.await.unwrap());
|
||||||
|
let update = high_rx.recv().await.unwrap();
|
||||||
|
assert_eq!(update.msg_type, MsgType::WindowUpdate);
|
||||||
|
let payload: aether_contracts::tunnel::WindowUpdatePayload =
|
||||||
|
serde_json::from_slice(&update.payload).unwrap();
|
||||||
|
assert_eq!(payload.delta_bytes, 1024);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test(start_paused = true)]
|
||||||
|
async fn stalled_body_delivery_emits_a_reset() {
|
||||||
|
let (high_tx, mut high_rx) = aether_runtime::bounded_queue(4);
|
||||||
|
let (normal_tx, _normal_rx) = aether_runtime::bounded_queue(1);
|
||||||
|
let sender = FrameSender::from_test_queues(high_tx, normal_tx);
|
||||||
|
sender
|
||||||
|
.try_send(TunnelFrame::new(
|
||||||
|
7,
|
||||||
|
MsgType::ResponseBody,
|
||||||
|
0,
|
||||||
|
Bytes::from_static(b"first"),
|
||||||
|
))
|
||||||
|
.unwrap();
|
||||||
|
assert!(
|
||||||
|
!send_frame(
|
||||||
|
&sender,
|
||||||
|
TunnelFrame::new(7, MsgType::ResponseBody, 0, Bytes::from_static(b"second"))
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
);
|
||||||
|
let reset = high_rx.recv().await.unwrap();
|
||||||
|
assert_eq!(reset.msg_type, MsgType::ResetStream);
|
||||||
|
assert_eq!(reset.stream_id, 7);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn request_credit_follows_consumption_without_redirect_replay() {
|
||||||
|
let (body_tx, body_rx) = mpsc::channel(4);
|
||||||
|
let (high_tx, mut high_rx) = aether_runtime::bounded_queue(4);
|
||||||
|
let (normal_tx, _normal_rx) = aether_runtime::bounded_queue(4);
|
||||||
|
let sender = FrameSender::from_test_queues(high_tx, normal_tx);
|
||||||
|
let mut prepared = prepare_request_body(
|
||||||
|
7,
|
||||||
|
body_rx,
|
||||||
|
Arc::new(AtomicUsize::new(0)),
|
||||||
|
Instant::now() + Duration::from_secs(10),
|
||||||
|
false,
|
||||||
|
sender,
|
||||||
|
);
|
||||||
|
body_tx
|
||||||
|
.send(TunnelFrame::new(
|
||||||
|
7,
|
||||||
|
MsgType::RequestBody,
|
||||||
|
flags::END_STREAM,
|
||||||
|
Bytes::from_static(b"body"),
|
||||||
|
))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
assert!(high_rx.try_recv().is_err());
|
||||||
|
let mut body = prepared.take_first_request_body();
|
||||||
|
assert!(body.frame().await.unwrap().is_ok());
|
||||||
|
assert_eq!(
|
||||||
|
high_rx.recv().await.unwrap().msg_type,
|
||||||
|
MsgType::WindowUpdate
|
||||||
|
);
|
||||||
|
assert!(body.frame().await.is_none());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn dropping_prepared_body_cancels_its_spooler() {
|
||||||
|
let (body_tx, body_rx) = mpsc::channel(4);
|
||||||
|
let (high_tx, _high_rx) = aether_runtime::bounded_queue(4);
|
||||||
|
let (normal_tx, _normal_rx) = aether_runtime::bounded_queue(4);
|
||||||
|
let sender = FrameSender::from_test_queues(high_tx, normal_tx);
|
||||||
|
let prepared = prepare_request_body(
|
||||||
|
7,
|
||||||
|
body_rx,
|
||||||
|
Arc::new(AtomicUsize::new(0)),
|
||||||
|
Instant::now() + Duration::from_secs(3600),
|
||||||
|
false,
|
||||||
|
sender,
|
||||||
|
);
|
||||||
|
drop(prepared);
|
||||||
|
tokio::time::timeout(Duration::from_secs(1), body_tx.closed())
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
use std::collections::HashMap;
|
use std::collections::HashMap;
|
||||||
use std::net::SocketAddr;
|
use std::net::SocketAddr;
|
||||||
use std::pin::Pin;
|
use std::pin::Pin;
|
||||||
@@ -2379,7 +2555,7 @@ mod tests {
|
|||||||
let (tx, rx) = mpsc::channel(4);
|
let (tx, rx) = mpsc::channel(4);
|
||||||
let (frame_tx, sent, writer_handle) = spawn_test_writer();
|
let (frame_tx, sent, writer_handle) = spawn_test_writer();
|
||||||
let body_size = Arc::new(AtomicUsize::new(0));
|
let body_size = Arc::new(AtomicUsize::new(0));
|
||||||
let prepared = prepare_request_body(
|
let mut prepared = prepare_request_body(
|
||||||
1,
|
1,
|
||||||
rx,
|
rx,
|
||||||
Arc::clone(&body_size),
|
Arc::clone(&body_size),
|
||||||
@@ -2389,6 +2565,7 @@ mod tests {
|
|||||||
);
|
);
|
||||||
let mut body = prepared
|
let mut body = prepared
|
||||||
.first_request_body
|
.first_request_body
|
||||||
|
.take()
|
||||||
.expect("first request body should be present");
|
.expect("first request body should be present");
|
||||||
|
|
||||||
tx.send(TunnelFrame::new(
|
tx.send(TunnelFrame::new(
|
||||||
|
|||||||
@@ -0,0 +1,52 @@
|
|||||||
|
use std::future::Future;
|
||||||
|
use std::pin::Pin;
|
||||||
|
use std::task::{Context, Poll};
|
||||||
|
|
||||||
|
use tokio::task::{JoinError, JoinHandle};
|
||||||
|
|
||||||
|
pub(super) struct SessionTask<T>(JoinHandle<T>);
|
||||||
|
|
||||||
|
impl<T> SessionTask<T> {
|
||||||
|
pub(super) fn new(handle: JoinHandle<T>) -> Self {
|
||||||
|
Self(handle)
|
||||||
|
}
|
||||||
|
pub(super) fn abort(&self) {
|
||||||
|
self.0.abort();
|
||||||
|
}
|
||||||
|
pub(super) fn is_finished(&self) -> bool {
|
||||||
|
self.0.is_finished()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<T> Future for SessionTask<T> {
|
||||||
|
type Output = Result<T, JoinError>;
|
||||||
|
|
||||||
|
fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
|
||||||
|
Pin::new(&mut self.0).poll(context)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<T> Drop for SessionTask<T> {
|
||||||
|
fn drop(&mut self) {
|
||||||
|
self.0.abort();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn dropping_a_session_task_aborts_its_child() {
|
||||||
|
let child = tokio::spawn(std::future::pending::<()>());
|
||||||
|
let abort = child.abort_handle();
|
||||||
|
drop(SessionTask::new(child));
|
||||||
|
tokio::time::timeout(std::time::Duration::from_secs(1), async {
|
||||||
|
while !abort.is_finished() {
|
||||||
|
tokio::task::yield_now().await;
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -13,6 +13,7 @@ use aether_contracts::tunnel::{MsgType, HEADER_SIZE};
|
|||||||
use aether_runtime::QueueSnapshot;
|
use aether_runtime::QueueSnapshot;
|
||||||
use aether_runtime::{bounded_queue, BoundedQueueSender, QueueSendError};
|
use aether_runtime::{bounded_queue, BoundedQueueSender, QueueSendError};
|
||||||
use futures_util::SinkExt;
|
use futures_util::SinkExt;
|
||||||
|
use tokio::sync::watch;
|
||||||
use tokio::task::JoinHandle;
|
use tokio::task::JoinHandle;
|
||||||
use tokio_tungstenite::tungstenite::Message;
|
use tokio_tungstenite::tungstenite::Message;
|
||||||
use tracing::{debug, error, trace};
|
use tracing::{debug, error, trace};
|
||||||
@@ -24,6 +25,8 @@ use aether_contracts::tunnel_security::SecureFrameCodec;
|
|||||||
|
|
||||||
const HIGH_PRIORITY_QUEUE_CAPACITY: usize = 64;
|
const HIGH_PRIORITY_QUEUE_CAPACITY: usize = 64;
|
||||||
const NORMAL_PRIORITY_QUEUE_CAPACITY: usize = 256;
|
const NORMAL_PRIORITY_QUEUE_CAPACITY: usize = 256;
|
||||||
|
const WRITE_TIMEOUT: Duration = Duration::from_secs(15);
|
||||||
|
const CLOSE_TIMEOUT: Duration = Duration::from_secs(1);
|
||||||
|
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||||
enum FramePriority {
|
enum FramePriority {
|
||||||
@@ -43,9 +46,18 @@ pub struct FrameQueueSnapshots {
|
|||||||
pub struct FrameSender {
|
pub struct FrameSender {
|
||||||
high_tx: BoundedQueueSender<Frame>,
|
high_tx: BoundedQueueSender<Frame>,
|
||||||
normal_tx: BoundedQueueSender<Frame>,
|
normal_tx: BoundedQueueSender<Frame>,
|
||||||
|
close_tx: watch::Sender<bool>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl FrameSender {
|
impl FrameSender {
|
||||||
|
pub fn close(&self) {
|
||||||
|
let _ = self.close_tx.send(true);
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) fn subscribe_close(&self) -> watch::Receiver<bool> {
|
||||||
|
self.close_tx.subscribe()
|
||||||
|
}
|
||||||
|
|
||||||
pub async fn send(&self, frame: Frame) -> Result<(), QueueSendError<Frame>> {
|
pub async fn send(&self, frame: Frame) -> Result<(), QueueSendError<Frame>> {
|
||||||
match classify_frame_priority(&frame) {
|
match classify_frame_priority(&frame) {
|
||||||
FramePriority::High => self.high_tx.send(frame).await,
|
FramePriority::High => self.high_tx.send(frame).await,
|
||||||
@@ -73,7 +85,12 @@ impl FrameSender {
|
|||||||
high_tx: BoundedQueueSender<Frame>,
|
high_tx: BoundedQueueSender<Frame>,
|
||||||
normal_tx: BoundedQueueSender<Frame>,
|
normal_tx: BoundedQueueSender<Frame>,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
Self { high_tx, normal_tx }
|
let (close_tx, _) = watch::channel(false);
|
||||||
|
Self {
|
||||||
|
high_tx,
|
||||||
|
normal_tx,
|
||||||
|
close_tx,
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -113,15 +130,24 @@ where
|
|||||||
{
|
{
|
||||||
let (high_tx, mut high_rx) = bounded_queue::<Frame>(HIGH_PRIORITY_QUEUE_CAPACITY);
|
let (high_tx, mut high_rx) = bounded_queue::<Frame>(HIGH_PRIORITY_QUEUE_CAPACITY);
|
||||||
let (normal_tx, mut normal_rx) = bounded_queue::<Frame>(NORMAL_PRIORITY_QUEUE_CAPACITY);
|
let (normal_tx, mut normal_rx) = bounded_queue::<Frame>(NORMAL_PRIORITY_QUEUE_CAPACITY);
|
||||||
let tx = FrameSender { high_tx, normal_tx };
|
let (close_tx, mut close_rx) = watch::channel(false);
|
||||||
|
let tx = FrameSender {
|
||||||
|
high_tx,
|
||||||
|
normal_tx,
|
||||||
|
close_tx,
|
||||||
|
};
|
||||||
|
|
||||||
let handle = tokio::spawn(async move {
|
let handle = tokio::spawn(async move {
|
||||||
let mut ping_ticker = tokio::time::interval(ping_interval);
|
let mut ping_ticker = tokio::time::interval(ping_interval);
|
||||||
let mut high_open = true;
|
let mut high_open = true;
|
||||||
let mut normal_open = true;
|
let mut normal_open = true;
|
||||||
|
let mut close_open = true;
|
||||||
ping_ticker.tick().await; // skip first immediate tick
|
ping_ticker.tick().await; // skip first immediate tick
|
||||||
|
|
||||||
loop {
|
loop {
|
||||||
|
if *close_rx.borrow() {
|
||||||
|
break;
|
||||||
|
}
|
||||||
if let Ok(frame) = high_rx.try_recv() {
|
if let Ok(frame) = high_rx.try_recv() {
|
||||||
if !write_frame(
|
if !write_frame(
|
||||||
&mut sink,
|
&mut sink,
|
||||||
@@ -141,6 +167,10 @@ where
|
|||||||
|
|
||||||
tokio::select! {
|
tokio::select! {
|
||||||
biased;
|
biased;
|
||||||
|
changed = close_rx.changed(), if close_open => {
|
||||||
|
if changed.is_err() { close_open = false; }
|
||||||
|
if *close_rx.borrow() { break; }
|
||||||
|
},
|
||||||
frame = high_rx.recv(), if high_open => {
|
frame = high_rx.recv(), if high_open => {
|
||||||
match frame {
|
match frame {
|
||||||
Some(frame) => {
|
Some(frame) => {
|
||||||
@@ -152,7 +182,7 @@ where
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
_ = ping_ticker.tick(), if high_open || normal_open => {
|
_ = ping_ticker.tick(), if high_open || normal_open => {
|
||||||
if let Err(e) = sink.send(Message::Ping(vec![])).await {
|
if let Err(e) = send_message(&mut sink, Message::Ping(vec![])).await {
|
||||||
error!(error = %e, "failed to send WebSocket ping");
|
error!(error = %e, "failed to send WebSocket ping");
|
||||||
if let Some(metrics) = tunnel_metrics.as_deref() {
|
if let Some(metrics) = tunnel_metrics.as_deref() {
|
||||||
metrics.record_error("ws_ping_error", &e.to_string());
|
metrics.record_error("ws_ping_error", &e.to_string());
|
||||||
@@ -174,7 +204,7 @@ where
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
debug!("writer task exiting");
|
debug!("writer task exiting");
|
||||||
let _ = sink.close().await;
|
let _ = tokio::time::timeout(CLOSE_TIMEOUT, sink.close()).await;
|
||||||
});
|
});
|
||||||
|
|
||||||
(tx, handle)
|
(tx, handle)
|
||||||
@@ -228,7 +258,7 @@ where
|
|||||||
None => frame.encode(),
|
None => frame.encode(),
|
||||||
};
|
};
|
||||||
let wire_len = data.len().max(HEADER_SIZE);
|
let wire_len = data.len().max(HEADER_SIZE);
|
||||||
if let Err(e) = sink.send(Message::Binary(data.into())).await {
|
if let Err(e) = send_message(sink, Message::Binary(data.into())).await {
|
||||||
error!(
|
error!(
|
||||||
stream_id = stream_id,
|
stream_id = stream_id,
|
||||||
msg_type = ?msg_type,
|
msg_type = ?msg_type,
|
||||||
@@ -248,8 +278,92 @@ where
|
|||||||
true
|
true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn send_message<S>(
|
||||||
|
sink: &mut S,
|
||||||
|
message: Message,
|
||||||
|
) -> Result<(), tokio_tungstenite::tungstenite::Error>
|
||||||
|
where
|
||||||
|
S: SinkExt<Message, Error = tokio_tungstenite::tungstenite::Error> + Unpin,
|
||||||
|
{
|
||||||
|
tokio::time::timeout(WRITE_TIMEOUT, sink.send(message))
|
||||||
|
.await
|
||||||
|
.map_err(|_| {
|
||||||
|
tokio_tungstenite::tungstenite::Error::Io(std::io::Error::new(
|
||||||
|
std::io::ErrorKind::TimedOut,
|
||||||
|
"tunnel WebSocket write timed out",
|
||||||
|
))
|
||||||
|
})?
|
||||||
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
|
#[tokio::test]
|
||||||
|
async fn dropping_last_sender_flushes_queued_body_and_end_frames() {
|
||||||
|
let sink = VecSink::default();
|
||||||
|
let sent = Arc::clone(&sink.sent);
|
||||||
|
let (sender, task) = spawn_writer(sink, Duration::from_secs(60));
|
||||||
|
sender
|
||||||
|
.send(Frame::new(
|
||||||
|
7,
|
||||||
|
MsgType::ResponseBody,
|
||||||
|
0,
|
||||||
|
bytes::Bytes::from_static(b"late"),
|
||||||
|
))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
sender
|
||||||
|
.send(Frame::new(7, MsgType::StreamEnd, 0, bytes::Bytes::new()))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
drop(sender);
|
||||||
|
task.await.unwrap();
|
||||||
|
let frames = sent.lock().unwrap();
|
||||||
|
assert_eq!(frames.len(), 2);
|
||||||
|
let Message::Binary(body) = &frames[0] else {
|
||||||
|
panic!("expected body")
|
||||||
|
};
|
||||||
|
assert_eq!(
|
||||||
|
Frame::decode(body.clone().into()).unwrap().payload,
|
||||||
|
b"late".as_slice()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
struct StalledSink;
|
||||||
|
|
||||||
|
impl futures_util::Sink<Message> for StalledSink {
|
||||||
|
type Error = Error;
|
||||||
|
fn poll_ready(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Error>> {
|
||||||
|
Poll::Pending
|
||||||
|
}
|
||||||
|
fn start_send(self: Pin<&mut Self>, _: Message) -> Result<(), Error> {
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Error>> {
|
||||||
|
Poll::Pending
|
||||||
|
}
|
||||||
|
fn poll_close(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Error>> {
|
||||||
|
Poll::Pending
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test(start_paused = true)]
|
||||||
|
async fn stalled_socket_write_and_close_are_bounded() {
|
||||||
|
let (sender, task) = spawn_writer(StalledSink, Duration::from_secs(60));
|
||||||
|
sender
|
||||||
|
.send(Frame::new(
|
||||||
|
1,
|
||||||
|
MsgType::ResponseBody,
|
||||||
|
0,
|
||||||
|
bytes::Bytes::from_static(b"data"),
|
||||||
|
))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
tokio::time::timeout(Duration::from_secs(20), task)
|
||||||
|
.await
|
||||||
|
.expect("writer should time out")
|
||||||
|
.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
use std::pin::Pin;
|
use std::pin::Pin;
|
||||||
use std::sync::{Arc, Mutex};
|
use std::sync::{Arc, Mutex};
|
||||||
use std::task::{Context, Poll};
|
use std::task::{Context, Poll};
|
||||||
|
|||||||
@@ -588,6 +588,68 @@ pub struct SettingsPayload {
|
|||||||
pub drain_deadline_ms: u64,
|
pub drain_deadline_ms: u64,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl SettingsPayload {
|
||||||
|
pub fn is_valid(&self) -> bool {
|
||||||
|
self.initial_stream_window_bytes > 0
|
||||||
|
&& u64::from(self.initial_stream_window_bytes)
|
||||||
|
<= MAX_TUNNEL_DECOMPRESSED_PAYLOAD_BYTES as u64
|
||||||
|
&& self.min_window_update_bytes > 0
|
||||||
|
&& self.min_window_update_bytes <= self.initial_stream_window_bytes
|
||||||
|
&& self.drain_deadline_ms > 0
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn negotiate(&self, initial_window_bytes: u32, drain_deadline_ms: u64) -> Self {
|
||||||
|
let window = self
|
||||||
|
.initial_stream_window_bytes
|
||||||
|
.min(initial_window_bytes)
|
||||||
|
.max(1);
|
||||||
|
Self {
|
||||||
|
initial_stream_window_bytes: window,
|
||||||
|
min_window_update_bytes: self.min_window_update_bytes.min((window / 4).max(1)),
|
||||||
|
drain_deadline_ms: self.drain_deadline_ms.min(drain_deadline_ms),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod settings_tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn negotiation_bounds_window_updates_by_the_smaller_window() {
|
||||||
|
let settings = SettingsPayload {
|
||||||
|
initial_stream_window_bytes: 512 * 1024,
|
||||||
|
min_window_update_bytes: 128 * 1024,
|
||||||
|
drain_deadline_ms: 30_000,
|
||||||
|
};
|
||||||
|
let negotiated = settings.negotiate(4 * 1024 * 1024, 1000);
|
||||||
|
assert_eq!(negotiated.initial_stream_window_bytes, 512 * 1024);
|
||||||
|
assert_eq!(negotiated.min_window_update_bytes, 128 * 1024);
|
||||||
|
assert_eq!(negotiated.drain_deadline_ms, 1000);
|
||||||
|
assert!(negotiated.is_valid());
|
||||||
|
let tiny = settings.negotiate(1, 1);
|
||||||
|
assert_eq!(tiny.initial_stream_window_bytes, 1);
|
||||||
|
assert_eq!(tiny.min_window_update_bytes, 1);
|
||||||
|
assert!(tiny.is_valid());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn invalid_window_settings_are_rejected() {
|
||||||
|
let mut settings = SettingsPayload {
|
||||||
|
initial_stream_window_bytes: 1024,
|
||||||
|
min_window_update_bytes: 256,
|
||||||
|
drain_deadline_ms: 1,
|
||||||
|
};
|
||||||
|
settings.min_window_update_bytes = 1025;
|
||||||
|
assert!(!settings.is_valid());
|
||||||
|
settings.min_window_update_bytes = 0;
|
||||||
|
assert!(!settings.is_valid());
|
||||||
|
settings.initial_stream_window_bytes = u32::MAX;
|
||||||
|
settings.min_window_update_bytes = 1;
|
||||||
|
assert!(!settings.is_valid());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||||
pub struct WindowUpdatePayload {
|
pub struct WindowUpdatePayload {
|
||||||
pub delta_bytes: u32,
|
pub delta_bytes: u32,
|
||||||
|
|||||||
@@ -0,0 +1,323 @@
|
|||||||
|
<template>
|
||||||
|
<Card
|
||||||
|
variant="interactive"
|
||||||
|
class="flex min-w-0 flex-col cursor-pointer overflow-hidden"
|
||||||
|
@mousedown="$emit('mousedown', $event)"
|
||||||
|
@click="$emit('rowClick', $event, provider.id)"
|
||||||
|
>
|
||||||
|
<div class="flex items-start gap-2 p-4 pb-3">
|
||||||
|
<slot name="drag-handle" />
|
||||||
|
<div
|
||||||
|
class="flex h-10 w-10 shrink-0 items-center justify-center rounded-xl text-base font-semibold"
|
||||||
|
:class="provider.is_active ? 'bg-primary/10 text-primary' : 'bg-muted text-muted-foreground'"
|
||||||
|
aria-hidden="true"
|
||||||
|
>
|
||||||
|
{{ provider.name.slice(0, 1).toUpperCase() }}
|
||||||
|
</div>
|
||||||
|
<div class="min-w-0 flex-1 space-y-1">
|
||||||
|
<div class="flex items-center gap-1.5">
|
||||||
|
<button
|
||||||
|
type="button"
|
||||||
|
class="min-w-0 truncate rounded text-left text-sm font-semibold text-foreground hover:text-primary focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-ring"
|
||||||
|
:title="provider.name"
|
||||||
|
@click.stop="$emit('viewDetail', provider.id)"
|
||||||
|
>
|
||||||
|
{{ provider.name }}
|
||||||
|
</button>
|
||||||
|
<a
|
||||||
|
v-if="safeProviderWebsite"
|
||||||
|
:href="safeProviderWebsite"
|
||||||
|
target="_blank"
|
||||||
|
rel="noopener noreferrer"
|
||||||
|
class="shrink-0 text-muted-foreground transition-colors hover:text-primary"
|
||||||
|
:title="safeProviderWebsite"
|
||||||
|
@click.stop
|
||||||
|
>
|
||||||
|
<ExternalLink class="h-3.5 w-3.5" />
|
||||||
|
</a>
|
||||||
|
</div>
|
||||||
|
<div
|
||||||
|
v-if="editingDescriptionId === provider.id"
|
||||||
|
data-desc-editor
|
||||||
|
class="flex items-center gap-1"
|
||||||
|
@click.stop
|
||||||
|
>
|
||||||
|
<input
|
||||||
|
v-model="localDescriptionValue"
|
||||||
|
v-auto-focus
|
||||||
|
class="min-w-0 flex-1 rounded border border-border bg-background px-1.5 py-0.5 text-xs text-foreground focus:outline-none focus:ring-1 focus:ring-primary/50"
|
||||||
|
:placeholder="legacyT('输入备注...')"
|
||||||
|
:aria-label="legacyT('输入备注...')"
|
||||||
|
@keydown="handleDescriptionKeydown"
|
||||||
|
>
|
||||||
|
<button
|
||||||
|
type="button"
|
||||||
|
class="shrink-0 rounded p-0.5 text-primary hover:bg-muted"
|
||||||
|
:title="legacyT('保存')"
|
||||||
|
:aria-label="legacyT('保存')"
|
||||||
|
@click="handleSave"
|
||||||
|
>
|
||||||
|
<Check class="h-3.5 w-3.5" />
|
||||||
|
</button>
|
||||||
|
<button
|
||||||
|
type="button"
|
||||||
|
class="shrink-0 rounded p-0.5 text-muted-foreground hover:bg-muted"
|
||||||
|
:title="legacyT('取消')"
|
||||||
|
:aria-label="legacyT('取消')"
|
||||||
|
@click="handleCancel"
|
||||||
|
>
|
||||||
|
<X class="h-3.5 w-3.5" />
|
||||||
|
</button>
|
||||||
|
</div>
|
||||||
|
<button
|
||||||
|
v-else
|
||||||
|
type="button"
|
||||||
|
class="group/desc flex max-w-full items-center gap-1 rounded text-xs text-muted-foreground transition-colors hover:text-foreground focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-ring"
|
||||||
|
:title="provider.description || legacyT('添加备注')"
|
||||||
|
@click="handleStartEdit"
|
||||||
|
>
|
||||||
|
<span class="truncate">{{ provider.description || legacyT('添加备注') }}</span>
|
||||||
|
<Pencil class="h-3 w-3 shrink-0 opacity-0 transition-opacity group-hover/desc:opacity-60 group-focus-visible/desc:opacity-60" />
|
||||||
|
</button>
|
||||||
|
</div>
|
||||||
|
<Badge
|
||||||
|
:variant="provider.is_active ? 'success' : 'secondary'"
|
||||||
|
class="shrink-0 text-xs"
|
||||||
|
>
|
||||||
|
{{ legacyT(provider.is_active ? '活跃' : '停用') }}
|
||||||
|
</Badge>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<div class="flex flex-1 flex-col gap-4 px-4 pb-4">
|
||||||
|
<div class="space-y-2 rounded-xl border border-border/40 bg-muted/20 p-3">
|
||||||
|
<div class="flex flex-wrap items-center justify-between gap-2">
|
||||||
|
<span class="text-xs text-muted-foreground">{{ legacyT('余额监控') }}</span>
|
||||||
|
<Badge
|
||||||
|
variant="outline"
|
||||||
|
class="border-border/50 text-[10px] font-normal"
|
||||||
|
>
|
||||||
|
{{ formatBillingType(provider.billing_type || 'pay_as_you_go') }}
|
||||||
|
</Badge>
|
||||||
|
</div>
|
||||||
|
<ProviderBalanceCell
|
||||||
|
:provider="provider"
|
||||||
|
:is-balance-loading="isBalanceLoading"
|
||||||
|
:get-provider-balance="getProviderBalance"
|
||||||
|
:get-provider-balance-breakdown="getProviderBalanceBreakdown"
|
||||||
|
:get-provider-balance-error="getProviderBalanceError"
|
||||||
|
:get-provider-checkin="getProviderCheckin"
|
||||||
|
:get-provider-cookie-expired="getProviderCookieExpired"
|
||||||
|
:get-provider-balance-extra="getProviderBalanceExtra"
|
||||||
|
:format-balance-display="formatBalanceDisplay"
|
||||||
|
:format-reset-countdown="formatResetCountdown"
|
||||||
|
:get-quota-used-color-class="getQuotaUsedColorClass"
|
||||||
|
/>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<dl class="grid grid-cols-3 divide-x divide-border/50 text-center">
|
||||||
|
<div class="min-w-0 space-y-1 px-1">
|
||||||
|
<dt class="text-xs text-muted-foreground">
|
||||||
|
{{ legacyT('端点') }}
|
||||||
|
</dt>
|
||||||
|
<dd class="text-sm font-semibold tabular-nums">
|
||||||
|
{{ provider.active_endpoints }}<span class="ml-0.5 text-xs font-normal text-muted-foreground">/ {{ provider.total_endpoints }}</span>
|
||||||
|
</dd>
|
||||||
|
</div>
|
||||||
|
<div class="min-w-0 space-y-1 px-1">
|
||||||
|
<dt class="text-xs text-muted-foreground">
|
||||||
|
{{ legacyT(isKeyManagedProviderType(provider.provider_type) ? '密钥' : '账号') }}
|
||||||
|
</dt>
|
||||||
|
<dd class="text-sm font-semibold tabular-nums">
|
||||||
|
{{ provider.active_keys }}<span class="ml-0.5 text-xs font-normal text-muted-foreground">/ {{ provider.total_keys }}</span>
|
||||||
|
</dd>
|
||||||
|
</div>
|
||||||
|
<div class="min-w-0 space-y-1 px-1">
|
||||||
|
<dt class="text-xs text-muted-foreground">
|
||||||
|
{{ legacyT('模型') }}
|
||||||
|
</dt>
|
||||||
|
<dd class="text-sm font-semibold tabular-nums">
|
||||||
|
{{ provider.active_models }}<span class="ml-0.5 text-xs font-normal text-muted-foreground">/ {{ provider.total_models }}</span>
|
||||||
|
</dd>
|
||||||
|
</div>
|
||||||
|
</dl>
|
||||||
|
|
||||||
|
<div class="mt-auto space-y-2 border-t border-border/40 pt-3">
|
||||||
|
<div class="text-xs text-muted-foreground">
|
||||||
|
{{ legacyT('端点健康') }}
|
||||||
|
</div>
|
||||||
|
<div
|
||||||
|
v-if="provider.endpoint_health_details?.length"
|
||||||
|
class="grid grid-cols-3 gap-x-3 gap-y-2"
|
||||||
|
>
|
||||||
|
<div
|
||||||
|
v-for="endpoint in sortEndpoints(provider.endpoint_health_details)"
|
||||||
|
:key="endpoint.api_format"
|
||||||
|
class="flex min-w-0 flex-col gap-1.5"
|
||||||
|
:title="getEndpointTooltip(endpoint, locale)"
|
||||||
|
>
|
||||||
|
<div class="flex items-center justify-between gap-1 text-[10px] leading-none text-muted-foreground">
|
||||||
|
<span class="font-medium">{{ formatApiFormatShort(endpoint.api_format) }}</span>
|
||||||
|
<span class="tabular-nums">{{ getEndpointHealthLabel(endpoint) }}</span>
|
||||||
|
</div>
|
||||||
|
<div class="h-1.5 w-full overflow-hidden rounded-full bg-border dark:bg-border/80">
|
||||||
|
<div
|
||||||
|
class="h-full rounded-full transition-all duration-300"
|
||||||
|
:class="getEndpointDotColor(endpoint)"
|
||||||
|
:style="{ width: getEndpointHealthBarWidth(endpoint) }"
|
||||||
|
/>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
<span
|
||||||
|
v-else
|
||||||
|
class="text-xs text-muted-foreground/60"
|
||||||
|
>{{ legacyT('暂无端点') }}</span>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<div
|
||||||
|
class="flex items-center justify-between gap-2 border-t border-border/40 bg-muted/10 px-3 py-2"
|
||||||
|
@click.stop
|
||||||
|
>
|
||||||
|
<Button
|
||||||
|
variant="ghost"
|
||||||
|
size="icon"
|
||||||
|
class="h-8 w-8 text-muted-foreground hover:text-primary"
|
||||||
|
:title="legacyT('查看详情')"
|
||||||
|
:aria-label="legacyT('查看详情')"
|
||||||
|
@click="$emit('viewDetail', provider.id)"
|
||||||
|
>
|
||||||
|
<Eye class="h-3.5 w-3.5" />
|
||||||
|
</Button>
|
||||||
|
<div class="flex shrink-0 items-center gap-0.5">
|
||||||
|
<Button
|
||||||
|
variant="ghost"
|
||||||
|
size="icon"
|
||||||
|
class="h-8 w-8 text-muted-foreground hover:text-foreground"
|
||||||
|
:title="legacyT('编辑提供商')"
|
||||||
|
:aria-label="legacyT('编辑提供商')"
|
||||||
|
@click="$emit('editProvider', provider)"
|
||||||
|
>
|
||||||
|
<Edit class="h-3.5 w-3.5" />
|
||||||
|
</Button>
|
||||||
|
<Button
|
||||||
|
variant="ghost"
|
||||||
|
size="icon"
|
||||||
|
class="h-8 w-8 text-muted-foreground hover:text-foreground"
|
||||||
|
:title="legacyT('扩展操作配置')"
|
||||||
|
:aria-label="legacyT('扩展操作配置')"
|
||||||
|
@click="$emit('openOpsConfig', provider)"
|
||||||
|
>
|
||||||
|
<KeyRound class="h-3.5 w-3.5" />
|
||||||
|
</Button>
|
||||||
|
<Button
|
||||||
|
variant="ghost"
|
||||||
|
size="icon"
|
||||||
|
class="h-8 w-8 text-muted-foreground hover:text-foreground"
|
||||||
|
:title="legacyT(provider.is_active ? '停用提供商' : '启用提供商')"
|
||||||
|
:aria-label="legacyT(provider.is_active ? '停用提供商' : '启用提供商')"
|
||||||
|
@click="$emit('toggleStatus', provider)"
|
||||||
|
>
|
||||||
|
<Power class="h-3.5 w-3.5" />
|
||||||
|
</Button>
|
||||||
|
<Button
|
||||||
|
variant="ghost"
|
||||||
|
size="icon"
|
||||||
|
class="h-8 w-8 text-muted-foreground hover:text-destructive"
|
||||||
|
:title="legacyT('删除提供商')"
|
||||||
|
:aria-label="legacyT('删除提供商')"
|
||||||
|
@click="$emit('deleteProvider', provider)"
|
||||||
|
>
|
||||||
|
<Trash2 class="h-3.5 w-3.5" />
|
||||||
|
</Button>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</Card>
|
||||||
|
</template>
|
||||||
|
|
||||||
|
<script setup lang="ts">
|
||||||
|
import { computed, ref, watch } from 'vue'
|
||||||
|
import { Check, Edit, ExternalLink, Eye, KeyRound, Pencil, Power, Trash2, X } from 'lucide-vue-next'
|
||||||
|
import Badge from '@/components/ui/badge.vue'
|
||||||
|
import Button from '@/components/ui/button.vue'
|
||||||
|
import Card from '@/components/ui/card.vue'
|
||||||
|
import ProviderBalanceCell from './ProviderBalanceCell.vue'
|
||||||
|
import { formatApiFormatShort, type ProviderWithEndpointsSummary } from '@/api/endpoints'
|
||||||
|
import type { BalanceExtraItem } from '@/features/providers/auth-templates'
|
||||||
|
import {
|
||||||
|
sortEndpoints,
|
||||||
|
getEndpointHealthLabel,
|
||||||
|
getEndpointHealthBarWidth,
|
||||||
|
getEndpointDotColor,
|
||||||
|
getEndpointTooltip,
|
||||||
|
} from '@/features/providers/composables/useEndpointStatus'
|
||||||
|
import { isKeyManagedProviderType } from '../utils/providerTypeUtils'
|
||||||
|
import { formatBillingType } from '@/utils/format'
|
||||||
|
import { safeExternalWebUrl } from '@/utils/navigationSecurity'
|
||||||
|
import { useI18n } from '@/i18n'
|
||||||
|
|
||||||
|
const props = defineProps<{
|
||||||
|
provider: ProviderWithEndpointsSummary
|
||||||
|
editingDescriptionId: string | null
|
||||||
|
isBalanceLoading: (providerId: string) => boolean
|
||||||
|
getProviderBalance: (providerId: string) => { available: number | null; currency: string } | null
|
||||||
|
getProviderBalanceBreakdown: (providerId: string) => { balance: number; points: number; currency: string } | null
|
||||||
|
getProviderBalanceError: (providerId: string) => { status: string; message: string } | null
|
||||||
|
getProviderCheckin: (providerId: string) => { success: boolean | null; message: string } | null
|
||||||
|
getProviderCookieExpired: (providerId: string) => { expired: boolean; message: string } | null
|
||||||
|
getProviderBalanceExtra: (providerId: string, architectureId?: string) => BalanceExtraItem[]
|
||||||
|
formatBalanceDisplay: (balance: { available: number | null; currency: string } | null) => string
|
||||||
|
formatResetCountdown: (resetsAt: number) => string
|
||||||
|
getQuotaUsedColorClass: (provider: ProviderWithEndpointsSummary) => string
|
||||||
|
}>()
|
||||||
|
|
||||||
|
const emit = defineEmits<{
|
||||||
|
'mousedown': [event: MouseEvent]
|
||||||
|
'rowClick': [event: MouseEvent, providerId: string]
|
||||||
|
'viewDetail': [providerId: string]
|
||||||
|
'editProvider': [provider: ProviderWithEndpointsSummary]
|
||||||
|
'openOpsConfig': [provider: ProviderWithEndpointsSummary]
|
||||||
|
'toggleStatus': [provider: ProviderWithEndpointsSummary]
|
||||||
|
'deleteProvider': [provider: ProviderWithEndpointsSummary]
|
||||||
|
'startEditDescription': [event: Event, provider: ProviderWithEndpointsSummary]
|
||||||
|
'saveDescription': [event: Event, provider: ProviderWithEndpointsSummary, value: string]
|
||||||
|
'cancelEditDescription': [event?: Event]
|
||||||
|
}>()
|
||||||
|
|
||||||
|
const { legacyT, locale } = useI18n()
|
||||||
|
const safeProviderWebsite = computed(() => safeExternalWebUrl(props.provider.website))
|
||||||
|
const localDescriptionValue = ref('')
|
||||||
|
const vAutoFocus = {
|
||||||
|
mounted: (element: HTMLElement) => element.focus(),
|
||||||
|
}
|
||||||
|
|
||||||
|
watch(() => props.editingDescriptionId, (providerId) => {
|
||||||
|
if (providerId === props.provider.id) {
|
||||||
|
localDescriptionValue.value = props.provider.description || ''
|
||||||
|
}
|
||||||
|
}, { immediate: true })
|
||||||
|
|
||||||
|
function handleStartEdit(event: Event) {
|
||||||
|
event.stopPropagation()
|
||||||
|
emit('startEditDescription', event, props.provider)
|
||||||
|
}
|
||||||
|
|
||||||
|
function handleSave(event: Event) {
|
||||||
|
event.stopPropagation()
|
||||||
|
emit('saveDescription', event, props.provider, localDescriptionValue.value)
|
||||||
|
}
|
||||||
|
|
||||||
|
function handleCancel(event: Event) {
|
||||||
|
event.stopPropagation()
|
||||||
|
emit('cancelEditDescription', event)
|
||||||
|
}
|
||||||
|
|
||||||
|
function handleDescriptionKeydown(event: KeyboardEvent) {
|
||||||
|
if (event.key === 'Enter') {
|
||||||
|
event.preventDefault()
|
||||||
|
handleSave(event)
|
||||||
|
} else if (event.key === 'Escape') {
|
||||||
|
handleCancel(event)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
</script>
|
||||||
@@ -0,0 +1,36 @@
|
|||||||
|
<template>
|
||||||
|
<Button
|
||||||
|
variant="ghost"
|
||||||
|
size="icon"
|
||||||
|
class="h-8 w-5 touch-none select-none cursor-grab text-muted-foreground/50 hover:text-primary active:cursor-grabbing disabled:cursor-default"
|
||||||
|
:disabled="disabled"
|
||||||
|
:title="legacyT('拖动调整当前页展示顺序,也可使用方向键移动')"
|
||||||
|
:aria-label="`${legacyT('调整展示顺序')}: ${providerName}`"
|
||||||
|
data-provider-drag-handle
|
||||||
|
@pointerdown.stop="$emit('pointerdown', $event)"
|
||||||
|
@keydown.stop="$emit('keydown', $event)"
|
||||||
|
@mousedown.stop
|
||||||
|
@click.stop
|
||||||
|
@dragstart.prevent
|
||||||
|
>
|
||||||
|
<GripVertical class="h-4 w-4" />
|
||||||
|
</Button>
|
||||||
|
</template>
|
||||||
|
|
||||||
|
<script setup lang="ts">
|
||||||
|
import { GripVertical } from 'lucide-vue-next'
|
||||||
|
import Button from '@/components/ui/button.vue'
|
||||||
|
import { useI18n } from '@/i18n'
|
||||||
|
|
||||||
|
defineProps<{
|
||||||
|
providerName: string
|
||||||
|
disabled: boolean
|
||||||
|
}>()
|
||||||
|
|
||||||
|
defineEmits<{
|
||||||
|
pointerdown: [event: PointerEvent]
|
||||||
|
keydown: [event: KeyboardEvent]
|
||||||
|
}>()
|
||||||
|
|
||||||
|
const { legacyT } = useI18n()
|
||||||
|
</script>
|
||||||
@@ -5,6 +5,7 @@
|
|||||||
>
|
>
|
||||||
<!-- 第一行:名称 + 状态 + 操作 -->
|
<!-- 第一行:名称 + 状态 + 操作 -->
|
||||||
<div class="flex items-start justify-between gap-3">
|
<div class="flex items-start justify-between gap-3">
|
||||||
|
<slot name="drag-handle" />
|
||||||
<div class="flex-1 min-w-0 space-y-0.5">
|
<div class="flex-1 min-w-0 space-y-0.5">
|
||||||
<div class="flex items-center gap-1.5">
|
<div class="flex items-center gap-1.5">
|
||||||
<span class="font-medium text-foreground truncate">{{ provider.name }}</span>
|
<span class="font-medium text-foreground truncate">{{ provider.name }}</span>
|
||||||
|
|||||||
@@ -22,7 +22,7 @@
|
|||||||
</div>
|
</div>
|
||||||
|
|
||||||
<!-- 状态筛选 -->
|
<!-- 状态筛选 -->
|
||||||
<div class="xl:hidden">
|
<div :class="{ 'xl:hidden': !cardView }">
|
||||||
<Select
|
<Select
|
||||||
:model-value="filterStatus"
|
:model-value="filterStatus"
|
||||||
@update:model-value="$emit('update:filterStatus', $event)"
|
@update:model-value="$emit('update:filterStatus', $event)"
|
||||||
@@ -43,7 +43,7 @@
|
|||||||
</div>
|
</div>
|
||||||
|
|
||||||
<!-- API 格式筛选 -->
|
<!-- API 格式筛选 -->
|
||||||
<div class="xl:hidden">
|
<div :class="{ 'xl:hidden': !cardView }">
|
||||||
<Select
|
<Select
|
||||||
:model-value="filterApiFormat"
|
:model-value="filterApiFormat"
|
||||||
@update:model-value="$emit('update:filterApiFormat', $event)"
|
@update:model-value="$emit('update:filterApiFormat', $event)"
|
||||||
@@ -64,7 +64,7 @@
|
|||||||
</div>
|
</div>
|
||||||
|
|
||||||
<!-- 模型筛选 -->
|
<!-- 模型筛选 -->
|
||||||
<div class="xl:hidden">
|
<div :class="{ 'xl:hidden': !cardView }">
|
||||||
<Select
|
<Select
|
||||||
:model-value="filterModel"
|
:model-value="filterModel"
|
||||||
@update:model-value="$emit('update:filterModel', $event)"
|
@update:model-value="$emit('update:filterModel', $event)"
|
||||||
@@ -122,13 +122,32 @@
|
|||||||
:loading="loading"
|
:loading="loading"
|
||||||
@click="$emit('refresh')"
|
@click="$emit('refresh')"
|
||||||
/>
|
/>
|
||||||
|
<Button
|
||||||
|
variant="ghost"
|
||||||
|
size="icon"
|
||||||
|
class="h-8 w-8"
|
||||||
|
:class="{ 'bg-primary/10 text-primary hover:bg-primary/15 hover:text-primary': cardView }"
|
||||||
|
:title="legacyT(cardView ? '切换到列表视图' : '切换到卡片视图')"
|
||||||
|
:aria-label="legacyT('卡片视图')"
|
||||||
|
:aria-pressed="cardView"
|
||||||
|
@click="$emit('toggleView')"
|
||||||
|
>
|
||||||
|
<List
|
||||||
|
v-if="cardView"
|
||||||
|
class="w-3.5 h-3.5"
|
||||||
|
/>
|
||||||
|
<LayoutGrid
|
||||||
|
v-else
|
||||||
|
class="w-3.5 h-3.5"
|
||||||
|
/>
|
||||||
|
</Button>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
</template>
|
</template>
|
||||||
|
|
||||||
<script setup lang="ts">
|
<script setup lang="ts">
|
||||||
import { Search, Plus, FilterX, Users } from 'lucide-vue-next'
|
import { Search, Plus, FilterX, Users, LayoutGrid, List } from 'lucide-vue-next'
|
||||||
import Button from '@/components/ui/button.vue'
|
import Button from '@/components/ui/button.vue'
|
||||||
import Input from '@/components/ui/input.vue'
|
import Input from '@/components/ui/input.vue'
|
||||||
import Select from '@/components/ui/select.vue'
|
import Select from '@/components/ui/select.vue'
|
||||||
@@ -150,6 +169,7 @@ defineProps<{
|
|||||||
modelFilters: FilterOption[]
|
modelFilters: FilterOption[]
|
||||||
hasActiveFilters: boolean
|
hasActiveFilters: boolean
|
||||||
loading: boolean
|
loading: boolean
|
||||||
|
cardView: boolean
|
||||||
}>()
|
}>()
|
||||||
|
|
||||||
defineEmits<{
|
defineEmits<{
|
||||||
@@ -161,6 +181,7 @@ defineEmits<{
|
|||||||
'batchProcess': []
|
'batchProcess': []
|
||||||
'addProvider': []
|
'addProvider': []
|
||||||
'refresh': []
|
'refresh': []
|
||||||
|
'toggleView': []
|
||||||
}>()
|
}>()
|
||||||
|
|
||||||
const { legacyT } = useI18n()
|
const { legacyT } = useI18n()
|
||||||
|
|||||||
@@ -4,6 +4,13 @@
|
|||||||
@mousedown="$emit('mousedown', $event)"
|
@mousedown="$emit('mousedown', $event)"
|
||||||
@click="$emit('rowClick', $event, provider.id)"
|
@click="$emit('rowClick', $event, provider.id)"
|
||||||
>
|
>
|
||||||
|
<TableCell
|
||||||
|
v-if="$slots['drag-handle']"
|
||||||
|
class="w-9 px-2 py-3.5"
|
||||||
|
@click.stop
|
||||||
|
>
|
||||||
|
<slot name="drag-handle" />
|
||||||
|
</TableCell>
|
||||||
<TableCell class="py-3.5">
|
<TableCell class="py-3.5">
|
||||||
<div class="space-y-0.5">
|
<div class="space-y-0.5">
|
||||||
<div class="flex items-center gap-1.5">
|
<div class="flex items-center gap-1.5">
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ import type { ProviderWithEndpointsSummary } from '@/api/endpoints'
|
|||||||
import { createI18n } from '@/i18n'
|
import { createI18n } from '@/i18n'
|
||||||
import ProviderTableRow from '../ProviderTableRow.vue'
|
import ProviderTableRow from '../ProviderTableRow.vue'
|
||||||
import ProviderMobileCard from '../ProviderMobileCard.vue'
|
import ProviderMobileCard from '../ProviderMobileCard.vue'
|
||||||
|
import ProviderCard from '../ProviderCard.vue'
|
||||||
|
|
||||||
vi.mock('../ProviderBalanceCell.vue', () => ({
|
vi.mock('../ProviderBalanceCell.vue', () => ({
|
||||||
default: { render: () => null },
|
default: { render: () => null },
|
||||||
@@ -74,6 +75,7 @@ function mountProvider(component: Component, healthScore: number | null) {
|
|||||||
describe.each([
|
describe.each([
|
||||||
['desktop provider row', ProviderTableRow],
|
['desktop provider row', ProviderTableRow],
|
||||||
['mobile provider card', ProviderMobileCard],
|
['mobile provider card', ProviderMobileCard],
|
||||||
|
['grid provider card', ProviderCard],
|
||||||
] as const)('%s endpoint health', (_name, component) => {
|
] as const)('%s endpoint health', (_name, component) => {
|
||||||
it.each([
|
it.each([
|
||||||
{ score: null, label: '-', width: '100%', color: 'bg-muted-foreground/40' },
|
{ score: null, label: '-', width: '100%', color: 'bg-muted-foreground/40' },
|
||||||
|
|||||||
@@ -0,0 +1,224 @@
|
|||||||
|
import { computed, nextTick, onScopeDispose, ref, watch, type Ref } from 'vue'
|
||||||
|
import { useEventListener, useLocalStorage, useRafFn } from '@vueuse/core'
|
||||||
|
import { useI18n } from '@/i18n'
|
||||||
|
|
||||||
|
interface SortableProvider {
|
||||||
|
id: string
|
||||||
|
name: string
|
||||||
|
}
|
||||||
|
|
||||||
|
interface ProviderPointerDrag {
|
||||||
|
providerId: string
|
||||||
|
pointerId: number
|
||||||
|
startX: number
|
||||||
|
startY: number
|
||||||
|
handle: HTMLElement
|
||||||
|
scrollContainer: HTMLElement | null
|
||||||
|
}
|
||||||
|
|
||||||
|
export function useProviderDisplayOrder<Provider extends SortableProvider>(
|
||||||
|
providers: () => Provider[],
|
||||||
|
container: Ref<HTMLElement | null>,
|
||||||
|
) {
|
||||||
|
const { legacyT } = useI18n()
|
||||||
|
const savedOrder = useLocalStorage<string[]>('aether-provider-display-order', [])
|
||||||
|
const knownOrder = ref<string[]>([])
|
||||||
|
const draggingProviderId = ref<string | null>(null)
|
||||||
|
const dropTargetId = ref<string | null>(null)
|
||||||
|
const pointerPosition = ref({ clientX: 0, clientY: 0 })
|
||||||
|
const announcement = ref('')
|
||||||
|
let pointerDrag: ProviderPointerDrag | null = null
|
||||||
|
let suppressClickUntil = 0
|
||||||
|
|
||||||
|
const normalizedOrder = computed(() => Array.isArray(savedOrder.value)
|
||||||
|
? [...new Set(savedOrder.value.filter((providerId): providerId is string => typeof providerId === 'string'))]
|
||||||
|
: [])
|
||||||
|
|
||||||
|
const orderedProviders = computed(() => {
|
||||||
|
const ranks = new Map(normalizedOrder.value.map((providerId, index) => [providerId, index]))
|
||||||
|
return [...providers()].sort((first, second) => (
|
||||||
|
(ranks.get(first.id) ?? ranks.size) - (ranks.get(second.id) ?? ranks.size)
|
||||||
|
))
|
||||||
|
})
|
||||||
|
|
||||||
|
const draggingProvider = computed(() => orderedProviders.value.find(provider => provider.id === draggingProviderId.value))
|
||||||
|
const dragPreviewStyle = computed(() => ({
|
||||||
|
left: `${Math.max(8, Math.min(pointerPosition.value.clientX + 12, window.innerWidth - 208))}px`,
|
||||||
|
top: `${Math.max(8, Math.min(pointerPosition.value.clientY + 12, window.innerHeight - 48))}px`,
|
||||||
|
}))
|
||||||
|
|
||||||
|
function moveProvider(providerId: string, targetId: string) {
|
||||||
|
const visibleIds = orderedProviders.value.map(provider => provider.id)
|
||||||
|
const sourceIndex = visibleIds.indexOf(providerId)
|
||||||
|
const targetIndex = visibleIds.indexOf(targetId)
|
||||||
|
if (sourceIndex < 0 || targetIndex < 0 || sourceIndex === targetIndex) return
|
||||||
|
|
||||||
|
visibleIds.splice(sourceIndex, 1)
|
||||||
|
visibleIds.splice(targetIndex, 0, providerId)
|
||||||
|
const visibleSet = new Set(visibleIds)
|
||||||
|
const allIds = [...new Set([...normalizedOrder.value, ...knownOrder.value, ...visibleIds])]
|
||||||
|
let visibleIndex = 0
|
||||||
|
savedOrder.value = allIds.map(currentId => visibleSet.has(currentId) ? visibleIds[visibleIndex++] ?? currentId : currentId)
|
||||||
|
announcement.value = `${legacyT('展示顺序已更新')}: ${orderedProviders.value[targetIndex]?.name} (${targetIndex + 1}/${visibleIds.length})`
|
||||||
|
}
|
||||||
|
|
||||||
|
function updateDropTarget() {
|
||||||
|
const target = document.elementFromPoint(pointerPosition.value.clientX, pointerPosition.value.clientY)
|
||||||
|
?.closest<HTMLElement>('[data-provider-sort-id]')
|
||||||
|
const targetId = target?.dataset.providerSortId
|
||||||
|
dropTargetId.value = target && container.value?.contains(target)
|
||||||
|
&& targetId !== draggingProviderId.value
|
||||||
|
&& orderedProviders.value.some(provider => provider.id === targetId)
|
||||||
|
? targetId ?? null
|
||||||
|
: null
|
||||||
|
}
|
||||||
|
|
||||||
|
function findScrollContainer(handle: HTMLElement): HTMLElement | null {
|
||||||
|
let ancestor = handle.parentElement
|
||||||
|
while (ancestor && ancestor !== document.body) {
|
||||||
|
if (/(auto|scroll)/.test(getComputedStyle(ancestor).overflowY) && ancestor.scrollHeight > ancestor.clientHeight) {
|
||||||
|
return ancestor
|
||||||
|
}
|
||||||
|
ancestor = ancestor.parentElement
|
||||||
|
}
|
||||||
|
return document.scrollingElement as HTMLElement | null
|
||||||
|
}
|
||||||
|
|
||||||
|
const { pause, resume } = useRafFn(() => {
|
||||||
|
const scrollContainer = pointerDrag?.scrollContainer
|
||||||
|
if (scrollContainer) {
|
||||||
|
const bounds = scrollContainer.getBoundingClientRect()
|
||||||
|
const isDocument = scrollContainer === document.scrollingElement
|
||||||
|
const top = isDocument ? 0 : Math.max(0, bounds.top)
|
||||||
|
const bottom = isDocument ? window.innerHeight : Math.min(window.innerHeight, bounds.bottom)
|
||||||
|
const pointerY = pointerPosition.value.clientY
|
||||||
|
const pointerX = pointerPosition.value.clientX
|
||||||
|
if (isDocument || (pointerX >= bounds.left && pointerX <= bounds.right)) {
|
||||||
|
if (pointerY < top + 48) {
|
||||||
|
scrollContainer.scrollTop -= Math.min(12, (top + 48 - pointerY) / 4)
|
||||||
|
} else if (pointerY > bottom - 48) {
|
||||||
|
scrollContainer.scrollTop += Math.min(12, (pointerY - bottom + 48) / 4)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
updateDropTarget()
|
||||||
|
}, { immediate: false })
|
||||||
|
|
||||||
|
function cancelDrag() {
|
||||||
|
const previous = pointerDrag
|
||||||
|
pointerDrag = null
|
||||||
|
pause()
|
||||||
|
if (draggingProviderId.value) suppressClickUntil = Date.now() + 250
|
||||||
|
draggingProviderId.value = null
|
||||||
|
dropTargetId.value = null
|
||||||
|
if (previous?.handle.hasPointerCapture?.(previous.pointerId)) {
|
||||||
|
previous.handle.releasePointerCapture(previous.pointerId)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
function startDrag(providerId: string, event: PointerEvent) {
|
||||||
|
if (event.button !== 0 || event.isPrimary === false || orderedProviders.value.length < 2) return
|
||||||
|
if (!orderedProviders.value.some(provider => provider.id === providerId)) return
|
||||||
|
const handle = event.currentTarget
|
||||||
|
if (!(handle instanceof HTMLElement)) return
|
||||||
|
|
||||||
|
cancelDrag()
|
||||||
|
event.preventDefault()
|
||||||
|
handle.focus({ preventScroll: true })
|
||||||
|
pointerDrag = {
|
||||||
|
providerId,
|
||||||
|
pointerId: event.pointerId,
|
||||||
|
startX: event.clientX,
|
||||||
|
startY: event.clientY,
|
||||||
|
handle,
|
||||||
|
scrollContainer: findScrollContainer(handle),
|
||||||
|
}
|
||||||
|
pointerPosition.value = { clientX: event.clientX, clientY: event.clientY }
|
||||||
|
handle.setPointerCapture?.(event.pointerId)
|
||||||
|
}
|
||||||
|
|
||||||
|
function handlePointerMove(event: PointerEvent) {
|
||||||
|
if (!pointerDrag || pointerDrag.pointerId !== event.pointerId) return
|
||||||
|
pointerPosition.value = { clientX: event.clientX, clientY: event.clientY }
|
||||||
|
if (!draggingProviderId.value) {
|
||||||
|
if (Math.hypot(event.clientX - pointerDrag.startX, event.clientY - pointerDrag.startY) < 5) return
|
||||||
|
draggingProviderId.value = pointerDrag.providerId
|
||||||
|
resume()
|
||||||
|
}
|
||||||
|
event.preventDefault()
|
||||||
|
updateDropTarget()
|
||||||
|
}
|
||||||
|
|
||||||
|
function handlePointerUp(event: PointerEvent) {
|
||||||
|
if (!pointerDrag || pointerDrag.pointerId !== event.pointerId) return
|
||||||
|
const providerId = draggingProviderId.value
|
||||||
|
if (providerId) {
|
||||||
|
pointerPosition.value = { clientX: event.clientX, clientY: event.clientY }
|
||||||
|
updateDropTarget()
|
||||||
|
}
|
||||||
|
const targetId = dropTargetId.value
|
||||||
|
cancelDrag()
|
||||||
|
if (providerId && targetId) moveProvider(providerId, targetId)
|
||||||
|
}
|
||||||
|
|
||||||
|
function handleSortKeydown(providerId: string, event: KeyboardEvent) {
|
||||||
|
if (event.key === 'Escape') {
|
||||||
|
cancelDrag()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
const directions: Record<string, number> = { ArrowUp: -1, ArrowLeft: -1, ArrowDown: 1, ArrowRight: 1 }
|
||||||
|
const direction = directions[event.key]
|
||||||
|
if (direction === undefined || pointerDrag) return
|
||||||
|
event.preventDefault()
|
||||||
|
const index = orderedProviders.value.findIndex(provider => provider.id === providerId)
|
||||||
|
const target = orderedProviders.value[index + direction]
|
||||||
|
if (index < 0 || !target) return
|
||||||
|
const handle = event.currentTarget as HTMLElement
|
||||||
|
moveProvider(providerId, target.id)
|
||||||
|
void nextTick(() => handle.focus({ preventScroll: true }))
|
||||||
|
}
|
||||||
|
|
||||||
|
function handleSortClick(event: MouseEvent) {
|
||||||
|
if (!draggingProviderId.value && Date.now() >= suppressClickUntil) return
|
||||||
|
suppressClickUntil = 0
|
||||||
|
event.preventDefault()
|
||||||
|
event.stopPropagation()
|
||||||
|
}
|
||||||
|
|
||||||
|
function sortItemClass(providerId: string) {
|
||||||
|
return {
|
||||||
|
'opacity-40': draggingProviderId.value === providerId,
|
||||||
|
'ring-2 ring-inset ring-primary/60 bg-primary/5': dropTargetId.value === providerId,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
watch(() => providers().map(provider => provider.id), (providerIds) => {
|
||||||
|
knownOrder.value = [...new Set([...knownOrder.value, ...providerIds])]
|
||||||
|
cancelDrag()
|
||||||
|
}, { immediate: true })
|
||||||
|
useEventListener(window, 'pointermove', handlePointerMove, { passive: false })
|
||||||
|
useEventListener(window, 'pointerup', handlePointerUp)
|
||||||
|
useEventListener(window, 'pointercancel', (event) => {
|
||||||
|
if (pointerDrag?.pointerId === event.pointerId) cancelDrag()
|
||||||
|
})
|
||||||
|
useEventListener(window, 'lostpointercapture', (event) => {
|
||||||
|
if (pointerDrag?.pointerId === event.pointerId) cancelDrag()
|
||||||
|
})
|
||||||
|
useEventListener(window, 'blur', cancelDrag)
|
||||||
|
useEventListener(window, 'keydown', (event) => {
|
||||||
|
if (event.key === 'Escape') cancelDrag()
|
||||||
|
})
|
||||||
|
onScopeDispose(cancelDrag)
|
||||||
|
|
||||||
|
return {
|
||||||
|
orderedProviders,
|
||||||
|
draggingProvider,
|
||||||
|
dragPreviewStyle,
|
||||||
|
announcement,
|
||||||
|
startDrag,
|
||||||
|
cancelDrag,
|
||||||
|
handleSortKeydown,
|
||||||
|
handleSortClick,
|
||||||
|
sortItemClass,
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1734,6 +1734,12 @@ const legacyExactEnglishMessages: Record<string, string> = {
|
|||||||
'启用提供商': 'Enable provider',
|
'启用提供商': 'Enable provider',
|
||||||
'停用提供商': 'Disable provider',
|
'停用提供商': 'Disable provider',
|
||||||
'扩展操作配置': 'Ops configuration',
|
'扩展操作配置': 'Ops configuration',
|
||||||
|
'卡片视图': 'Card view',
|
||||||
|
'切换到卡片视图': 'Switch to card view',
|
||||||
|
'切换到列表视图': 'Switch to list view',
|
||||||
|
'拖动调整当前页展示顺序,也可使用方向键移动': 'Drag to reorder the current page, or use the arrow keys',
|
||||||
|
'调整展示顺序': 'Adjust display order',
|
||||||
|
'展示顺序已更新': 'Display order updated',
|
||||||
'输入备注...': 'Enter notes...',
|
'输入备注...': 'Enter notes...',
|
||||||
'添加备注': 'Add notes',
|
'添加备注': 'Add notes',
|
||||||
'配额': 'Quota',
|
'配额': 'Quota',
|
||||||
|
|||||||
@@ -1,5 +1,10 @@
|
|||||||
<template>
|
<template>
|
||||||
<div class="space-y-4">
|
<div
|
||||||
|
ref="providerListRef"
|
||||||
|
class="space-y-4"
|
||||||
|
:class="{ 'select-none [&_*]:!cursor-grabbing': draggingProvider }"
|
||||||
|
@click.capture="handleSortClick"
|
||||||
|
>
|
||||||
<ProviderDeleteProgressCard
|
<ProviderDeleteProgressCard
|
||||||
:progress="providerDeleteProgress"
|
:progress="providerDeleteProgress"
|
||||||
:stage-label="providerDeleteStageLabel"
|
:stage-label="providerDeleteStageLabel"
|
||||||
@@ -25,6 +30,7 @@
|
|||||||
:model-filters="modelFilters"
|
:model-filters="modelFilters"
|
||||||
:has-active-filters="hasActiveFilters"
|
:has-active-filters="hasActiveFilters"
|
||||||
:loading="loading"
|
:loading="loading"
|
||||||
|
:card-view="cardView"
|
||||||
@update:search-query="searchQuery = $event"
|
@update:search-query="searchQuery = $event"
|
||||||
@update:filter-status="filterStatus = $event"
|
@update:filter-status="filterStatus = $event"
|
||||||
@update:filter-api-format="filterApiFormat = $event"
|
@update:filter-api-format="filterApiFormat = $event"
|
||||||
@@ -33,6 +39,7 @@
|
|||||||
@batch-process="openProviderBatchDialog"
|
@batch-process="openProviderBatchDialog"
|
||||||
@add-provider="openAddProviderDialog"
|
@add-provider="openAddProviderDialog"
|
||||||
@refresh="loadProviders"
|
@refresh="loadProviders"
|
||||||
|
@toggle-view="cardView = !cardView"
|
||||||
/>
|
/>
|
||||||
|
|
||||||
<!-- 加载状态 -->
|
<!-- 加载状态 -->
|
||||||
@@ -54,6 +61,50 @@
|
|||||||
/>
|
/>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
|
<div
|
||||||
|
v-else-if="cardView"
|
||||||
|
class="grid grid-cols-1 gap-4 p-4 sm:p-6 md:grid-cols-2 2xl:grid-cols-3"
|
||||||
|
>
|
||||||
|
<ProviderCard
|
||||||
|
v-for="provider in displayedProviders"
|
||||||
|
:key="provider.id"
|
||||||
|
:provider="provider"
|
||||||
|
:data-provider-sort-id="provider.id"
|
||||||
|
:class="sortItemClass(provider.id)"
|
||||||
|
:editing-description-id="editingDescriptionId"
|
||||||
|
:is-balance-loading="isBalanceLoading"
|
||||||
|
:get-provider-balance="getProviderBalance"
|
||||||
|
:get-provider-balance-breakdown="getProviderBalanceBreakdown"
|
||||||
|
:get-provider-balance-error="getProviderBalanceError"
|
||||||
|
:get-provider-checkin="getProviderCheckin"
|
||||||
|
:get-provider-cookie-expired="getProviderCookieExpired"
|
||||||
|
:get-provider-balance-extra="getProviderBalanceExtra"
|
||||||
|
:format-balance-display="formatBalanceDisplay"
|
||||||
|
:format-reset-countdown="formatResetCountdown"
|
||||||
|
:get-quota-used-color-class="getQuotaUsedColorClass"
|
||||||
|
@mousedown="handleMouseDown"
|
||||||
|
@row-click="handleRowClick"
|
||||||
|
@view-detail="openProviderDrawer"
|
||||||
|
@edit-provider="openEditProviderDialog"
|
||||||
|
@open-ops-config="openOpsConfigDialog"
|
||||||
|
@toggle-status="toggleProviderStatus"
|
||||||
|
@delete-provider="handleDeleteProvider"
|
||||||
|
@start-edit-description="startEditDescription"
|
||||||
|
@save-description="saveDescription"
|
||||||
|
@cancel-edit-description="cancelEditDescription"
|
||||||
|
>
|
||||||
|
<template #drag-handle>
|
||||||
|
<ProviderDragHandle
|
||||||
|
class="-ml-2 h-10 w-4"
|
||||||
|
:provider-name="provider.name"
|
||||||
|
:disabled="loading || displayedProviders.length < 2"
|
||||||
|
@pointerdown="startDrag(provider.id, $event)"
|
||||||
|
@keydown="handleSortKeydown(provider.id, $event)"
|
||||||
|
/>
|
||||||
|
</template>
|
||||||
|
</ProviderCard>
|
||||||
|
</div>
|
||||||
|
|
||||||
<!-- 桌面端表格 -->
|
<!-- 桌面端表格 -->
|
||||||
<div
|
<div
|
||||||
v-else
|
v-else
|
||||||
@@ -62,6 +113,9 @@
|
|||||||
<Table>
|
<Table>
|
||||||
<TableHeader>
|
<TableHeader>
|
||||||
<TableRow>
|
<TableRow>
|
||||||
|
<TableHead class="w-9 px-2">
|
||||||
|
<span class="sr-only">{{ legacyT('调整展示顺序') }}</span>
|
||||||
|
</TableHead>
|
||||||
<TableHead class="w-[18%] min-w-[140px]">
|
<TableHead class="w-[18%] min-w-[140px]">
|
||||||
{{ legacyT('提供商信息') }}
|
{{ legacyT('提供商信息') }}
|
||||||
</TableHead>
|
</TableHead>
|
||||||
@@ -131,6 +185,8 @@
|
|||||||
v-for="provider in displayedProviders"
|
v-for="provider in displayedProviders"
|
||||||
:key="provider.id"
|
:key="provider.id"
|
||||||
:provider="provider"
|
:provider="provider"
|
||||||
|
:data-provider-sort-id="provider.id"
|
||||||
|
:class="sortItemClass(provider.id)"
|
||||||
:editing-description-id="editingDescriptionId"
|
:editing-description-id="editingDescriptionId"
|
||||||
:is-balance-loading="isBalanceLoading"
|
:is-balance-loading="isBalanceLoading"
|
||||||
:get-provider-balance="getProviderBalance"
|
:get-provider-balance="getProviderBalance"
|
||||||
@@ -152,20 +208,31 @@
|
|||||||
@start-edit-description="startEditDescription"
|
@start-edit-description="startEditDescription"
|
||||||
@save-description="saveDescription"
|
@save-description="saveDescription"
|
||||||
@cancel-edit-description="cancelEditDescription"
|
@cancel-edit-description="cancelEditDescription"
|
||||||
/>
|
>
|
||||||
|
<template #drag-handle>
|
||||||
|
<ProviderDragHandle
|
||||||
|
:provider-name="provider.name"
|
||||||
|
:disabled="loading || displayedProviders.length < 2"
|
||||||
|
@pointerdown="startDrag(provider.id, $event)"
|
||||||
|
@keydown="handleSortKeydown(provider.id, $event)"
|
||||||
|
/>
|
||||||
|
</template>
|
||||||
|
</ProviderTableRow>
|
||||||
</TableBody>
|
</TableBody>
|
||||||
</Table>
|
</Table>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<!-- 移动端卡片列表 -->
|
<!-- 移动端卡片列表 -->
|
||||||
<div
|
<div
|
||||||
v-if="!loading && providers.length > 0"
|
v-if="!cardView && !loading && providers.length > 0"
|
||||||
class="xl:hidden divide-y divide-border/40"
|
class="xl:hidden divide-y divide-border/40"
|
||||||
>
|
>
|
||||||
<ProviderMobileCard
|
<ProviderMobileCard
|
||||||
v-for="provider in displayedProviders"
|
v-for="provider in displayedProviders"
|
||||||
:key="provider.id"
|
:key="provider.id"
|
||||||
:provider="provider"
|
:provider="provider"
|
||||||
|
:data-provider-sort-id="provider.id"
|
||||||
|
:class="sortItemClass(provider.id)"
|
||||||
:editing-description-id="editingDescriptionId"
|
:editing-description-id="editingDescriptionId"
|
||||||
:is-balance-loading="isBalanceLoading"
|
:is-balance-loading="isBalanceLoading"
|
||||||
:get-provider-balance="getProviderBalance"
|
:get-provider-balance="getProviderBalance"
|
||||||
@@ -182,7 +249,17 @@
|
|||||||
@start-edit-description="startEditDescription"
|
@start-edit-description="startEditDescription"
|
||||||
@save-description="saveDescription"
|
@save-description="saveDescription"
|
||||||
@cancel-edit-description="cancelEditDescription"
|
@cancel-edit-description="cancelEditDescription"
|
||||||
/>
|
>
|
||||||
|
<template #drag-handle>
|
||||||
|
<ProviderDragHandle
|
||||||
|
class="-ml-2 w-4"
|
||||||
|
:provider-name="provider.name"
|
||||||
|
:disabled="loading || displayedProviders.length < 2"
|
||||||
|
@pointerdown="startDrag(provider.id, $event)"
|
||||||
|
@keydown="handleSortKeydown(provider.id, $event)"
|
||||||
|
/>
|
||||||
|
</template>
|
||||||
|
</ProviderMobileCard>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<!-- 分页 -->
|
<!-- 分页 -->
|
||||||
@@ -196,8 +273,25 @@
|
|||||||
@update:page-size="pageSize = $event"
|
@update:page-size="pageSize = $event"
|
||||||
/>
|
/>
|
||||||
</Card>
|
</Card>
|
||||||
|
<span
|
||||||
|
class="sr-only"
|
||||||
|
role="status"
|
||||||
|
aria-live="polite"
|
||||||
|
aria-atomic="true"
|
||||||
|
>{{ announcement }}</span>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
|
<Teleport to="body">
|
||||||
|
<div
|
||||||
|
v-if="draggingProvider"
|
||||||
|
class="pointer-events-none fixed z-[100] w-[200px] truncate rounded-xl border border-primary/40 bg-card px-3 py-2 text-sm font-medium text-foreground shadow-lg"
|
||||||
|
:style="dragPreviewStyle"
|
||||||
|
aria-hidden="true"
|
||||||
|
>
|
||||||
|
{{ draggingProvider.name }}
|
||||||
|
</div>
|
||||||
|
</Teleport>
|
||||||
|
|
||||||
<!-- 对话框 -->
|
<!-- 对话框 -->
|
||||||
<ProviderFormDialog
|
<ProviderFormDialog
|
||||||
v-model="providerDialogOpen"
|
v-model="providerDialogOpen"
|
||||||
@@ -234,6 +328,7 @@
|
|||||||
|
|
||||||
<script setup lang="ts">
|
<script setup lang="ts">
|
||||||
import { ref, computed, watch, onMounted, onUnmounted, defineAsyncComponent } from 'vue'
|
import { ref, computed, watch, onMounted, onUnmounted, defineAsyncComponent } from 'vue'
|
||||||
|
import { useLocalStorage } from '@vueuse/core'
|
||||||
import Card from '@/components/ui/card.vue'
|
import Card from '@/components/ui/card.vue'
|
||||||
import Table from '@/components/ui/table.vue'
|
import Table from '@/components/ui/table.vue'
|
||||||
import TableHeader from '@/components/ui/table-header.vue'
|
import TableHeader from '@/components/ui/table-header.vue'
|
||||||
@@ -248,6 +343,8 @@ import ProviderBatchActionDialog from '@/features/providers/components/ProviderB
|
|||||||
import ProviderTableHeader from '@/features/providers/components/ProviderTableHeader.vue'
|
import ProviderTableHeader from '@/features/providers/components/ProviderTableHeader.vue'
|
||||||
import ProviderTableRow from '@/features/providers/components/ProviderTableRow.vue'
|
import ProviderTableRow from '@/features/providers/components/ProviderTableRow.vue'
|
||||||
import ProviderMobileCard from '@/features/providers/components/ProviderMobileCard.vue'
|
import ProviderMobileCard from '@/features/providers/components/ProviderMobileCard.vue'
|
||||||
|
import ProviderCard from '@/features/providers/components/ProviderCard.vue'
|
||||||
|
import ProviderDragHandle from '@/features/providers/components/ProviderDragHandle.vue'
|
||||||
import ProviderDeleteProgressCard from '@/features/providers/components/ProviderDeleteProgressCard.vue'
|
import ProviderDeleteProgressCard from '@/features/providers/components/ProviderDeleteProgressCard.vue'
|
||||||
import ProviderEmptyState from '@/features/providers/components/ProviderEmptyState.vue'
|
import ProviderEmptyState from '@/features/providers/components/ProviderEmptyState.vue'
|
||||||
import { useToast } from '@/composables/useToast'
|
import { useToast } from '@/composables/useToast'
|
||||||
@@ -255,6 +352,7 @@ import { useConfirm } from '@/composables/useConfirm'
|
|||||||
import { useRowClick } from '@/composables/useRowClick'
|
import { useRowClick } from '@/composables/useRowClick'
|
||||||
import { useProviderFilters } from '@/features/providers/composables/useProviderFilters'
|
import { useProviderFilters } from '@/features/providers/composables/useProviderFilters'
|
||||||
import { useProviderBalance } from '@/features/providers/composables/useProviderBalance'
|
import { useProviderBalance } from '@/features/providers/composables/useProviderBalance'
|
||||||
|
import { useProviderDisplayOrder } from '@/features/providers/composables/useProviderDisplayOrder'
|
||||||
import {
|
import {
|
||||||
getProvidersSummary,
|
getProvidersSummary,
|
||||||
getProvider,
|
getProvider,
|
||||||
@@ -294,6 +392,7 @@ function showLegacyError(err: unknown, fallback: string, title = '错误') {
|
|||||||
|
|
||||||
// 状态
|
// 状态
|
||||||
const loading = ref(false)
|
const loading = ref(false)
|
||||||
|
const cardView = useLocalStorage('aether-provider-card-view', false, { flush: 'sync' })
|
||||||
const providers = ref<ProviderWithEndpointsSummary[]>([])
|
const providers = ref<ProviderWithEndpointsSummary[]>([])
|
||||||
let providersRequestId = 0
|
let providersRequestId = 0
|
||||||
const providerDialogOpen = ref(false)
|
const providerDialogOpen = ref(false)
|
||||||
@@ -475,7 +574,20 @@ function sortProvidersByActiveAndPriority(items: ProviderWithEndpointsSummary[])
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
const displayedProviders = computed(() => sortProvidersByActiveAndPriority(providers.value))
|
const providerListRef = ref<HTMLElement | null>(null)
|
||||||
|
const {
|
||||||
|
orderedProviders: displayedProviders,
|
||||||
|
draggingProvider,
|
||||||
|
dragPreviewStyle,
|
||||||
|
announcement,
|
||||||
|
startDrag,
|
||||||
|
cancelDrag,
|
||||||
|
handleSortKeydown,
|
||||||
|
handleSortClick,
|
||||||
|
sortItemClass,
|
||||||
|
} = useProviderDisplayOrder(() => sortProvidersByActiveAndPriority(providers.value), providerListRef)
|
||||||
|
|
||||||
|
watch([loading, cardView, queryParams], cancelDrag)
|
||||||
|
|
||||||
function startEditDescription(_event: Event, provider: ProviderWithEndpointsSummary) {
|
function startEditDescription(_event: Event, provider: ProviderWithEndpointsSummary) {
|
||||||
editingDescriptionId.value = provider.id
|
editingDescriptionId.value = provider.id
|
||||||
|
|||||||
@@ -0,0 +1,539 @@
|
|||||||
|
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
|
||||||
|
import { createApp, nextTick, type App } from 'vue'
|
||||||
|
import type { ProviderWithEndpointsSummary } from '@/api/endpoints'
|
||||||
|
import { createI18n, setI18nLocale } from '@/i18n'
|
||||||
|
import ProviderManagement from '../ProviderManagement.vue'
|
||||||
|
|
||||||
|
const apiMocks = vi.hoisted(() => ({
|
||||||
|
getProvidersSummary: vi.fn(),
|
||||||
|
getGlobalModels: vi.fn(),
|
||||||
|
getProvider: vi.fn(),
|
||||||
|
updateProvider: vi.fn(),
|
||||||
|
}))
|
||||||
|
|
||||||
|
vi.mock('@/api/endpoints', async (importOriginal) => ({
|
||||||
|
...await importOriginal<typeof import('@/api/endpoints')>(),
|
||||||
|
...apiMocks,
|
||||||
|
}))
|
||||||
|
|
||||||
|
vi.mock('@/composables/useConfirm', () => ({
|
||||||
|
useConfirm: () => ({ confirmDanger: vi.fn().mockResolvedValue(false) }),
|
||||||
|
}))
|
||||||
|
|
||||||
|
vi.mock('@/composables/useToast', () => ({
|
||||||
|
useToast: () => ({ success: vi.fn(), error: vi.fn(), info: vi.fn() }),
|
||||||
|
}))
|
||||||
|
|
||||||
|
vi.mock('@/features/providers/composables/useProviderBalance', () => ({
|
||||||
|
useProviderBalance: () => ({
|
||||||
|
loadArchitectureSchemas: vi.fn(),
|
||||||
|
loadBalances: vi.fn(),
|
||||||
|
getProviderBalance: () => ({ available: 125, currency: 'USD' }),
|
||||||
|
getProviderBalanceBreakdown: () => null,
|
||||||
|
getProviderBalanceError: () => null,
|
||||||
|
isBalanceLoading: () => false,
|
||||||
|
getProviderCheckin: () => null,
|
||||||
|
getProviderCookieExpired: () => null,
|
||||||
|
formatBalanceDisplay: () => '$125.00',
|
||||||
|
formatResetCountdown: () => '',
|
||||||
|
getProviderBalanceExtra: () => [],
|
||||||
|
getQuotaUsedColorClass: () => '',
|
||||||
|
startTick: vi.fn(),
|
||||||
|
stopTick: vi.fn(),
|
||||||
|
}),
|
||||||
|
}))
|
||||||
|
|
||||||
|
vi.mock('@/features/providers/components', () => ({
|
||||||
|
ProviderFormDialog: { render: () => null },
|
||||||
|
ProviderAuthDialog: { render: () => null },
|
||||||
|
}))
|
||||||
|
|
||||||
|
vi.mock('@/features/providers/components/ProviderBatchActionDialog.vue', () => ({
|
||||||
|
default: { render: () => null },
|
||||||
|
}))
|
||||||
|
|
||||||
|
vi.mock('@/features/providers/components/ProviderDetailDrawer.vue', async () => {
|
||||||
|
const { h } = await import('vue')
|
||||||
|
return {
|
||||||
|
__esModule: true,
|
||||||
|
default: {
|
||||||
|
props: ['open', 'providerId'],
|
||||||
|
setup: (props: { open: boolean; providerId: string }) => () => props.open
|
||||||
|
? h('div', { 'data-provider-detail': props.providerId })
|
||||||
|
: null,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
function createProvider(overrides: Partial<ProviderWithEndpointsSummary> = {}): ProviderWithEndpointsSummary {
|
||||||
|
return {
|
||||||
|
id: 'provider-1',
|
||||||
|
name: 'Provider One',
|
||||||
|
description: 'Primary provider',
|
||||||
|
provider_type: 'custom',
|
||||||
|
provider_priority: 10,
|
||||||
|
keep_priority_on_conversion: false,
|
||||||
|
enable_format_conversion: true,
|
||||||
|
is_active: true,
|
||||||
|
total_endpoints: 2,
|
||||||
|
active_endpoints: 1,
|
||||||
|
total_keys: 3,
|
||||||
|
active_keys: 2,
|
||||||
|
total_models: 4,
|
||||||
|
active_models: 3,
|
||||||
|
global_model_ids: ['model-1'],
|
||||||
|
avg_health_score: 0.8,
|
||||||
|
unhealthy_endpoints: 0,
|
||||||
|
api_formats: ['openai:chat'],
|
||||||
|
endpoint_health_details: [{
|
||||||
|
api_format: 'openai:chat',
|
||||||
|
health_score: 0.8,
|
||||||
|
is_active: true,
|
||||||
|
total_keys: 3,
|
||||||
|
active_keys: 2,
|
||||||
|
}],
|
||||||
|
ops_configured: true,
|
||||||
|
created_at: '2026-09-07T00:00:00Z',
|
||||||
|
updated_at: '2026-09-07T00:00:00Z',
|
||||||
|
...overrides,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
let mountedApp: App | null = null
|
||||||
|
let mountedRoot: HTMLElement | null = null
|
||||||
|
|
||||||
|
async function settle() {
|
||||||
|
for (let index = 0; index < 8; index += 1) {
|
||||||
|
await Promise.resolve()
|
||||||
|
await nextTick()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
async function mountView() {
|
||||||
|
const root = document.createElement('div')
|
||||||
|
document.body.appendChild(root)
|
||||||
|
mountedRoot = root
|
||||||
|
mountedApp = createApp(ProviderManagement)
|
||||||
|
mountedApp.use(createI18n())
|
||||||
|
mountedApp.mount(root)
|
||||||
|
await settle()
|
||||||
|
return root
|
||||||
|
}
|
||||||
|
|
||||||
|
function unmountView() {
|
||||||
|
mountedApp?.unmount()
|
||||||
|
mountedRoot?.remove()
|
||||||
|
mountedApp = null
|
||||||
|
mountedRoot = null
|
||||||
|
}
|
||||||
|
|
||||||
|
function findButton(root: HTMLElement, title: string): HTMLButtonElement {
|
||||||
|
const button = root.querySelector<HTMLButtonElement>(`button[title="${title}"]`)
|
||||||
|
expect(button, `Missing button: ${title}`).not.toBeNull()
|
||||||
|
return button!
|
||||||
|
}
|
||||||
|
|
||||||
|
beforeEach(() => {
|
||||||
|
vi.clearAllMocks()
|
||||||
|
apiMocks.getProvidersSummary.mockResolvedValue({
|
||||||
|
items: [createProvider()],
|
||||||
|
total: 40,
|
||||||
|
})
|
||||||
|
apiMocks.getGlobalModels.mockResolvedValue({ models: [{ id: 'model-1', name: 'Model One' }] })
|
||||||
|
apiMocks.getProvider.mockResolvedValue(createProvider())
|
||||||
|
apiMocks.updateProvider.mockResolvedValue(createProvider({ is_active: false }))
|
||||||
|
})
|
||||||
|
|
||||||
|
const originalElementFromPoint = Object.getOwnPropertyDescriptor(document, 'elementFromPoint')
|
||||||
|
|
||||||
|
afterEach(() => {
|
||||||
|
unmountView()
|
||||||
|
if (originalElementFromPoint) {
|
||||||
|
Object.defineProperty(document, 'elementFromPoint', originalElementFromPoint)
|
||||||
|
} else {
|
||||||
|
Reflect.deleteProperty(document, 'elementFromPoint')
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
describe('ProviderManagement card view', () => {
|
||||||
|
it('places the view toggle immediately after refresh and switches layouts without reloading data', async () => {
|
||||||
|
const root = await mountView()
|
||||||
|
const toggle = findButton(root, '切换到卡片视图')
|
||||||
|
const filters = [...root.querySelectorAll('[role="combobox"]')].slice(0, 3)
|
||||||
|
|
||||||
|
expect(toggle.previousElementSibling).toBe(findButton(root, '刷新'))
|
||||||
|
expect(toggle.getAttribute('aria-pressed')).toBe('false')
|
||||||
|
expect(root.querySelector('table')).not.toBeNull()
|
||||||
|
expect(root.querySelector('dl')).toBeNull()
|
||||||
|
for (const filter of filters) {
|
||||||
|
expect(filter.closest('.xl\\:hidden')).not.toBeNull()
|
||||||
|
}
|
||||||
|
|
||||||
|
const requests = apiMocks.getProvidersSummary.mock.calls.length
|
||||||
|
toggle.click()
|
||||||
|
await settle()
|
||||||
|
|
||||||
|
expect(root.querySelector('table')).toBeNull()
|
||||||
|
expect(root.querySelectorAll('dl')).toHaveLength(1)
|
||||||
|
expect(root.textContent).toContain('$125.00')
|
||||||
|
expect(root.textContent).toContain('Primary provider')
|
||||||
|
expect(root.textContent).toContain('80%')
|
||||||
|
expect(toggle.getAttribute('aria-pressed')).toBe('true')
|
||||||
|
for (const filter of filters) {
|
||||||
|
expect(filter.closest('.xl\\:hidden')).toBeNull()
|
||||||
|
}
|
||||||
|
|
||||||
|
findButton(root, '切换到列表视图').click()
|
||||||
|
await settle()
|
||||||
|
expect(root.querySelector('table')).not.toBeNull()
|
||||||
|
expect(root.querySelector('dl')).toBeNull()
|
||||||
|
expect(apiMocks.getProvidersSummary).toHaveBeenCalledTimes(requests)
|
||||||
|
})
|
||||||
|
|
||||||
|
it('remembers the chosen layout across remounts', async () => {
|
||||||
|
let root = await mountView()
|
||||||
|
findButton(root, '切换到卡片视图').click()
|
||||||
|
await settle()
|
||||||
|
expect(localStorage.getItem('aether-provider-card-view')).toBe('true')
|
||||||
|
|
||||||
|
unmountView()
|
||||||
|
root = await mountView()
|
||||||
|
expect(root.querySelector('table')).toBeNull()
|
||||||
|
expect(findButton(root, '切换到列表视图').getAttribute('aria-pressed')).toBe('true')
|
||||||
|
|
||||||
|
findButton(root, '切换到列表视图').click()
|
||||||
|
await settle()
|
||||||
|
expect(localStorage.getItem('aether-provider-card-view')).toBe('false')
|
||||||
|
})
|
||||||
|
|
||||||
|
it.each([
|
||||||
|
{ name: 'cards', initial: false, selected: true, title: '切换到卡片视图' },
|
||||||
|
{ name: 'table', initial: true, selected: false, title: '切换到列表视图' },
|
||||||
|
])('immediately saves $name and restores it while reloading data', async ({ initial, selected, title }) => {
|
||||||
|
localStorage.setItem('aether-provider-card-view', String(initial))
|
||||||
|
let root = await mountView()
|
||||||
|
findButton(root, title).click()
|
||||||
|
expect(localStorage.getItem('aether-provider-card-view')).toBe(String(selected))
|
||||||
|
|
||||||
|
unmountView()
|
||||||
|
let finishLoading!: (value: { items: ProviderWithEndpointsSummary[]; total: number }) => void
|
||||||
|
apiMocks.getProvidersSummary.mockReturnValue(new Promise((resolve) => {
|
||||||
|
finishLoading = resolve
|
||||||
|
}))
|
||||||
|
root = await mountView()
|
||||||
|
const restoredTitle = selected ? '切换到列表视图' : '切换到卡片视图'
|
||||||
|
expect(findButton(root, restoredTitle).getAttribute('aria-pressed')).toBe(String(selected))
|
||||||
|
expect(findButton(root, '刷新').disabled).toBe(true)
|
||||||
|
|
||||||
|
finishLoading({ items: [createProvider()], total: 1 })
|
||||||
|
await settle()
|
||||||
|
expect(root.querySelector('table') === null).toBe(selected)
|
||||||
|
expect(root.querySelector('dl') !== null).toBe(selected)
|
||||||
|
expect(localStorage.getItem('aether-provider-card-view')).toBe(String(selected))
|
||||||
|
})
|
||||||
|
|
||||||
|
it('keeps the current search and page when switching views', async () => {
|
||||||
|
const root = await mountView()
|
||||||
|
const search = root.querySelector<HTMLInputElement>('#provider-search')!
|
||||||
|
search.value = 'Provider'
|
||||||
|
search.dispatchEvent(new Event('input', { bubbles: true }))
|
||||||
|
await vi.waitFor(() => {
|
||||||
|
expect(apiMocks.getProvidersSummary).toHaveBeenLastCalledWith(
|
||||||
|
expect.objectContaining({ search: 'Provider' }),
|
||||||
|
expect.any(Object),
|
||||||
|
)
|
||||||
|
})
|
||||||
|
|
||||||
|
const secondPage = [...root.querySelectorAll<HTMLButtonElement>('button')]
|
||||||
|
.find(button => button.textContent?.trim() === '2')!
|
||||||
|
secondPage.click()
|
||||||
|
await settle()
|
||||||
|
const requests = apiMocks.getProvidersSummary.mock.calls.length
|
||||||
|
|
||||||
|
findButton(root, '切换到卡片视图').click()
|
||||||
|
await settle()
|
||||||
|
expect(search.value).toBe('Provider')
|
||||||
|
expect(root.querySelector('[aria-current="page"]')?.textContent?.trim()).toBe('2')
|
||||||
|
expect(apiMocks.getProvidersSummary).toHaveBeenCalledTimes(requests)
|
||||||
|
expect(apiMocks.getProvidersSummary).toHaveBeenLastCalledWith(
|
||||||
|
expect.objectContaining({ search: 'Provider', page: 2 }),
|
||||||
|
expect.any(Object),
|
||||||
|
)
|
||||||
|
})
|
||||||
|
|
||||||
|
it('supports note editing, status actions, and details from cards', async () => {
|
||||||
|
localStorage.setItem('aether-provider-card-view', 'true')
|
||||||
|
const root = await mountView()
|
||||||
|
findButton(root, 'Primary provider').click()
|
||||||
|
await settle()
|
||||||
|
|
||||||
|
const input = root.querySelector<HTMLInputElement>('[data-desc-editor] input')!
|
||||||
|
expect(input.value).toBe('Primary provider')
|
||||||
|
input.value = 'Updated note'
|
||||||
|
input.dispatchEvent(new Event('input', { bubbles: true }))
|
||||||
|
input.dispatchEvent(new KeyboardEvent('keydown', { key: 'Enter', bubbles: true }))
|
||||||
|
await settle()
|
||||||
|
|
||||||
|
expect(apiMocks.updateProvider).toHaveBeenCalledWith('provider-1', { description: 'Updated note' })
|
||||||
|
expect(root.textContent).toContain('Updated note')
|
||||||
|
expect(root.querySelector('[data-provider-detail]')).toBeNull()
|
||||||
|
|
||||||
|
findButton(root, '停用提供商').click()
|
||||||
|
await settle()
|
||||||
|
expect(apiMocks.updateProvider).toHaveBeenCalledWith('provider-1', { is_active: false })
|
||||||
|
expect(root.querySelector('[data-provider-detail]')).toBeNull()
|
||||||
|
|
||||||
|
findButton(root, 'Provider One').click()
|
||||||
|
await vi.waitFor(() => {
|
||||||
|
expect(root.querySelector('[data-provider-detail="provider-1"]')).not.toBeNull()
|
||||||
|
})
|
||||||
|
})
|
||||||
|
|
||||||
|
it('uses account labels and handles providers without endpoints', async () => {
|
||||||
|
localStorage.setItem('aether-provider-card-view', 'true')
|
||||||
|
apiMocks.getProvidersSummary.mockResolvedValue({
|
||||||
|
items: [createProvider({ provider_type: 'codex', is_active: false, endpoint_health_details: [] })],
|
||||||
|
total: 1,
|
||||||
|
})
|
||||||
|
const root = await mountView()
|
||||||
|
|
||||||
|
expect(root.querySelector('dl')?.textContent).toContain('账号')
|
||||||
|
expect(root.textContent).toContain('暂无端点')
|
||||||
|
expect(findButton(root, '启用提供商')).not.toBeNull()
|
||||||
|
})
|
||||||
|
|
||||||
|
it('does not display cards during loading or with an empty result', async () => {
|
||||||
|
localStorage.setItem('aether-provider-card-view', 'true')
|
||||||
|
let resolveRequest!: (value: { items: ProviderWithEndpointsSummary[]; total: number }) => void
|
||||||
|
apiMocks.getProvidersSummary.mockReturnValue(new Promise((resolve) => {
|
||||||
|
resolveRequest = resolve
|
||||||
|
}))
|
||||||
|
const root = await mountView()
|
||||||
|
expect(root.querySelector('dl')).toBeNull()
|
||||||
|
expect(findButton(root, '刷新').disabled).toBe(true)
|
||||||
|
|
||||||
|
resolveRequest({ items: [], total: 0 })
|
||||||
|
await settle()
|
||||||
|
expect(root.querySelector('dl')).toBeNull()
|
||||||
|
expect(root.textContent).toContain('暂无提供商,点击右上角添加')
|
||||||
|
expect(findButton(root, '刷新').disabled).toBe(false)
|
||||||
|
})
|
||||||
|
|
||||||
|
it('translates the view switch labels', async () => {
|
||||||
|
const root = await mountView()
|
||||||
|
setI18nLocale('en-US')
|
||||||
|
await settle()
|
||||||
|
expect(findButton(root, 'Switch to card view').getAttribute('aria-label')).toBe('Card view')
|
||||||
|
|
||||||
|
findButton(root, 'Switch to card view').click()
|
||||||
|
await settle()
|
||||||
|
expect(findButton(root, 'Switch to list view').getAttribute('aria-pressed')).toBe('true')
|
||||||
|
})
|
||||||
|
})
|
||||||
|
|
||||||
|
function mockSortableProviders() {
|
||||||
|
const providers = [1, 2, 3, 4].map(index => createProvider({
|
||||||
|
id: `provider-${index}`,
|
||||||
|
name: `Provider ${index}`,
|
||||||
|
provider_priority: index * 10,
|
||||||
|
}))
|
||||||
|
apiMocks.getProvidersSummary.mockResolvedValue({ items: providers, total: providers.length })
|
||||||
|
return providers
|
||||||
|
}
|
||||||
|
|
||||||
|
function providerElements(root: HTMLElement): HTMLElement[] {
|
||||||
|
const container = root.querySelector('table') ?? root
|
||||||
|
return [...container.querySelectorAll<HTMLElement>('[data-provider-sort-id]')]
|
||||||
|
}
|
||||||
|
|
||||||
|
function providerOrder(root: HTMLElement): string[] {
|
||||||
|
return providerElements(root).map(element => element.dataset.providerSortId!)
|
||||||
|
}
|
||||||
|
|
||||||
|
function pointerEvent(type: string, clientX: number, clientY: number, options: { button?: number; pointerType?: string } = {}) {
|
||||||
|
const event = new MouseEvent(type, { bubbles: true, cancelable: true, clientX, clientY, button: options.button ?? 0 })
|
||||||
|
Object.defineProperties(event, {
|
||||||
|
pointerId: { value: 1 },
|
||||||
|
isPrimary: { value: true },
|
||||||
|
pointerType: { value: options.pointerType ?? 'mouse' },
|
||||||
|
})
|
||||||
|
return event
|
||||||
|
}
|
||||||
|
|
||||||
|
function startProviderDrag(root: HTMLElement, sourceId: string, targetId: string, pointerType = 'mouse') {
|
||||||
|
const elements = providerElements(root)
|
||||||
|
const handle = elements.find(element => element.dataset.providerSortId === sourceId)!
|
||||||
|
.querySelector<HTMLButtonElement>('[data-provider-drag-handle]')!
|
||||||
|
const target = elements.find(element => element.dataset.providerSortId === targetId)!
|
||||||
|
const hitTest = vi.fn((): Element | null => target)
|
||||||
|
Object.defineProperty(document, 'elementFromPoint', { configurable: true, value: hitTest })
|
||||||
|
handle.dispatchEvent(pointerEvent('pointerdown', 40, 100, { pointerType }))
|
||||||
|
window.dispatchEvent(pointerEvent('pointermove', 100, 200, { pointerType }))
|
||||||
|
return { handle, target, hitTest }
|
||||||
|
}
|
||||||
|
|
||||||
|
async function dropProvider(handle: HTMLButtonElement) {
|
||||||
|
window.dispatchEvent(pointerEvent('pointerup', 100, 200))
|
||||||
|
handle.click()
|
||||||
|
await settle()
|
||||||
|
}
|
||||||
|
|
||||||
|
describe('ProviderManagement shared display order', () => {
|
||||||
|
it('drags table rows, synchronizes both card layouts, and leaves scheduling priorities unchanged', async () => {
|
||||||
|
const providers = mockSortableProviders()
|
||||||
|
const root = await mountView()
|
||||||
|
const { handle, target } = startProviderDrag(root, 'provider-1', 'provider-3')
|
||||||
|
await settle()
|
||||||
|
expect(target.classList.contains('ring-2')).toBe(true)
|
||||||
|
expect(providerElements(root)[0]?.classList.contains('opacity-40')).toBe(true)
|
||||||
|
|
||||||
|
await dropProvider(handle)
|
||||||
|
const expected = ['provider-2', 'provider-3', 'provider-1', 'provider-4']
|
||||||
|
expect(providerOrder(root)).toEqual(expected)
|
||||||
|
const mobileOrder = [...root.querySelectorAll<HTMLElement>('[data-provider-sort-id]')]
|
||||||
|
.filter(element => !element.closest('table'))
|
||||||
|
.map(element => element.dataset.providerSortId)
|
||||||
|
expect(mobileOrder).toEqual(expected)
|
||||||
|
expect(root.querySelector('[data-provider-detail]')).toBeNull()
|
||||||
|
expect(apiMocks.updateProvider).not.toHaveBeenCalled()
|
||||||
|
expect(providers.map(provider => provider.provider_priority)).toEqual([10, 20, 30, 40])
|
||||||
|
|
||||||
|
findButton(root, '切换到卡片视图').click()
|
||||||
|
await settle()
|
||||||
|
expect(providerOrder(root)).toEqual(expected)
|
||||||
|
expect(apiMocks.getProvidersSummary).toHaveBeenCalledTimes(1)
|
||||||
|
})
|
||||||
|
|
||||||
|
it('supports touch dragging on card headers and restores the order after refresh and remount', async () => {
|
||||||
|
mockSortableProviders()
|
||||||
|
localStorage.setItem('aether-provider-card-view', 'true')
|
||||||
|
let root = await mountView()
|
||||||
|
const { handle } = startProviderDrag(root, 'provider-4', 'provider-1', 'touch')
|
||||||
|
await dropProvider(handle)
|
||||||
|
const expected = ['provider-4', 'provider-1', 'provider-2', 'provider-3']
|
||||||
|
expect(providerOrder(root)).toEqual(expected)
|
||||||
|
expect(JSON.parse(localStorage.getItem('aether-provider-display-order')!)).toEqual(expected)
|
||||||
|
|
||||||
|
findButton(root, '刷新').click()
|
||||||
|
await settle()
|
||||||
|
expect(providerOrder(root)).toEqual(expected)
|
||||||
|
findButton(root, '切换到列表视图').click()
|
||||||
|
await settle()
|
||||||
|
expect(providerOrder(root)).toEqual(expected)
|
||||||
|
|
||||||
|
unmountView()
|
||||||
|
root = await mountView()
|
||||||
|
expect(providerOrder(root)).toEqual(expected)
|
||||||
|
})
|
||||||
|
|
||||||
|
it.each(['escape', 'pointercancel', 'outside'] as const)('cancels a drag without saving when cancelled by %s', async (reason) => {
|
||||||
|
mockSortableProviders()
|
||||||
|
const root = await mountView()
|
||||||
|
const original = providerOrder(root)
|
||||||
|
const { handle, hitTest } = startProviderDrag(root, 'provider-1', 'provider-3')
|
||||||
|
if (reason === 'escape') {
|
||||||
|
window.dispatchEvent(new KeyboardEvent('keydown', { key: 'Escape', bubbles: true }))
|
||||||
|
} else if (reason === 'pointercancel') {
|
||||||
|
window.dispatchEvent(pointerEvent('pointercancel', 100, 200))
|
||||||
|
} else {
|
||||||
|
hitTest.mockReturnValue(null)
|
||||||
|
}
|
||||||
|
await dropProvider(handle)
|
||||||
|
|
||||||
|
expect(providerOrder(root)).toEqual(original)
|
||||||
|
expect(JSON.parse(localStorage.getItem('aether-provider-display-order')!)).toEqual([])
|
||||||
|
expect(root.querySelector('.opacity-40')).toBeNull()
|
||||||
|
expect(root.querySelector('[data-provider-detail]')).toBeNull()
|
||||||
|
})
|
||||||
|
|
||||||
|
it('does not reorder on a handle click or a secondary mouse button', async () => {
|
||||||
|
mockSortableProviders()
|
||||||
|
const root = await mountView()
|
||||||
|
const original = providerOrder(root)
|
||||||
|
const handle = providerElements(root)[0]!.querySelector<HTMLButtonElement>('[data-provider-drag-handle]')!
|
||||||
|
handle.dispatchEvent(pointerEvent('pointerdown', 40, 100))
|
||||||
|
window.dispatchEvent(pointerEvent('pointermove', 42, 101))
|
||||||
|
window.dispatchEvent(pointerEvent('pointerup', 42, 101))
|
||||||
|
handle.click()
|
||||||
|
handle.dispatchEvent(pointerEvent('pointerdown', 40, 100, { button: 2 }))
|
||||||
|
window.dispatchEvent(pointerEvent('pointermove', 100, 200))
|
||||||
|
window.dispatchEvent(pointerEvent('pointerup', 100, 200))
|
||||||
|
await settle()
|
||||||
|
|
||||||
|
expect(providerOrder(root)).toEqual(original)
|
||||||
|
expect(root.querySelector('[data-provider-detail]')).toBeNull()
|
||||||
|
})
|
||||||
|
|
||||||
|
it('moves only visible providers while preserving filtered-out positions', async () => {
|
||||||
|
const providers = mockSortableProviders()
|
||||||
|
const root = await mountView()
|
||||||
|
apiMocks.getProvidersSummary.mockResolvedValue({ items: [providers[0], providers[2]], total: 2 })
|
||||||
|
const search = root.querySelector<HTMLInputElement>('#provider-search')!
|
||||||
|
search.value = 'filtered'
|
||||||
|
search.dispatchEvent(new Event('input', { bubbles: true }))
|
||||||
|
await vi.waitFor(() => expect(providerOrder(root)).toEqual(['provider-1', 'provider-3']))
|
||||||
|
|
||||||
|
const { handle } = startProviderDrag(root, 'provider-1', 'provider-3')
|
||||||
|
await dropProvider(handle)
|
||||||
|
expect(providerOrder(root)).toEqual(['provider-3', 'provider-1'])
|
||||||
|
|
||||||
|
apiMocks.getProvidersSummary.mockResolvedValue({ items: providers, total: 4 })
|
||||||
|
findButton(root, '重置筛选').click()
|
||||||
|
await vi.waitFor(() => {
|
||||||
|
expect(providerOrder(root)).toEqual(['provider-3', 'provider-2', 'provider-1', 'provider-4'])
|
||||||
|
})
|
||||||
|
})
|
||||||
|
|
||||||
|
it('supports keyboard ordering and keeps focus on the moved handle', async () => {
|
||||||
|
mockSortableProviders()
|
||||||
|
const root = await mountView()
|
||||||
|
const handle = providerElements(root)[0]!.querySelector<HTMLButtonElement>('[data-provider-drag-handle]')!
|
||||||
|
handle.focus()
|
||||||
|
handle.dispatchEvent(new KeyboardEvent('keydown', { key: 'ArrowDown', bubbles: true, cancelable: true }))
|
||||||
|
await settle()
|
||||||
|
|
||||||
|
expect(providerOrder(root)).toEqual(['provider-2', 'provider-1', 'provider-3', 'provider-4'])
|
||||||
|
expect(document.activeElement).toBe(handle)
|
||||||
|
expect(root.querySelector('[role="status"]')?.textContent).toContain('展示顺序已更新')
|
||||||
|
expect(root.querySelector('[data-provider-detail]')).toBeNull()
|
||||||
|
|
||||||
|
handle.dispatchEvent(new KeyboardEvent('keydown', { key: 'ArrowLeft', bubbles: true, cancelable: true }))
|
||||||
|
await settle()
|
||||||
|
expect(providerOrder(root)).toEqual(['provider-1', 'provider-2', 'provider-3', 'provider-4'])
|
||||||
|
})
|
||||||
|
|
||||||
|
it('preserves saved ordering on other pages when reordering the current page', async () => {
|
||||||
|
const providers = mockSortableProviders()
|
||||||
|
apiMocks.getProvidersSummary.mockImplementation(async ({ page }: { page: number }) => ({
|
||||||
|
items: page === 1 ? providers.slice(0, 2) : providers.slice(2),
|
||||||
|
total: 40,
|
||||||
|
}))
|
||||||
|
const root = await mountView()
|
||||||
|
const firstDrag = startProviderDrag(root, 'provider-1', 'provider-2')
|
||||||
|
await dropProvider(firstDrag.handle)
|
||||||
|
|
||||||
|
const secondPage = [...root.querySelectorAll<HTMLButtonElement>('button')]
|
||||||
|
.find(button => button.textContent?.trim() === '2')!
|
||||||
|
secondPage.click()
|
||||||
|
await settle()
|
||||||
|
expect(providerOrder(root)).toEqual(['provider-3', 'provider-4'])
|
||||||
|
const secondDrag = startProviderDrag(root, 'provider-4', 'provider-3')
|
||||||
|
await dropProvider(secondDrag.handle)
|
||||||
|
expect(providerOrder(root)).toEqual(['provider-4', 'provider-3'])
|
||||||
|
|
||||||
|
const firstPage = [...root.querySelectorAll<HTMLButtonElement>('button')]
|
||||||
|
.find(button => button.textContent?.trim() === '1')!
|
||||||
|
firstPage.click()
|
||||||
|
await settle()
|
||||||
|
expect(providerOrder(root)).toEqual(['provider-2', 'provider-1'])
|
||||||
|
expect(JSON.parse(localStorage.getItem('aether-provider-display-order')!))
|
||||||
|
.toEqual(['provider-2', 'provider-1', 'provider-4', 'provider-3'])
|
||||||
|
})
|
||||||
|
|
||||||
|
it('ignores stale IDs and appends providers that are not in the saved order', async () => {
|
||||||
|
mockSortableProviders()
|
||||||
|
localStorage.setItem('aether-provider-display-order', JSON.stringify(['deleted-provider', 'provider-3', 'provider-1']))
|
||||||
|
const root = await mountView()
|
||||||
|
expect(providerOrder(root)).toEqual(['provider-3', 'provider-1', 'provider-2', 'provider-4'])
|
||||||
|
})
|
||||||
|
})
|
||||||
Reference in New Issue
Block a user