use alloc::sync::Arc;
use core::future::Future;
use log::{debug, error, warn};
use routers_network::{Entry, Network};
use routers_transition::{
Continuation, MatchError, Matcher,
costing::{CostingStrategies, DefaultEmissionCost, DefaultTransitionCost},
layer::generation::StandardGenerator,
primitives::PredicateCache,
weigh::AllCompute,
};
use tracing::{field, info_span};
use crate::event::MatchedDiff;
use crate::protocol::job::SolveJob;
use crate::protocol::result::SolveOutcome;
pub struct Engine<N: Network> {
network: Arc<N>,
runtime: N::Runtime,
costing: CostingStrategies<DefaultEmissionCost, DefaultTransitionCost, N::Entry>,
cache: Arc<PredicateCache<N>>,
search_distance: Option<f64>,
max_candidates: Option<usize>,
window_layers: Option<usize>,
}
impl<N: Network> Engine<N> {
pub fn new(network: Arc<N>, runtime: N::Runtime, search_distance: Option<f64>) -> Self {
Self {
network,
runtime,
costing: CostingStrategies::default(),
cache: Arc::new(PredicateCache::default()),
search_distance,
max_candidates: None,
window_layers: None,
}
}
#[must_use]
pub fn with_max_candidates(mut self, max_candidates: Option<usize>) -> Self {
self.max_candidates = max_candidates;
self
}
#[must_use]
pub fn with_window_layers(mut self, window_layers: Option<usize>) -> Self {
self.window_layers = window_layers;
self
}
pub fn solve(&self, job: &SolveJob<N::Entry>) -> SolveOutcome<N::Entry> {
let vehicle_id = job.identity.vehicle_id;
let mut generator = StandardGenerator::new(self.network.as_ref(), &self.costing.emission);
if let Some(distance) = self.search_distance {
generator = generator.with_search_distance(distance);
}
generator = generator.with_max_candidates(self.max_candidates);
let weigher = AllCompute::default().use_cache(self.cache.clone());
let matcher = Matcher::new(
self.network.as_ref(),
&self.costing,
generator,
weigher,
&self.runtime,
);
let span = info_span!(
"match_event",
outcome = field::Empty,
severity = field::Empty,
continuation = field::Empty,
converged = field::Empty,
window_cut = field::Empty,
emitted = field::Empty,
);
let _entered = span.enter();
let (mut trip, fresh, downgraded) = match job.context.clone() {
Continuation::Resume { trip, fresh } if !matcher.supports(&trip) => {
span.record("continuation", "downgrade");
warn!("{vehicle_id}: resume references a foreign shard; restarting");
let fresh = trip.origins().iter().copied().chain(fresh).collect();
(matcher.begin(), fresh, true)
}
Continuation::Resume { trip, fresh } => {
span.record("continuation", "resume");
(trip, fresh, false)
}
Continuation::Restart { fresh } => {
span.record("continuation", "restart");
(matcher.begin(), fresh, false)
}
};
info_span!("push", points = fresh.len()).in_scope(|| {
for origin in fresh {
match matcher.push(&mut trip, origin) {
Ok(_) => {}
Err(MatchError::Unanchored(err)) => {
info_span!("point_drop", reason = "unanchored")
.in_scope(|| debug!("{vehicle_id}: dropped off-network point ({err})"));
}
Err(err) => {
info_span!("point_drop", reason = "push_error")
.in_scope(|| error!("{vehicle_id}: could not push point: {err}"));
}
}
}
});
if trip.is_empty() {
span.record("outcome", "no_anchor");
span.record("severity", "nominal");
debug!("{vehicle_id}: no anchored layers to solve");
return SolveOutcome::Unanchored;
}
if let Err(err) = info_span!("solve").in_scope(|| matcher.solve(&mut trip)) {
let (outcome, severity) = classify(&err);
span.record("outcome", outcome);
span.record("severity", severity);
if severity == "nominal" {
debug!("{vehicle_id}: unable to solve trip: {err}");
} else {
error!("{vehicle_id}: unable to solve trip: {err}");
}
return terminal_outcome(err);
}
let origins = trip.origins().to_vec();
let solution = match info_span!("snapshot").in_scope(|| matcher.snapshot(&mut trip)) {
Ok(solution) => solution,
Err(err) => {
let (outcome, severity) = classify(&err);
span.record("outcome", outcome);
span.record("severity", severity);
return terminal_outcome(err);
}
};
let mut diff = info_span!("emit")
.in_scope(|| MatchedDiff::new(&solution, &origins, self.network.as_ref(), 0));
diff.downgraded = downgraded;
drop(solution);
span.record("emitted", diff.layers.len());
let converged_through = match matcher.convergence(&trip) {
Ok(Some(layer)) => {
span.record("converged", layer.index() as u64);
let timestamp = origins[layer.index()].timestamp;
trip.tail(trip.layers() - layer.index());
Some(timestamp)
}
Ok(None) => None,
Err(err) => {
error!("{vehicle_id}: convergence query failed: {err}");
None
}
};
let converged_through = match self.window_layers {
Some(window) if trip.layers() > window => {
let cut = trip.origins()[trip.layers() - window - 1].timestamp;
trip.tail(window);
span.record("window_cut", true);
Some(converged_through.map_or(cut, |c| c.max(cut)))
}
_ => converged_through,
};
span.record("outcome", "success");
span.record("severity", "ok");
SolveOutcome::Solved {
diff,
trip,
converged_through,
}
}
}
impl<N> Engine<N>
where
N: Network + 'static,
N::Runtime: 'static,
{
pub fn solve_blocking(
self: &Arc<Self>,
job: SolveJob<N::Entry>,
) -> impl Future<Output = SolveOutcome<N::Entry>> {
let engine = Arc::clone(self);
async move {
tokio::task::spawn_blocking(move || engine.solve(&job))
.await
.unwrap_or_else(|err| {
error!("solve task panicked: {err}");
SolveOutcome::Internal {
reason: "panic".to_owned(),
}
})
}
}
}
fn classify(err: &MatchError) -> (&'static str, &'static str) {
match err {
MatchError::Unanchored(_) => ("unanchored", "nominal"),
MatchError::Disconnected(_) => ("disconnected", "nominal"),
MatchError::TrellisError(_) | MatchError::SolveError(_) => ("internal", "fatal"),
}
}
fn terminal_outcome<E: Entry>(err: MatchError) -> SolveOutcome<E> {
match err {
MatchError::Unanchored(_) => SolveOutcome::Unanchored,
MatchError::Disconnected(_) => SolveOutcome::Disconnected,
err @ (MatchError::TrellisError(_) | MatchError::SolveError(_)) => SolveOutcome::Internal {
reason: err.to_string(),
},
}
}
#[cfg(test)]
mod tests {
use super::*;
use geo::{Point, point};
use routers_network::mock::{MockEntryId, MockNetwork, MockNetworkBuilder};
use routers_transition::Origin;
use crate::event::VehicleId;
use crate::protocol::ids::{GraphVersion, Lane, ObservationId, RegionId, SCHEMA_VERSION};
use crate::protocol::job::{JobIdentity, SolveJob};
fn bent_road() -> MockNetwork {
MockNetworkBuilder::new()
.node(1, point!(x: -118.15, y: 34.15))
.node(2, point!(x: -118.16, y: 34.15))
.node(3, point!(x: -118.17, y: 34.15))
.node(4, point!(x: -118.17, y: 34.14))
.node(5, point!(x: -118.18, y: 34.14))
.edge(1, 2)
.edge(2, 3)
.edge(3, 4)
.edge(4, 5)
.build()
}
fn bent_road_disjoint() -> MockNetwork {
MockNetworkBuilder::new()
.node(1001, point!(x: -118.15, y: 34.15))
.node(1002, point!(x: -118.16, y: 34.15))
.node(1003, point!(x: -118.17, y: 34.15))
.node(1004, point!(x: -118.17, y: 34.14))
.node(1005, point!(x: -118.18, y: 34.14))
.edge(1001, 1002)
.edge(1002, 1003)
.edge(1003, 1004)
.edge(1004, 1005)
.build()
}
fn observations() -> Vec<Origin> {
[
point!(x: -118.151, y: 34.1503),
point!(x: -118.155, y: 34.1503),
point!(x: -118.165, y: 34.1503),
point!(x: -118.170, y: 34.1490),
point!(x: -118.172, y: 34.1403),
point!(x: -118.179, y: 34.1403),
]
.into_iter()
.enumerate()
.map(|(index, point)| Origin::new(point, 1_775_000_000_000_000 + index as i64 * 5_000_000))
.collect()
}
fn engine(network: MockNetwork) -> Arc<Engine<MockNetwork>> {
Arc::new(Engine::new(Arc::new(network), (), None))
}
fn identity() -> JobIdentity {
JobIdentity {
schema: SCHEMA_VERSION,
vehicle_id: VehicleId(7),
observation: ObservationId {
partition: 0,
sequence: 1,
},
base: None,
graph: GraphVersion::new("v1").unwrap(),
region: RegionId::new("region").unwrap(),
}
}
fn job(context: Continuation<MockEntryId>) -> SolveJob<MockEntryId> {
SolveJob::new(identity(), Lane::DEFAULT, i64::MAX, context)
}
#[test]
fn restart_solves_every_layer() {
let engine = engine(bent_road());
let origins = observations();
let outcome = engine.solve(&job(Continuation::Restart {
fresh: origins.clone(),
}));
let SolveOutcome::Solved {
diff,
trip,
converged_through,
} = outcome
else {
panic!("a restart over an anchored trace must solve, got {outcome:?}");
};
assert!(!diff.downgraded, "a fresh restart is never a downgrade");
assert_eq!(
diff.layers.len(),
origins.len(),
"one emitted layer per observation"
);
match converged_through {
Some(timestamp) => {
assert!(
origins.iter().any(|o| o.timestamp == timestamp),
"the convergence stamp is one of the observations'"
);
assert_eq!(
trip.origins().first().map(|o| o.timestamp),
Some(timestamp),
"the cut trip resumes from the convergence layer"
);
}
None => assert!(
!trip.is_empty(),
"an unfused trip stays whole as the resume state"
),
}
}
#[test]
fn resume_extends_without_downgrade() {
let engine = engine(bent_road());
let SolveOutcome::Solved { trip, .. } = engine.solve(&job(Continuation::Restart {
fresh: observations(),
})) else {
panic!("the seed restart must solve");
};
let next = Origin::new(
point!(x: -118.1795, y: 34.1401),
1_775_000_000_000_000 + 6 * 5_000_000,
);
let outcome = engine.solve(&job(Continuation::Resume {
trip,
fresh: vec![next],
}));
let SolveOutcome::Solved { diff, .. } = outcome else {
panic!("resuming a supported trip must solve, got {outcome:?}");
};
assert!(
!diff.downgraded,
"a trip this engine supports resumes rather than downgrades"
);
}
#[test]
fn resume_of_foreign_trip_downgrades() {
let foreign = engine(bent_road_disjoint());
let SolveOutcome::Solved { trip, .. } = foreign.solve(&job(Continuation::Restart {
fresh: observations(),
})) else {
panic!("the foreign restart must solve");
};
let local = engine(bent_road());
let outcome = local.solve(&job(Continuation::Resume {
trip,
fresh: Vec::new(),
}));
let SolveOutcome::Solved { diff, .. } = outcome else {
panic!("a downgraded resume over an anchored trace must solve, got {outcome:?}");
};
assert!(
diff.downgraded,
"a foreign trip forces the downgrade flag on the emission"
);
}
#[test]
fn all_points_off_network_are_unanchored() {
let engine = engine(bent_road());
let off: Vec<Origin> = (0..4)
.map(|i| Origin::new(Point::new(0.0, 0.0), 1_775_000_000_000_000 + i * 5_000_000))
.collect();
let outcome = engine.solve(&job(Continuation::Restart { fresh: off }));
assert!(
matches!(outcome, SolveOutcome::Unanchored),
"off-network points solve to Unanchored, got {outcome:?}"
);
}
#[tokio::test]
async fn solve_blocking_matches_direct() {
let engine = engine(bent_road());
let outcome = engine
.solve_blocking(job(Continuation::Restart {
fresh: observations(),
}))
.await;
assert!(
matches!(outcome, SolveOutcome::Solved { .. }),
"the blocking solve mirrors the direct one, got {outcome:?}"
);
}
}