Bound local HTTP resource use

This commit is contained in:
Georg Bauer
2026-07-25 12:56:23 +02:00
parent 429cbc78de
commit 90400701fa
2 changed files with 114 additions and 12 deletions

View File

@@ -34,13 +34,16 @@ use std::fs::File;
use std::io::{Read, Write};
use std::net::{TcpListener, TcpStream};
use std::path::PathBuf;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex, RwLock};
use std::thread::{self, JoinHandle};
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
const MAX_HEADER_BYTES: usize = 64 * 1024;
const MAX_BODY_BYTES: usize = 64 * 1024 * 1024;
const MAX_CONNECTIONS: usize = 64;
const HTTP_IO_TIMEOUT: Duration = Duration::from_secs(10);
const HTTP_REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
const PREFILL_KEEPALIVE_INTERVAL: Duration = Duration::from_secs(5);
const TOOLS_PROMPT: &str = "## Tools\n\n\
You have access to a set of tools to help answer the user question. You can invoke tools by writing a \"<DSMLtool_calls>\" block like the following:\n\n\
@@ -71,6 +74,7 @@ struct State {
models_path: PathBuf,
cache_path: PathBuf,
sequence: AtomicU64,
connections: Arc<AtomicUsize>,
tool_memory: Mutex<HashMap<String, String>>,
metrics: Arc<Metrics>,
}
@@ -226,6 +230,7 @@ impl ServerHandle {
models_path,
cache_path,
sequence: AtomicU64::new(0),
connections: Arc::new(AtomicUsize::new(0)),
tool_memory: Mutex::new(HashMap::new()),
metrics: Arc::clone(&metrics),
});
@@ -257,10 +262,20 @@ fn serve(listener: TcpListener, state: Arc<State>, stop: Arc<AtomicBool>) {
while !stop.load(Ordering::Relaxed) {
match listener.accept() {
Ok((stream, _)) => {
let Some(slot) = ConnectionSlot::acquire(&state.connections) else {
continue;
};
let state = Arc::clone(&state);
let _ = thread::Builder::new()
.name("http-request".into())
.spawn(move || handle(stream, &state));
if let Err(error) =
thread::Builder::new()
.name("http-request".into())
.spawn(move || {
let _slot = slot;
handle(stream, &state);
})
{
eprintln!("DS4Server: endpoint request thread failed: {error}");
}
}
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(50));
@@ -273,10 +288,29 @@ fn serve(listener: TcpListener, state: Arc<State>, stop: Arc<AtomicBool>) {
}
}
struct ConnectionSlot(Arc<AtomicUsize>);
impl ConnectionSlot {
fn acquire(active: &Arc<AtomicUsize>) -> Option<Self> {
active
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |count| {
(count < MAX_CONNECTIONS).then_some(count + 1)
})
.ok()?;
Some(Self(Arc::clone(active)))
}
}
impl Drop for ConnectionSlot {
fn drop(&mut self) {
self.0.fetch_sub(1, Ordering::Relaxed);
}
}
fn handle(mut stream: TcpStream, state: &State) {
let _ = stream.set_read_timeout(Some(Duration::from_secs(30)));
let _ = stream.set_write_timeout(Some(HTTP_IO_TIMEOUT));
let started = Instant::now();
let request = match read_request(&mut stream) {
let request = match read_request(&mut stream, started + HTTP_REQUEST_TIMEOUT) {
Ok(request) => request,
Err(error) => {
state.metrics.http_started("", 0);
@@ -518,6 +552,41 @@ fn empty_arguments() -> String {
mod tests {
use super::*;
#[test]
fn connection_slots_enforce_the_limit_and_release_on_drop() {
let active = Arc::new(AtomicUsize::new(0));
let mut slots = (0..MAX_CONNECTIONS)
.map(|_| ConnectionSlot::acquire(&active).unwrap())
.collect::<Vec<_>>();
assert!(ConnectionSlot::acquire(&active).is_none());
slots.pop();
assert!(ConnectionSlot::acquire(&active).is_some());
}
#[test]
fn request_deadline_stops_slow_clients() {
let listener = match TcpListener::bind("127.0.0.1:0") {
Ok(listener) => listener,
Err(error) if error.kind() == std::io::ErrorKind::PermissionDenied => return,
Err(error) => panic!("could not start test server: {error}"),
};
let address = listener.local_addr().unwrap();
let client = thread::spawn(move || {
let mut stream = TcpStream::connect(address).unwrap();
stream.write_all(b"GET /v1/models HTTP/1.1\r\n").unwrap();
thread::sleep(Duration::from_millis(200));
});
let (mut stream, _) = listener.accept().unwrap();
let error = read_request(&mut stream, Instant::now() + Duration::from_millis(50))
.err()
.unwrap();
client.join().unwrap();
assert_eq!(error, "HTTP request timed out");
}
#[test]
fn tool_prompt_matches_ds4_server_text() {
assert_eq!(TOOLS_PROMPT.len(), 1_183);
@@ -612,6 +681,7 @@ mod tests {
models_path: PathBuf::new(),
cache_path: PathBuf::new(),
sequence: AtomicU64::new(0),
connections: Arc::new(AtomicUsize::new(0)),
tool_memory: Mutex::new(HashMap::new()),
metrics,
};
@@ -647,6 +717,7 @@ mod tests {
models_path: PathBuf::new(),
cache_path: PathBuf::new(),
sequence: AtomicU64::new(0),
connections: Arc::new(AtomicUsize::new(0)),
tool_memory: Mutex::new(HashMap::new()),
metrics,
};
@@ -673,6 +744,7 @@ mod tests {
models_path: PathBuf::new(),
cache_path: PathBuf::new(),
sequence: AtomicU64::new(0),
connections: Arc::new(AtomicUsize::new(0)),
tool_memory: Mutex::new(HashMap::new()),
metrics,
};