use std::collections::BTreeMap;
use std::ops::RangeInclusive;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use arc_swap::ArcSwap;
use tokio::sync::watch;
use tokio::task::JoinHandle;
use crate::core::supervision::task::{supervise, RestartPolicy, TaskHealth, TaskSpec};
use crate::error::{OlError, ERR_MODEL_RELAY_ENDPOINT_PORTS};
use super::wire_format::WireFormat;
use super::ModelRelayState;
pub const ENDPOINT_BLOCK: u16 = 32;
pub const ENDPOINT_TASK: &str = "model-relay-endpoint";
const RELEASE_DRAIN: Duration = Duration::from_secs(5);
pub fn endpoint_port_block(main: u16) -> RangeInclusive<u16> {
if main == 0 {
return empty_block();
}
match (main.checked_add(1), main.checked_add(ENDPOINT_BLOCK)) {
(Some(lo), Some(hi)) => lo..=hi,
_ => empty_block(),
}
}
#[allow(clippy::reversed_empty_ranges)]
fn empty_block() -> RangeInclusive<u16> {
1..=0
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct RelayPorts {
pub daemon: u16,
pub main: u16,
}
impl RelayPorts {
pub fn contains(&self, port: u16) -> bool {
port == self.daemon || port == self.main || endpoint_port_block(self.main).contains(&port)
}
pub fn in_block(&self, port: u16) -> bool {
endpoint_port_block(self.main).contains(&port)
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum OriginRefused {
#[error("scheme \"{0}\" is not http or https")]
Scheme(String),
#[error("the URL names no host")]
NoHost,
#[error("the URL names this client's own listener on port {0}")]
OwnListener(u16),
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Origin(reqwest::Url);
impl Origin {
pub fn as_url(&self) -> &reqwest::Url {
&self.0
}
}
impl std::fmt::Display for Origin {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.0.as_str())
}
}
pub fn normalize_origin(url: &reqwest::Url, own: &RelayPorts) -> Result<Origin, OriginRefused> {
let scheme = url.scheme();
if scheme != "http" && scheme != "https" {
return Err(OriginRefused::Scheme(scheme.to_string()));
}
let Some(host) = url.host() else {
return Err(OriginRefused::NoHost);
};
let loopback = match host {
url::Host::Ipv4(ip) => ip.is_loopback() || ip.is_unspecified(),
url::Host::Ipv6(ip) => ip.is_loopback() || ip.is_unspecified(),
url::Host::Domain(d) => d.eq_ignore_ascii_case("localhost"),
};
if let Some(port) = url.port_or_known_default() {
if loopback && own.contains(port) {
return Err(OriginRefused::OwnListener(port));
}
}
let mut origin = url.clone();
origin.set_path("/");
origin.set_query(None);
origin.set_fragment(None);
let _ = origin.set_username("");
let _ = origin.set_password(None);
Ok(Origin(origin))
}
#[derive(Clone, Debug)]
pub struct EndpointSpec {
pub key: String,
pub agent: &'static str,
pub family: Option<WireFormat>,
pub port: u16,
pub origin: Origin,
}
#[derive(Debug)]
pub struct RelayEndpoint {
key: String,
agent: &'static str,
family: Option<WireFormat>,
port: u16,
origin: ArcSwap<Origin>,
requests: AtomicU64,
last_request_unix: AtomicU64,
upstream_failures: AtomicU64,
contested: std::sync::atomic::AtomicBool,
}
impl RelayEndpoint {
pub fn new(spec: EndpointSpec) -> Self {
Self {
key: spec.key,
agent: spec.agent,
family: spec.family,
port: spec.port,
origin: ArcSwap::from_pointee(spec.origin),
requests: AtomicU64::new(0),
last_request_unix: AtomicU64::new(0),
upstream_failures: AtomicU64::new(0),
contested: std::sync::atomic::AtomicBool::new(false),
}
}
pub fn contested(&self) -> bool {
self.contested.load(Ordering::Relaxed)
}
pub fn set_contested(&self, contested: bool) {
self.contested.store(contested, Ordering::Relaxed);
}
pub fn key(&self) -> &str {
&self.key
}
pub fn agent(&self) -> &'static str {
self.agent
}
pub fn family(&self) -> Option<WireFormat> {
self.family
}
pub fn port(&self) -> u16 {
self.port
}
pub fn origin(&self) -> Arc<Origin> {
self.origin.load_full()
}
pub fn set_origin(&self, origin: Origin) -> bool {
if *self.origin.load_full() == origin {
return false;
}
self.origin.store(Arc::new(origin));
true
}
pub fn requests(&self) -> u64 {
self.requests.load(Ordering::Relaxed)
}
pub fn last_request_unix(&self) -> Option<u64> {
match self.last_request_unix.load(Ordering::Relaxed) {
0 => None,
t => Some(t),
}
}
pub fn upstream_failures(&self) -> u64 {
self.upstream_failures.load(Ordering::Relaxed)
}
pub(crate) fn note_request(&self) {
self.requests.fetch_add(1, Ordering::Relaxed);
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0);
self.last_request_unix.store(now, Ordering::Relaxed);
}
pub(crate) fn note_upstream_failure(&self) {
self.upstream_failures.fetch_add(1, Ordering::Relaxed);
}
pub fn to_json(&self) -> serde_json::Value {
serde_json::json!({
"key": self.key,
"agent": self.agent,
"port": self.port,
"origin": self.origin().to_string(),
"family": self.family.map(WireFormat::as_str),
"requests": self.requests(),
"last_request_unix": self.last_request_unix(),
"upstream_failures": self.upstream_failures(),
"contested": self.contested(),
})
}
}
pub type StateFactory = Arc<dyn Fn(Arc<RelayEndpoint>) -> ModelRelayState + Send + Sync>;
struct Running {
endpoint: Arc<RelayEndpoint>,
stop: watch::Sender<bool>,
join: JoinHandle<()>,
}
pub struct EndpointListeners {
make: StateFactory,
running: std::sync::Mutex<BTreeMap<String, Running>>,
}
impl EndpointListeners {
pub fn new(make: StateFactory) -> Self {
Self {
make,
running: std::sync::Mutex::new(BTreeMap::new()),
}
}
fn lock(&self) -> std::sync::MutexGuard<'_, BTreeMap<String, Running>> {
self.running.lock().unwrap_or_else(|e| e.into_inner())
}
pub async fn ensure(&self, spec: EndpointSpec) -> Result<Arc<RelayEndpoint>, OlError> {
let moved = {
let running = self.lock();
match running.get(&spec.key) {
Some(r) if r.endpoint.port() == spec.port && !r.join.is_finished() => {
r.endpoint.set_origin(spec.origin.clone());
return Ok(r.endpoint.clone());
}
Some(_) => true,
None => false,
}
};
if moved {
self.release(&spec.key).await;
}
let listener = super::bind_pinned(spec.port).await.map_err(|e| {
OlError::new(
ERR_MODEL_RELAY_ENDPOINT_PORTS,
format!(
"model relay endpoint {} could not bind loopback port {}: {}",
spec.key, spec.port, e.message
),
)
})?;
let endpoint = Arc::new(RelayEndpoint::new(spec));
let (stop, stop_rx) = watch::channel(false);
let pre_bound = Arc::new(tokio::sync::Mutex::new(Some(listener)));
let make = self.make.clone();
let served = endpoint.clone();
let serve_stop = stop_rx.clone();
let join = supervise(
TaskSpec::new(ENDPOINT_TASK, RestartPolicy::Always),
Arc::new(TaskHealth::new(ENDPOINT_TASK, RestartPolicy::Always)),
stop_rx,
move || {
let state = Arc::new(make(served.clone()).with_endpoint(served.clone()));
super::serve_attempt(pre_bound.clone(), state, serve_stop.clone())
},
);
tracing::info!(
key = endpoint.key(),
port = endpoint.port(),
origin = %endpoint.origin(),
"model relay endpoint serving (loopback only)"
);
self.lock().insert(
endpoint.key().to_string(),
Running {
endpoint: endpoint.clone(),
stop,
join,
},
);
Ok(endpoint)
}
pub async fn release(&self, key: &str) {
let Some(running) = self.lock().remove(key) else {
return;
};
stop_and_join(running).await;
}
pub async fn shutdown_all(&self) {
let all: Vec<Running> = std::mem::take(&mut *self.lock()).into_values().collect();
for running in all {
stop_and_join(running).await;
}
}
pub fn get(&self, key: &str) -> Option<Arc<RelayEndpoint>> {
self.lock().get(key).map(|r| r.endpoint.clone())
}
pub fn snapshot(&self) -> Vec<Arc<RelayEndpoint>> {
self.lock().values().map(|r| r.endpoint.clone()).collect()
}
}
impl Drop for EndpointListeners {
fn drop(&mut self) {
for running in self.lock().values() {
let _ = running.stop.send(true);
}
}
}
async fn stop_and_join(mut running: Running) {
let _ = running.stop.send(true);
if tokio::time::timeout(RELEASE_DRAIN, &mut running.join)
.await
.is_err()
{
tracing::warn!(
key = running.endpoint.key(),
port = running.endpoint.port(),
"model relay endpoint did not stop within the drain window — aborting it"
);
running.join.abort();
}
}
#[cfg(test)]
mod tests {
use super::*;
fn ports(main: u16) -> RelayPorts {
RelayPorts { daemon: 7443, main }
}
fn url(s: &str) -> reqwest::Url {
reqwest::Url::parse(s).expect("test url parses")
}
#[test]
fn endpoint_port_block_follows_the_main_port() {
assert_eq!(endpoint_port_block(7600), 7601..=7632);
assert_eq!(endpoint_port_block(17601), 17602..=17633);
let ports = ports(7600);
assert!(ports.in_block(7601) && ports.in_block(7632));
assert!(!ports.in_block(7600) && !ports.in_block(7633));
}
#[test]
fn a_main_port_at_the_top_of_the_range_yields_no_block() {
assert!(endpoint_port_block(u16::MAX - 10).is_empty());
assert!(endpoint_port_block(u16::MAX).is_empty());
assert!(!endpoint_port_block(u16::MAX - ENDPOINT_BLOCK).is_empty());
assert!(endpoint_port_block(0).is_empty());
assert!(!ports(0).in_block(1));
}
#[test]
fn origin_is_normalised_to_scheme_host_port() {
let own = ports(7600);
let cases = [
("https://gw.corp/anthropic/v1?x=1#f", "https://gw.corp/"),
("http://127.0.0.1:11434", "http://127.0.0.1:11434/"),
("http://localhost:1234/v1", "http://localhost:1234/"),
(
"https://generativelanguage.googleapis.com/v1beta",
"https://generativelanguage.googleapis.com/",
),
("https://user:pw@gw.corp:8443/x", "https://gw.corp:8443/"),
];
for (input, want) in cases {
let got = normalize_origin(&url(input), &own).expect(input);
assert_eq!(got.as_url().as_str(), want, "{input}");
}
}
#[test]
fn origin_that_is_our_own_listener_is_refused() {
let own = ports(7600);
for (input, port) in [
("http://127.0.0.1:7600", 7600),
("http://127.0.0.1:7601/v1", 7601),
("http://localhost:7632", 7632),
("http://[::1]:7443", 7443),
("http://0.0.0.0:7610", 7610),
] {
assert_eq!(
normalize_origin(&url(input), &own),
Err(OriginRefused::OwnListener(port)),
"{input}"
);
}
assert!(normalize_origin(&url("http://127.0.0.1:11434"), &own).is_ok());
assert!(normalize_origin(&url("http://127.0.0.1:7633"), &own).is_ok());
assert!(normalize_origin(&url("http://gw.corp:7601"), &own).is_ok());
}
#[test]
fn only_http_origins_are_forwardable() {
let own = ports(7600);
assert_eq!(
normalize_origin(&url("ftp://gw.corp/"), &own),
Err(OriginRefused::Scheme("ftp".into()))
);
assert_eq!(
normalize_origin(&url("file:///tmp/x"), &own),
Err(OriginRefused::Scheme("file".into()))
);
}
#[test]
fn an_origin_update_reports_whether_it_changed() {
let own = ports(7600);
let ep = RelayEndpoint::new(EndpointSpec {
key: "test:slot".into(),
agent: "cline",
family: None,
port: 7601,
origin: normalize_origin(&url("http://127.0.0.1:11434"), &own).expect("origin"),
});
let same = normalize_origin(&url("http://127.0.0.1:11434/api"), &own).expect("origin");
assert!(!ep.set_origin(same), "the same origin is not a change");
let moved = normalize_origin(&url("https://gw.corp"), &own).expect("origin");
assert!(ep.set_origin(moved.clone()));
assert_eq!(*ep.origin(), moved);
}
#[test]
fn counters_start_at_zero_and_record_the_request_time() {
let own = ports(7600);
let ep = RelayEndpoint::new(EndpointSpec {
key: "test:slot".into(),
agent: "cline",
family: Some(WireFormat::OllamaNative),
port: 7601,
origin: normalize_origin(&url("http://127.0.0.1:11434"), &own).expect("origin"),
});
assert_eq!(ep.requests(), 0);
assert_eq!(ep.last_request_unix(), None);
ep.note_request();
assert_eq!(ep.requests(), 1);
assert!(ep.last_request_unix().is_some());
let json = ep.to_json();
assert_eq!(json["family"], "ollama-native");
assert_eq!(json["origin"], "http://127.0.0.1:11434/");
}
}