diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index d0a07ef4..853382c0 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -55,6 +55,25 @@ jobs: tool: cargo-nextest - run: cargo nextest run --locked --workspace + loom: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v6.0.2 + - uses: dtolnay/rust-toolchain@stable + - uses: Swatinem/rust-cache@v2.9.1 + - run: cargo test --locked -p contextforge-gateway-rs-lib --features loom gateway::session_manager::concurrency + + miri: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v6.0.2 + - uses: dtolnay/rust-toolchain@nightly + with: + components: miri + - uses: Swatinem/rust-cache@v2.9.1 + - run: cargo miri setup + - run: cargo miri test --locked -p contextforge-gateway-rs-lib miri_checks + build: runs-on: ubuntu-latest steps: diff --git a/Cargo.lock b/Cargo.lock index eee54d5e..150356b4 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -673,6 +673,7 @@ dependencies = [ "hyper-util", "itertools", "jsonwebtoken", + "loom", "lru_time_cache", "mockito", "openid", @@ -1422,6 +1423,21 @@ dependencies = [ "slab", ] +[[package]] +name = "generator" +version = "0.8.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "52f04ae4152da20c76fe800fa48659201d5cf627c5149ca0b707b69d7eef6cf9" +dependencies = [ + "cc", + "cfg-if", + "libc", + "log", + "rustversion", + "windows-link", + "windows-result", +] + [[package]] name = "generic-array" version = "0.14.9" @@ -2120,6 +2136,19 @@ version = "0.4.29" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" +[[package]] +name = "loom" +version = "0.7.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "419e0dc8046cb947daa77eb95ae174acfbddb7673b4151f56d1eed8e93fbfaca" +dependencies = [ + "cfg-if", + "generator", + "scoped-tls", + "tracing", + "tracing-subscriber", +] + [[package]] name = "lru-slab" version = "0.1.2" @@ -3457,6 +3486,12 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "scoped-tls" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e1cf6437eb19a8f4a6cc0f7dca544973b0b78843adbfeb3683d1a94a0024a294" + [[package]] name = "scopeguard" version = "1.2.0" diff --git a/crates/contextforge-gateway-rs-lib/Cargo.toml b/crates/contextforge-gateway-rs-lib/Cargo.toml index 5c931826..52c88bf5 100644 --- a/crates/contextforge-gateway-rs-lib/Cargo.toml +++ b/crates/contextforge-gateway-rs-lib/Cargo.toml @@ -55,6 +55,7 @@ typed-builder.workspace = true [features] default = [] with_tools = [] +loom = [] [dev-dependencies] @@ -64,6 +65,7 @@ axum-test = "20.0.0" test-log = "0.2.20" axum-server = { version = "0.8.0", features = ["tls-rustls"] } futures.workspace = true +loom = "0.7.2" [lints] workspace = true diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs index 98fc8005..21709eae 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/mcp_gateway.rs @@ -29,7 +29,7 @@ use crate::{ SessionId, gateway::{ mcp_call_validator::InitializeCallValidator, - session_manager::SessionManager, + session_manager::{SessionManager, return_transport_entry}, session_store::{UserSession, UserSessionStore}, }, }; @@ -252,7 +252,8 @@ where let mut transports = self.transports.lock().await; for (name, svc) in backend_services { - transports.entry(BackendTransportKey::from((&name, session_id))).and_modify(|e| e.service = svc); + let key = BackendTransportKey::from((&name, session_id)); + return_transport_entry(&mut transports, &key, svc, |entry| &mut entry.service); } drop(transports); @@ -399,7 +400,8 @@ where let mut transports = self.transports.lock().await; for (name, svc) in backend_services { - transports.entry(BackendTransportKey::from((&name, session_id))).and_modify(|e| e.service = svc); + let key = BackendTransportKey::from((&name, session_id)); + return_transport_entry(&mut transports, &key, svc, |entry| &mut entry.service); } drop(transports); @@ -806,3 +808,7 @@ mod tests { assert_eq!(Some(pair), split_tool_name(&tool_name, &backend_names)); } } + +#[cfg(all(test, miri))] +#[path = "miri_checks/namespace_routing.rs"] +mod miri_checks; diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/miri_checks/config_serialization.rs b/crates/contextforge-gateway-rs-lib/src/gateway/miri_checks/config_serialization.rs new file mode 100644 index 00000000..d19c9d8b --- /dev/null +++ b/crates/contextforge-gateway-rs-lib/src/gateway/miri_checks/config_serialization.rs @@ -0,0 +1,60 @@ +use std::sync::Arc; + +use serde::Deserialize; + +use super::{SessionMapping, UserSession}; + +#[derive(Debug, Deserialize, PartialEq, Eq)] +struct OwnedUserSession { + name: String, + principal: String, + downstream_session_id: Arc, +} + +fn round_trip_session_mapping(entries: &[(&str, Option<&str>)]) -> Vec<(String, Option)> { + let mut mapping = SessionMapping::new(); + for (backend_name, upstream_session_id) in entries { + let upstream_session_id = upstream_session_id.map(Arc::::from); + mapping.push((*backend_name).to_owned(), upstream_session_id.as_ref()); + } + + let encoded = rmp_serde::encode::to_vec(&mapping).expect("session mapping should encode"); + let decoded: SessionMapping = rmp_serde::decode::from_slice(&encoded).expect("session mapping should decode"); + + entries + .iter() + .map(|(backend_name, _)| { + let upstream_session_id = + decoded.get(backend_name).and_then(super::SessionMap::session).map(|id| id.to_string()); + ((*backend_name).to_owned(), upstream_session_id) + }) + .collect() +} + +fn user_session_msgpack_round_trip(principal: &str, downstream_session_id: &str) -> bool { + let session = UserSession::new(principal.to_owned(), Arc::::from(downstream_session_id)); + let encoded = rmp_serde::encode::to_vec(&session).expect("user session should encode"); + let decoded: OwnedUserSession = rmp_serde::decode::from_slice(&encoded).expect("user session should decode"); + decoded + == OwnedUserSession { + name: "UserSession".to_owned(), + principal: principal.to_owned(), + downstream_session_id: Arc::::from(downstream_session_id), + } +} + +#[test] +fn session_mapping_msgpack_round_trip() { + let entries = [("backend-a", Some("upstream-a")), ("backend-b", None)]; + let round_tripped = round_trip_session_mapping(&entries); + + assert_eq!( + round_tripped, + vec![("backend-a".to_owned(), Some("upstream-a".to_owned())), ("backend-b".to_owned(), None)] + ); +} + +#[test] +fn user_session_key_msgpack_round_trip() { + assert!(user_session_msgpack_round_trip("principal-a", "downstream-session-a")); +} diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/miri_checks/namespace_routing.rs b/crates/contextforge-gateway-rs-lib/src/gateway/miri_checks/namespace_routing.rs new file mode 100644 index 00000000..e144ad95 --- /dev/null +++ b/crates/contextforge-gateway-rs-lib/src/gateway/miri_checks/namespace_routing.rs @@ -0,0 +1,41 @@ +use super::{split_resource_name, split_tool_name}; + +fn split_tool_name_owned(tool_name: &str, backend_names: &[&str]) -> Option<(String, String)> { + split_tool_name(&tool_name, backend_names).map(|pair| (pair.backend_name.to_owned(), pair.tool_name.to_owned())) +} + +fn split_resource_name_owned(resource_uri: &str, backend_names: &[&str]) -> Option<(String, String)> { + split_resource_name(&resource_uri, backend_names) + .map(|pair| (pair.backend_name.to_owned(), pair.resource_uri.to_owned())) +} + +#[test] +fn longest_backend_prefix_wins_for_tools() { + let backend_names = ["counter", "counter-one"]; + let parsed = split_tool_name_owned("counter-one-increment", &backend_names); + assert_eq!(parsed, Some(("counter-one".to_owned(), "increment".to_owned()))); +} + +#[test] +fn longest_backend_prefix_wins_for_resources() { + let backend_names = ["counter", "counter-one"]; + let parsed = split_resource_name_owned("counter-one-memo://insights", &backend_names); + assert_eq!(parsed, Some(("counter-one".to_owned(), "memo://insights".to_owned()))); +} + +#[test] +fn missing_separator_does_not_match() { + let backend_names = ["counter-one"]; + let parsed = split_tool_name_owned("counter-oneincrement", &backend_names); + assert_eq!(parsed, None); +} + +#[test] +fn hyphen_and_underscore_backend_names_are_distinct() { + let backend_names = ["counter-one", "counter_one"]; + let hyphenated = split_tool_name_owned("counter-one-increment", &backend_names); + let underscored = split_tool_name_owned("counter_one-increment", &backend_names); + + assert_eq!(hyphenated, Some(("counter-one".to_owned(), "increment".to_owned()))); + assert_eq!(underscored, Some(("counter_one".to_owned(), "increment".to_owned()))); +} diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/session_manager.rs b/crates/contextforge-gateway-rs-lib/src/gateway/session_manager.rs index 49e829ad..20f13f71 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/session_manager.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/session_manager.rs @@ -1,4 +1,8 @@ -use std::{collections::HashMap, sync::Arc}; +use std::{ + collections::HashMap, + hash::{BuildHasher, Hash}, + sync::Arc, +}; use contextforge_gateway_rs_apis::user_store::VirtualHost; use tokio::sync::Mutex; @@ -7,6 +11,39 @@ use tracing::{debug, info}; use super::mcp_gateway::{BackendTransportKey, BackendTransportService, ServiceHolder}; use crate::layers::session_id::SessionId; +pub(crate) fn borrow_transport_entry( + transports: &mut HashMap, + key: &K, + service_slot: impl FnOnce(&mut V) -> &mut Option, +) -> BorrowedTransport +where + K: Eq + Hash, + Hasher: BuildHasher, +{ + transports + .get_mut(key) + .map_or(BorrowedTransport::Missing, |entry| BorrowedTransport::Borrowed(service_slot(entry).take())) +} + +pub(crate) enum BorrowedTransport { + Borrowed(Option), + Missing, +} + +pub(crate) fn return_transport_entry( + transports: &mut HashMap, + key: &K, + running_service: Option, + service_slot: impl FnOnce(&mut V) -> &mut Option, +) where + K: Eq + Hash, + Hasher: BuildHasher, +{ + if let Some(entry) = transports.get_mut(key) { + *service_slot(entry) = running_service; + } +} + pub struct SessionManager<'a> { virtual_host: &'a VirtualHost, session_id: &'a SessionId, @@ -32,9 +69,11 @@ impl<'a> SessionManager<'a> { names .into_iter() .filter_map(|name| { - transports - .get_mut(&BackendTransportKey::from((&name, self.session_id))) - .map(|b| ServiceHolder::new(name, b.service.take())) + let key = BackendTransportKey::from((&name, self.session_id)); + match borrow_transport_entry(&mut transports, &key, |entry| &mut entry.service) { + BorrowedTransport::Borrowed(service) => Some(ServiceHolder::new(name, service)), + BorrowedTransport::Missing => None, + } }) .collect() } @@ -44,9 +83,8 @@ impl<'a> SessionManager<'a> { info!("Returning transports {:?} {backend_transports:?}", self.session_id); let mut transports = self.transports.lock().await; for svc_holder in backend_transports { - transports - .entry(BackendTransportKey::from((&svc_holder.name, self.session_id))) - .and_modify(|e| e.service = svc_holder.running_service); + let key = BackendTransportKey::from((&svc_holder.name, self.session_id)); + return_transport_entry(&mut transports, &key, svc_holder.running_service, |entry| &mut entry.service); } } @@ -61,3 +99,175 @@ impl<'a> SessionManager<'a> { } } } + +#[cfg(all(test, feature = "loom"))] +mod concurrency { + use std::collections::HashMap; + + use loom::sync::{Arc, Mutex}; + use loom::thread; + + use super::{BorrowedTransport, borrow_transport_entry, return_transport_entry}; + + #[derive(Clone, Debug, Hash, PartialEq, Eq)] + struct TransportKey { + backend: &'static str, + session: &'static str, + } + + #[derive(Debug)] + struct TestServiceHolder { + key: TransportKey, + running_service: Option, + } + + fn borrow_transports( + transports: &Arc>>>, + keys: &[TransportKey], + ) -> Vec { + let mut transports = transports.lock().expect("transport mutex should not be poisoned"); + keys.iter() + .filter_map(|key| match borrow_transport_entry(&mut transports, key, |service| service) { + BorrowedTransport::Borrowed(running_service) => { + Some(TestServiceHolder { key: key.clone(), running_service }) + }, + BorrowedTransport::Missing => None, + }) + .collect() + } + + fn return_transports( + transports: &Arc>>>, + holders: Vec, + ) { + let mut transports = transports.lock().expect("transport mutex should not be poisoned"); + for holder in holders { + return_transport_entry(&mut transports, &holder.key, holder.running_service, |service| service); + } + } + + fn cleanup_transports(transports: &Arc>>>, keys: &[TransportKey]) { + let mut transports = transports.lock().expect("transport mutex should not be poisoned"); + for key in keys { + transports.remove(key); + } + } + + #[test] + fn concurrent_borrowers_do_not_erase_returned_transport() { + loom::model(|| { + let key = TransportKey { backend: "backend-a", session: "session-one" }; + let transports = Arc::new(Mutex::new(HashMap::from([(key.clone(), Some(7))]))); + let keys = [key.clone()]; + + let first_transports = Arc::clone(&transports); + let first_keys = keys.clone(); + let first = thread::spawn(move || { + let borrowed = borrow_transports(&first_transports, &first_keys); + thread::yield_now(); + return_transports(&first_transports, borrowed); + }); + + let second_transports = Arc::clone(&transports); + let second = thread::spawn(move || { + let borrowed = borrow_transports(&second_transports, &keys); + thread::yield_now(); + return_transports(&second_transports, borrowed); + }); + + first.join().expect("first borrower should finish"); + second.join().expect("second borrower should finish"); + + let final_service = transports.lock().expect("transport mutex should not be poisoned").get(&key).copied(); + assert_eq!(final_service, Some(Some(7))); + }); + } + + #[test] + fn cleanup_does_not_resurrect_returned_transport() { + loom::model(|| { + let key = TransportKey { backend: "backend-a", session: "session-one" }; + let transports = Arc::new(Mutex::new(HashMap::from([(key.clone(), Some(7))]))); + let keys = [key.clone()]; + + let borrower_transports = Arc::clone(&transports); + let borrower_keys = keys.clone(); + let borrower = thread::spawn(move || { + let borrowed = borrow_transports(&borrower_transports, &borrower_keys); + thread::yield_now(); + return_transports(&borrower_transports, borrowed); + }); + + let cleanup_transports_ref = Arc::clone(&transports); + let cleanup = thread::spawn(move || { + thread::yield_now(); + cleanup_transports(&cleanup_transports_ref, &keys); + }); + + borrower.join().expect("borrower should finish"); + cleanup.join().expect("cleanup should finish"); + + let final_service = transports.lock().expect("transport mutex should not be poisoned").get(&key).copied(); + assert_ne!(final_service, Some(None)); + }); + } + + #[test] + fn different_sessions_do_not_interfere() { + loom::model(|| { + let first_key = TransportKey { backend: "backend-a", session: "session-one" }; + let second_key = TransportKey { backend: "backend-a", session: "session-two" }; + let transports = + Arc::new(Mutex::new(HashMap::from([(first_key.clone(), Some(7)), (second_key.clone(), Some(11))]))); + + let first_transports = Arc::clone(&transports); + let first_keys = [first_key.clone()]; + let first = thread::spawn(move || { + let borrowed = borrow_transports(&first_transports, &first_keys); + thread::yield_now(); + return_transports(&first_transports, borrowed); + }); + + let second_transports = Arc::clone(&transports); + let second_keys = [second_key.clone()]; + let second = thread::spawn(move || { + let borrowed = borrow_transports(&second_transports, &second_keys); + thread::yield_now(); + return_transports(&second_transports, borrowed); + }); + + first.join().expect("first session should finish"); + second.join().expect("second session should finish"); + + let transports = transports.lock().expect("transport mutex should not be poisoned"); + assert_eq!(transports.get(&first_key).copied(), Some(Some(7))); + assert_eq!(transports.get(&second_key).copied(), Some(Some(11))); + }); + } + + #[test] + fn multiple_backends_preserve_unborrowed_services() { + loom::model(|| { + let borrowed_key = TransportKey { backend: "backend-a", session: "session-one" }; + let unborrowed_key = TransportKey { backend: "backend-b", session: "session-one" }; + let transports = Arc::new(Mutex::new(HashMap::from([ + (borrowed_key.clone(), Some(7)), + (unborrowed_key.clone(), Some(11)), + ]))); + let keys = [borrowed_key.clone()]; + + let borrower_transports = Arc::clone(&transports); + let borrower = thread::spawn(move || { + let borrowed = borrow_transports(&borrower_transports, &keys); + thread::yield_now(); + return_transports(&borrower_transports, borrowed); + }); + + borrower.join().expect("borrower should finish"); + + let transports = transports.lock().expect("transport mutex should not be poisoned"); + assert_eq!(transports.get(&borrowed_key).copied(), Some(Some(7))); + assert_eq!(transports.get(&unborrowed_key).copied(), Some(Some(11))); + }); + } +} diff --git a/crates/contextforge-gateway-rs-lib/src/gateway/session_store/mod.rs b/crates/contextforge-gateway-rs-lib/src/gateway/session_store/mod.rs index 4b6b94f6..9983c97f 100644 --- a/crates/contextforge-gateway-rs-lib/src/gateway/session_store/mod.rs +++ b/crates/contextforge-gateway-rs-lib/src/gateway/session_store/mod.rs @@ -78,3 +78,7 @@ pub trait UserSessionStore: Send + Sync { session_mapping: &'a SessionMapping, ) -> Result<(), SessionStoreError>; } + +#[cfg(all(test, miri))] +#[path = "../miri_checks/config_serialization.rs"] +mod miri_checks;