#![allow(dead_code)]
use crate::routing::connection_cache::ConnectionCache;
use crate::routing::endpoint_cooldown::EndpointCooldownTracker;
use crate::routing::key_range_cache::{KeyRangeCache, RangeMode};
use crate::routing::server_connection::ServerConnection;
use std::collections::HashMap;
use std::fmt::{self, Debug, Formatter};
use std::sync::{Arc, RwLock};
#[derive(Clone, Debug, Default)]
pub(crate) struct RoutingContext<'a> {
pub transaction_id: Option<&'a [u8]>,
pub routing_key: Option<&'a [u8]>,
pub prefer_leader: bool,
pub use_transaction_affinity: bool,
}
#[derive(Debug, Default)]
struct AffinityTracker {
entries: HashMap<Vec<u8>, Arc<str>>,
}
impl AffinityTracker {
fn new() -> Self {
Self {
entries: HashMap::new(),
}
}
fn get(&self, transaction_id: &[u8]) -> Option<Arc<str>> {
if transaction_id.is_empty() {
return None;
}
self.entries.get(transaction_id).cloned()
}
fn insert(&mut self, transaction_id: &[u8], address: &str) {
if transaction_id.is_empty() {
return;
}
if let Some(entry) = self.entries.get_mut(transaction_id) {
if entry.as_ref() != address {
*entry = Arc::from(address);
}
return;
}
self.entries
.insert(transaction_id.to_vec(), Arc::from(address));
}
fn remove(&mut self, transaction_id: &[u8]) -> Option<Arc<str>> {
if transaction_id.is_empty() {
return None;
}
self.entries.remove(transaction_id)
}
fn len(&self) -> usize {
self.entries.len()
}
}
#[derive(Clone)]
pub(crate) struct LocationRouter {
key_range_cache: Arc<KeyRangeCache>,
connection_cache: Arc<ConnectionCache>,
cooldown_tracker: Arc<EndpointCooldownTracker>,
affinity_tracker: Arc<RwLock<AffinityTracker>>,
}
impl Debug for LocationRouter {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
f.debug_struct("LocationRouter")
.field("connection_cache", &self.connection_cache)
.field("cooldown_tracker", &self.cooldown_tracker)
.finish_non_exhaustive()
}
}
impl LocationRouter {
pub(crate) fn new(
key_range_cache: Arc<KeyRangeCache>,
connection_cache: Arc<ConnectionCache>,
cooldown_tracker: Arc<EndpointCooldownTracker>,
) -> Self {
Self {
key_range_cache,
connection_cache,
cooldown_tracker,
affinity_tracker: Arc::new(RwLock::new(AffinityTracker::new())),
}
}
pub(crate) fn key_range_cache(&self) -> &Arc<KeyRangeCache> {
&self.key_range_cache
}
pub(crate) fn connection_cache(&self) -> &Arc<ConnectionCache> {
&self.connection_cache
}
pub(crate) fn cooldown_tracker(&self) -> &Arc<EndpointCooldownTracker> {
&self.cooldown_tracker
}
pub(crate) fn resolve_connection(&self, context: &RoutingContext<'_>) -> ServerConnection {
if context.use_transaction_affinity
&& let Some(transaction_id) = context.transaction_id
&& let Some(affinity_address) = self.get_transaction_affinity(transaction_id)
&& !self.cooldown_tracker.is_cooling_down(&affinity_address)
&& let Some(connection) = self.connection_cache.get_if_present(&affinity_address)
{
return connection;
}
if let Some(routing_key) = context.routing_key
&& let Some(range) =
self.key_range_cache
.find_range(routing_key, &[], RangeMode::CoveringSplit)
&& let Some(tablet) = self
.key_range_cache
.select_tablet(&range, context.prefer_leader)
&& !self
.cooldown_tracker
.is_cooling_down(&tablet.server_address)
&& let Some(connection) = self.connection_cache.get_if_present(&tablet.server_address)
{
if context.use_transaction_affinity
&& let Some(transaction_id) = context.transaction_id
{
self.record_transaction_affinity(transaction_id, &tablet.server_address);
}
return connection;
}
self.connection_cache.default_connection().clone()
}
pub(crate) fn record_transaction_affinity(&self, transaction_id: &[u8], address: &str) {
if transaction_id.is_empty() {
return;
}
let mut tracker = self
.affinity_tracker
.write()
.expect("affinity tracker lock poisoned");
tracker.insert(transaction_id, address);
}
pub(crate) fn get_transaction_affinity(&self, transaction_id: &[u8]) -> Option<Arc<str>> {
if transaction_id.is_empty() {
return None;
}
let tracker = self
.affinity_tracker
.read()
.expect("affinity tracker lock poisoned");
tracker.get(transaction_id)
}
pub(crate) fn clear_transaction_affinity(&self, transaction_id: &[u8]) {
if transaction_id.is_empty() {
return;
}
let mut tracker = self
.affinity_tracker
.write()
.expect("affinity tracker lock poisoned");
let _ = tracker.remove(transaction_id);
}
pub(crate) fn affinity_count(&self) -> usize {
let tracker = self
.affinity_tracker
.read()
.expect("affinity tracker lock poisoned");
tracker.len()
}
pub(crate) fn record_failure(&self, address: &str) {
self.cooldown_tracker.record_failure(address);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::client::Channel;
use crate::generated::gapic_dataplane::stub::Spanner as SpannerStub;
use crate::model::{CacheUpdate, Group, Range, Tablet};
use crate::routing::server_connection::ServerConnection;
use gaxi::options::ClientConfig;
#[test]
fn location_router_implements_send_sync_debug_clone() {
static_assertions::assert_impl_all!(LocationRouter: Send, Sync, Debug, Clone);
}
#[derive(Debug)]
struct DummyStub;
impl SpannerStub for DummyStub {}
fn create_test_connection(address: &str) -> ServerConnection {
let channel = Channel::new_for_test(DummyStub);
ServerConnection::new(address.to_string(), channel)
}
fn make_test_router() -> LocationRouter {
let default_connection = create_test_connection("spanner.googleapis.com:443");
let connection_cache = Arc::new(ConnectionCache::new(default_connection));
let key_range_cache = Arc::new(KeyRangeCache::new());
let cooldown_tracker = Arc::new(EndpointCooldownTracker::new());
LocationRouter::new(key_range_cache, connection_cache, cooldown_tracker)
}
async fn populate_test_routing_table(
router: &LocationRouter,
address: &str,
start_key: Vec<u8>,
limit_key: Vec<u8>,
) {
let group = Group::new().set_group_uid(100u64).set_tablets(vec![
Tablet::default()
.set_tablet_uid(10u64)
.set_server_address(address),
]);
let range = Range::new()
.set_group_uid(100u64)
.set_start_key(start_key)
.set_limit_key(limit_key);
let update = CacheUpdate::new()
.set_database_id(1u64)
.set_group(vec![group])
.set_range(vec![range]);
router.key_range_cache().add_ranges(&update);
let _ = router
.connection_cache()
.get(address, &ClientConfig::default())
.await
.expect("should initialize connection");
}
#[test]
fn location_router_new_and_accessors() {
let router = make_test_router();
assert!(router.key_range_cache().is_empty());
assert_eq!(router.connection_cache().len(), 1);
assert_eq!(router.affinity_count(), 0);
}
#[test]
fn location_router_resolve_empty_cache_returns_default_connection() {
let router = make_test_router();
let context = RoutingContext::default();
let connection = router.resolve_connection(&context);
assert_eq!(connection.address(), "spanner.googleapis.com:443");
}
#[tokio::test]
async fn location_router_resolve_routing_key_returns_cached_connection() {
let router = make_test_router();
populate_test_routing_table(&router, "10.0.0.1:15000", vec![0x01], vec![0x09]).await;
let key = vec![0x05];
let context = RoutingContext {
transaction_id: None,
routing_key: Some(&key),
prefer_leader: false,
use_transaction_affinity: false,
};
let connection = router.resolve_connection(&context);
assert_eq!(connection.address(), "10.0.0.1:15000");
}
#[tokio::test]
async fn location_router_resolve_records_and_uses_transaction_affinity() {
let router = make_test_router();
populate_test_routing_table(&router, "10.0.0.1:15000", vec![0x01], vec![0x09]).await;
let transaction_id = b"tx1";
let key = vec![0x05];
let context = RoutingContext {
transaction_id: Some(transaction_id),
routing_key: Some(&key),
prefer_leader: false,
use_transaction_affinity: true,
};
let connection_first = router.resolve_connection(&context);
assert_eq!(connection_first.address(), "10.0.0.1:15000");
assert_eq!(
router.get_transaction_affinity(transaction_id).as_deref(),
Some("10.0.0.1:15000")
);
router.key_range_cache().clear();
let context_affinity = RoutingContext {
transaction_id: Some(transaction_id),
routing_key: None,
prefer_leader: false,
use_transaction_affinity: true,
};
let connection_second = router.resolve_connection(&context_affinity);
assert_eq!(connection_second.address(), "10.0.0.1:15000");
}
#[tokio::test]
async fn location_router_read_only_transaction_bypasses_affinity() {
let router = make_test_router();
populate_test_routing_table(&router, "10.0.0.1:15000", vec![0x01], vec![0x09]).await;
let transaction_id = b"tx_readonly";
let key = vec![0x05];
let context = RoutingContext {
transaction_id: Some(transaction_id),
routing_key: Some(&key),
prefer_leader: false,
use_transaction_affinity: false,
};
let connection = router.resolve_connection(&context);
assert_eq!(connection.address(), "10.0.0.1:15000");
assert_eq!(router.get_transaction_affinity(transaction_id), None);
}
#[tokio::test]
async fn location_router_skips_cooldown_endpoint() {
let router = make_test_router();
populate_test_routing_table(&router, "10.0.0.1:15000", vec![0x01], vec![0x09]).await;
router.record_failure("10.0.0.1:15000");
let key = vec![0x05];
let context = RoutingContext {
transaction_id: None,
routing_key: Some(&key),
prefer_leader: false,
use_transaction_affinity: false,
};
let connection = router.resolve_connection(&context);
assert_eq!(connection.address(), "spanner.googleapis.com:443");
}
#[tokio::test]
async fn location_router_skips_cooldown_affinity_endpoint() {
let router = make_test_router();
populate_test_routing_table(&router, "10.0.0.1:15000", vec![0x01], vec![0x09]).await;
router.record_transaction_affinity(b"tx1", "10.0.0.1:15000");
router.record_failure("10.0.0.1:15000");
let context = RoutingContext {
transaction_id: Some(b"tx1"),
routing_key: None,
prefer_leader: false,
use_transaction_affinity: true,
};
let connection = router.resolve_connection(&context);
assert_eq!(connection.address(), "spanner.googleapis.com:443");
}
#[test]
fn location_router_clear_transaction_affinity() {
let router = make_test_router();
router.record_transaction_affinity(b"tx1", "10.0.0.1:15000");
assert_eq!(router.affinity_count(), 1);
router.clear_transaction_affinity(b"tx1");
assert_eq!(router.affinity_count(), 0);
assert_eq!(router.get_transaction_affinity(b"tx1"), None);
}
#[test]
fn location_router_ignores_empty_transaction_id() {
let router = make_test_router();
router.record_transaction_affinity(&[], "10.0.0.1:15000");
assert_eq!(router.affinity_count(), 0);
assert_eq!(router.get_transaction_affinity(&[]), None);
router.clear_transaction_affinity(&[]);
assert_eq!(router.affinity_count(), 0);
}
#[test]
fn location_router_affinity_explicit_cleanup() {
let router = make_test_router();
router.record_transaction_affinity(b"tx1", "10.0.0.1:15000");
router.record_transaction_affinity(b"tx2", "10.0.0.2:15000");
router.record_transaction_affinity(b"tx3", "10.0.0.3:15000");
assert_eq!(router.affinity_count(), 3);
assert_eq!(
router.get_transaction_affinity(b"tx1").as_deref(),
Some("10.0.0.1:15000")
);
assert_eq!(
router.get_transaction_affinity(b"tx2").as_deref(),
Some("10.0.0.2:15000")
);
router.clear_transaction_affinity(b"tx1");
assert_eq!(router.affinity_count(), 2);
assert_eq!(router.get_transaction_affinity(b"tx1"), None);
}
#[tokio::test]
async fn location_router_prefer_leader_routing() {
let router = make_test_router();
let tablet_leader = Tablet::default()
.set_tablet_uid(10u64)
.set_server_address("10.0.0.1:15000");
let tablet_replica = Tablet::default()
.set_tablet_uid(11u64)
.set_server_address("10.0.0.2:15000");
let group = Group::new()
.set_group_uid(100u64)
.set_leader_index(0)
.set_tablets(vec![tablet_leader, tablet_replica]);
let range = Range::new()
.set_group_uid(100u64)
.set_start_key(vec![0x01])
.set_limit_key(vec![0x09]);
let update = CacheUpdate::new()
.set_database_id(1u64)
.set_group(vec![group])
.set_range(vec![range]);
router.key_range_cache().add_ranges(&update);
let _ = router
.connection_cache()
.get("10.0.0.1:15000", &ClientConfig::default())
.await
.expect("should initialize connection");
let key = vec![0x05];
let context = RoutingContext {
transaction_id: None,
routing_key: Some(&key),
prefer_leader: true,
use_transaction_affinity: false,
};
let connection = router.resolve_connection(&context);
assert_eq!(connection.address(), "10.0.0.1:15000");
}
#[test]
fn location_router_debug_formatting() {
let router = make_test_router();
let debug_str = format!("{:?}", router);
assert!(debug_str.contains("LocationRouter"));
assert!(debug_str.contains("connection_cache"));
assert!(debug_str.contains("cooldown_tracker"));
}
#[test]
fn location_router_record_cooldown_helper() {
let router = make_test_router();
router.record_failure("10.0.0.1:15000");
assert!(router.cooldown_tracker().is_cooling_down("10.0.0.1:15000"));
}
}