Skip to main content

canwu_routing/
lib.rs

1//! Deterministic routing over a versioned, time-dependent transfer network.
2//!
3//! This crate is deliberately pure: it does not read simulation state, mutate
4//! capacity, draw randomness, schedule work, or resolve information recipients.
5
6#![allow(clippy::missing_errors_doc)]
7
8use canwu_time::{SimDuration, SimTime};
9use serde::{Deserialize, Serialize};
10use std::cmp::Ordering;
11use std::collections::{BTreeMap, BTreeSet, BinaryHeap, VecDeque};
12use std::fmt;
13
14pub const ROUTING_ALGORITHM_VERSION: &str = "canwu-routing.v1";
15
16#[derive(Clone, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)]
17#[serde(transparent)]
18pub struct RoutingNodeRef(String);
19
20impl RoutingNodeRef {
21    #[must_use]
22    pub fn new(value: impl Into<String>) -> Self {
23        Self(value.into())
24    }
25
26    #[must_use]
27    pub fn as_str(&self) -> &str {
28        &self.0
29    }
30}
31
32#[derive(Clone, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)]
33#[serde(transparent)]
34pub struct RoutingConnectionRef(String);
35
36impl RoutingConnectionRef {
37    #[must_use]
38    pub fn new(value: impl Into<String>) -> Self {
39        Self(value.into())
40    }
41
42    #[must_use]
43    pub fn as_str(&self) -> &str {
44        &self.0
45    }
46}
47
48#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)]
49#[serde(rename_all = "snake_case")]
50pub enum RoutingEndpointKind {
51    Settlement,
52    RelayStation,
53    RailwayStation,
54    Port,
55    Airport,
56    TelegraphOffice,
57    DeliveryDistrict,
58    MilitaryPosition,
59}
60
61#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)]
62#[serde(rename_all = "snake_case")]
63pub enum TransferMode {
64    Foot,
65    Horse,
66    RoadVehicle,
67    RiverBoat,
68    Sea,
69    Rail,
70    Air,
71    Signal,
72}
73
74#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
75pub struct RoutingEndpoint {
76    pub id: RoutingNodeRef,
77    pub kind: RoutingEndpointKind,
78}
79
80#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
81pub struct DepartureSlot {
82    pub departure_at: SimTime,
83    pub duration: SimDuration,
84}
85
86#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
87pub struct DurationSample {
88    pub from: SimTime,
89    pub duration: SimDuration,
90}
91
92#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
93#[serde(tag = "type", rename_all = "snake_case")]
94pub enum TraversalModel {
95    Fixed { duration: SimDuration },
96    Departures { slots: Vec<DepartureSlot> },
97    Piecewise { samples: Vec<DurationSample> },
98}
99
100impl TraversalModel {
101    fn validate(&self) -> Result<(), RoutingError> {
102        match self {
103            Self::Fixed { duration } if duration.is_negative() => Err(
104                RoutingError::InvalidNetwork("negative traversal duration".to_owned()),
105            ),
106            Self::Fixed { .. } => Ok(()),
107            Self::Departures { slots } => {
108                if slots
109                    .windows(2)
110                    .any(|pair| pair[0].departure_at >= pair[1].departure_at)
111                    || slots.iter().any(|slot| slot.duration.is_negative())
112                {
113                    return Err(RoutingError::InvalidNetwork(
114                        "departure slots must be sorted and non-negative".to_owned(),
115                    ));
116                }
117                Ok(())
118            }
119            Self::Piecewise { samples } => {
120                if samples.windows(2).any(|pair| pair[0].from >= pair[1].from)
121                    || samples.iter().any(|sample| sample.duration.is_negative())
122                {
123                    return Err(RoutingError::InvalidNetwork(
124                        "duration samples must be sorted and non-negative".to_owned(),
125                    ));
126                }
127                Ok(())
128            }
129        }
130    }
131
132    fn traverse_after(&self, at: SimTime) -> Option<(SimTime, SimTime)> {
133        match self {
134            Self::Fixed { duration } => Some((at, at.checked_add(*duration)?)),
135            Self::Departures { slots } => slots.iter().find_map(|slot| {
136                (slot.departure_at >= at).then(|| {
137                    Some((
138                        slot.departure_at,
139                        slot.departure_at.checked_add(slot.duration)?,
140                    ))
141                })?
142            }),
143            Self::Piecewise { samples } => samples
144                .iter()
145                .rev()
146                .find(|sample| sample.from <= at)
147                .and_then(|sample| Some((at, at.checked_add(sample.duration)?))),
148        }
149    }
150}
151
152#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
153pub struct RoutingConnection {
154    pub id: RoutingConnectionRef,
155    pub from: RoutingNodeRef,
156    pub to: RoutingNodeRef,
157    pub mode: TransferMode,
158    pub traversal: TraversalModel,
159    pub available_from: Option<SimTime>,
160    pub available_until: Option<SimTime>,
161    pub risk_per_mille: u32,
162    pub resource_cost: u64,
163}
164
165impl RoutingConnection {
166    fn validate(&self, endpoints: &BTreeSet<RoutingNodeRef>) -> Result<(), RoutingError> {
167        if !endpoints.contains(&self.from) || !endpoints.contains(&self.to) || self.from == self.to
168        {
169            return Err(RoutingError::InvalidNetwork(format!(
170                "connection {} has invalid endpoints",
171                self.id.0
172            )));
173        }
174        if self.risk_per_mille > 1_000 {
175            return Err(RoutingError::InvalidNetwork(format!(
176                "connection {} risk exceeds 1000 per mille",
177                self.id.0
178            )));
179        }
180        if let (Some(start), Some(end)) = (self.available_from, self.available_until)
181            && start > end
182        {
183            return Err(RoutingError::InvalidNetwork(format!(
184                "connection {} availability is inverted",
185                self.id.0
186            )));
187        }
188        self.traversal.validate()
189    }
190
191    fn traverse_after(&self, at: SimTime) -> Option<(SimTime, SimTime)> {
192        let at = self.available_from.map_or(at, |start| at.max(start));
193        if self.available_until.is_some_and(|end| at > end) {
194            return None;
195        }
196        let result = self.traversal.traverse_after(at)?;
197        self.available_until
198            .is_none_or(|end| result.1 <= end)
199            .then_some(result)
200    }
201}
202
203#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
204pub struct RoutingNetwork {
205    pub version: String,
206    pub endpoints: Vec<RoutingEndpoint>,
207    pub connections: Vec<RoutingConnection>,
208}
209
210impl RoutingNetwork {
211    pub fn new(
212        version: impl Into<String>,
213        mut endpoints: Vec<RoutingEndpoint>,
214        mut connections: Vec<RoutingConnection>,
215    ) -> Result<Self, RoutingError> {
216        endpoints.sort_by(|left, right| left.id.cmp(&right.id));
217        connections.sort_by(|left, right| left.id.cmp(&right.id));
218        let endpoint_ids = endpoints
219            .iter()
220            .map(|endpoint| endpoint.id.clone())
221            .collect::<Vec<_>>();
222        if endpoint_ids.windows(2).any(|pair| pair[0] == pair[1]) {
223            return Err(RoutingError::InvalidNetwork(
224                "duplicate endpoint".to_owned(),
225            ));
226        }
227        let endpoint_set = endpoint_ids.iter().cloned().collect::<BTreeSet<_>>();
228        if connections.windows(2).any(|pair| pair[0].id == pair[1].id) {
229            return Err(RoutingError::InvalidNetwork(
230                "duplicate connection".to_owned(),
231            ));
232        }
233        for connection in &connections {
234            connection.validate(&endpoint_set)?;
235        }
236        Ok(Self {
237            version: version.into(),
238            endpoints,
239            connections,
240        })
241    }
242}
243
244#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
245pub struct PlanningSnapshot {
246    pub observer: String,
247    pub observed_at: SimTime,
248    pub valid_until: Option<SimTime>,
249    pub knowledge_cut: String,
250    pub topology_version: String,
251    pub timetable_version: Option<String>,
252    pub network: RoutingNetwork,
253}
254
255impl PlanningSnapshot {
256    pub fn validate(&self) -> Result<(), RoutingError> {
257        if self
258            .valid_until
259            .is_some_and(|until| until < self.observed_at)
260        {
261            return Err(RoutingError::InvalidSnapshot(
262                "planning snapshot expires before it is observed".to_owned(),
263            ));
264        }
265        if self.topology_version != self.network.version {
266            return Err(RoutingError::InvalidSnapshot(
267                "topology version does not match network version".to_owned(),
268            ));
269        }
270        Ok(())
271    }
272
273    #[must_use]
274    pub fn digest(&self) -> String {
275        canonical_digest(self)
276    }
277}
278
279#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, Ord, PartialEq, PartialOrd, Serialize)]
280#[serde(rename_all = "snake_case")]
281pub enum RoutingAlgorithm {
282    FifoDijkstraV1,
283    BoundedLabelCorrectingV1,
284}
285
286#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
287pub struct RoutingPolicy {
288    pub version: String,
289    pub algorithm: RoutingAlgorithm,
290    pub allowed_modes: BTreeSet<TransferMode>,
291    pub max_arrival_at: Option<SimTime>,
292    pub max_expanded_nodes: usize,
293    pub max_transfers: usize,
294    pub max_risk_per_mille: u64,
295}
296
297impl Default for RoutingPolicy {
298    fn default() -> Self {
299        Self {
300            version: "canwu-routing.policy.v1".to_owned(),
301            algorithm: RoutingAlgorithm::FifoDijkstraV1,
302            allowed_modes: BTreeSet::new(),
303            max_arrival_at: None,
304            max_expanded_nodes: 10_000,
305            max_transfers: 64,
306            max_risk_per_mille: 1_000,
307        }
308    }
309}
310
311impl RoutingPolicy {
312    fn allows(&self, mode: TransferMode) -> bool {
313        self.allowed_modes.is_empty() || self.allowed_modes.contains(&mode)
314    }
315}
316
317#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
318pub struct RoutingRequest {
319    pub origin: RoutingNodeRef,
320    pub destination: RoutingNodeRef,
321    pub departure_at: SimTime,
322    pub policy: RoutingPolicy,
323}
324
325#[derive(Clone, Debug, Deserialize, Eq, Ord, PartialEq, PartialOrd, Serialize)]
326pub struct RouteLeg {
327    pub connection: RoutingConnectionRef,
328    pub from: RoutingNodeRef,
329    pub to: RoutingNodeRef,
330    pub mode: TransferMode,
331    pub planned_departure_at: SimTime,
332    pub planned_arrival_at: SimTime,
333}
334
335#[derive(Clone, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
336pub struct RouteCost {
337    pub estimated_arrival_at: SimTime,
338    pub risk_per_mille: u64,
339    pub resource_cost: u64,
340    pub transfers: usize,
341}
342
343#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
344pub struct RoutePlan {
345    pub algorithm_version: String,
346    pub policy_version: String,
347    pub planning_snapshot_digest: String,
348    pub origin: RoutingNodeRef,
349    pub destination: RoutingNodeRef,
350    pub departure_at: SimTime,
351    pub estimated_arrival_at: SimTime,
352    pub cost: RouteCost,
353    pub legs: Vec<RouteLeg>,
354    pub digest: String,
355}
356
357#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
358pub enum RoutingError {
359    InvalidNetwork(String),
360    InvalidSnapshot(String),
361    UnknownOrigin,
362    UnknownDestination,
363    NoKnownRoute,
364    SearchHorizonExceeded,
365    ExpansionBudgetExceeded,
366    RequirementsUnsatisfied,
367    ArithmeticOverflow,
368}
369
370impl fmt::Display for RoutingError {
371    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
372        match self {
373            Self::InvalidNetwork(message) => {
374                write!(formatter, "invalid routing network: {message}")
375            }
376            Self::InvalidSnapshot(message) => {
377                write!(formatter, "invalid planning snapshot: {message}")
378            }
379            Self::UnknownOrigin => {
380                formatter.write_str("routing origin is not present in the network")
381            }
382            Self::UnknownDestination => {
383                formatter.write_str("routing destination is not present in the network")
384            }
385            Self::NoKnownRoute => formatter.write_str("no route satisfies the routing policy"),
386            Self::SearchHorizonExceeded => {
387                formatter.write_str("route departure is outside the planning snapshot horizon")
388            }
389            Self::ExpansionBudgetExceeded => {
390                formatter.write_str("routing expansion budget was exceeded")
391            }
392            Self::RequirementsUnsatisfied => {
393                formatter.write_str("routing requirements are not satisfied")
394            }
395            Self::ArithmeticOverflow => formatter.write_str("routing arithmetic overflowed"),
396        }
397    }
398}
399
400impl std::error::Error for RoutingError {}
401
402#[derive(Clone, Debug, Eq, PartialEq)]
403struct Label {
404    node: RoutingNodeRef,
405    arrival_at: SimTime,
406    risk_per_mille: u64,
407    resource_cost: u64,
408    transfers: usize,
409    legs: Vec<RouteLeg>,
410}
411
412impl Label {
413    fn key(&self) -> (&SimTime, u64, u64, usize, &Vec<RouteLeg>) {
414        (
415            &self.arrival_at,
416            self.risk_per_mille,
417            self.resource_cost,
418            self.transfers,
419            &self.legs,
420        )
421    }
422}
423
424#[derive(Clone, Debug, Eq, PartialEq)]
425struct QueueEntry {
426    node: RoutingNodeRef,
427    arrival_at: SimTime,
428    risk_per_mille: u64,
429    resource_cost: u64,
430    transfers: usize,
431    legs: Vec<RouteLeg>,
432}
433
434impl QueueEntry {
435    fn from_label(label: &Label) -> Self {
436        Self {
437            node: label.node.clone(),
438            arrival_at: label.arrival_at,
439            risk_per_mille: label.risk_per_mille,
440            resource_cost: label.resource_cost,
441            transfers: label.transfers,
442            legs: label.legs.clone(),
443        }
444    }
445
446    fn key(&self) -> (&SimTime, u64, u64, usize, &RoutingNodeRef, &Vec<RouteLeg>) {
447        (
448            &self.arrival_at,
449            self.risk_per_mille,
450            self.resource_cost,
451            self.transfers,
452            &self.node,
453            &self.legs,
454        )
455    }
456}
457
458impl Ord for QueueEntry {
459    fn cmp(&self, other: &Self) -> Ordering {
460        other.key().cmp(&self.key())
461    }
462}
463
464impl PartialOrd for QueueEntry {
465    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
466        Some(self.cmp(other))
467    }
468}
469
470fn better(left: &Label, right: &Label) -> bool {
471    left.key() < right.key()
472}
473
474fn adjacency(network: &RoutingNetwork) -> BTreeMap<RoutingNodeRef, Vec<&RoutingConnection>> {
475    let mut index = BTreeMap::new();
476    for connection in &network.connections {
477        index
478            .entry(connection.from.clone())
479            .or_insert_with(Vec::new)
480            .push(connection);
481    }
482    index
483}
484
485fn outgoing<'a>(
486    index: &'a BTreeMap<RoutingNodeRef, Vec<&'a RoutingConnection>>,
487    node: &RoutingNodeRef,
488) -> impl Iterator<Item = &'a RoutingConnection> {
489    index.get(node).into_iter().flatten().copied()
490}
491
492fn extend(label: &Label, connection: &RoutingConnection, policy: &RoutingPolicy) -> Option<Label> {
493    if !policy.allows(connection.mode) || label.transfers >= policy.max_transfers {
494        return None;
495    }
496    let (departure_at, arrival_at) = connection.traverse_after(label.arrival_at)?;
497    let risk_per_mille = label
498        .risk_per_mille
499        .checked_add(u64::from(connection.risk_per_mille))?;
500    let resource_cost = label.resource_cost.checked_add(connection.resource_cost)?;
501    let transfers = label.transfers + 1;
502    if risk_per_mille > policy.max_risk_per_mille
503        || policy
504            .max_arrival_at
505            .is_some_and(|limit| arrival_at > limit)
506    {
507        return None;
508    }
509    let mut legs = label.legs.clone();
510    legs.push(RouteLeg {
511        connection: connection.id.clone(),
512        from: connection.from.clone(),
513        to: connection.to.clone(),
514        mode: connection.mode,
515        planned_departure_at: departure_at,
516        planned_arrival_at: arrival_at,
517    });
518    Some(Label {
519        node: connection.to.clone(),
520        arrival_at,
521        risk_per_mille,
522        resource_cost,
523        transfers,
524        legs,
525    })
526}
527
528fn solve_dijkstra(
529    request: &RoutingRequest,
530    index: &BTreeMap<RoutingNodeRef, Vec<&RoutingConnection>>,
531) -> Result<Label, RoutingError> {
532    let mut labels = BTreeMap::<RoutingNodeRef, Label>::new();
533    let start = Label {
534        node: request.origin.clone(),
535        arrival_at: request.departure_at,
536        risk_per_mille: 0,
537        resource_cost: 0,
538        transfers: 0,
539        legs: Vec::new(),
540    };
541    labels.insert(request.origin.clone(), start.clone());
542    let mut queue = BinaryHeap::from([QueueEntry::from_label(&start)]);
543    let mut expanded = 0;
544    while let Some(entry) = queue.pop() {
545        expanded += 1;
546        if expanded > request.policy.max_expanded_nodes {
547            return Err(RoutingError::ExpansionBudgetExceeded);
548        }
549        let Some(current) = labels.get(&entry.node).cloned() else {
550            continue;
551        };
552        if current.arrival_at != entry.arrival_at
553            || current.risk_per_mille != entry.risk_per_mille
554            || current.resource_cost != entry.resource_cost
555            || current.transfers != entry.transfers
556            || current.legs != entry.legs
557        {
558            // The label may have been improved since this queue entry was added.
559            continue;
560        }
561        if entry.node == request.destination {
562            return Ok(current);
563        }
564        for connection in outgoing(index, &entry.node) {
565            let Some(candidate) = extend(&current, connection, &request.policy) else {
566                continue;
567            };
568            let replace = labels
569                .get(&candidate.node)
570                .is_none_or(|existing| better(&candidate, existing));
571            if replace {
572                queue.push(QueueEntry::from_label(&candidate));
573                labels.insert(candidate.node.clone(), candidate);
574            }
575        }
576    }
577    Err(RoutingError::NoKnownRoute)
578}
579
580fn solve_label_correcting(
581    request: &RoutingRequest,
582    index: &BTreeMap<RoutingNodeRef, Vec<&RoutingConnection>>,
583) -> Result<Label, RoutingError> {
584    let mut labels = BTreeMap::<RoutingNodeRef, Vec<Label>>::new();
585    let start = Label {
586        node: request.origin.clone(),
587        arrival_at: request.departure_at,
588        risk_per_mille: 0,
589        resource_cost: 0,
590        transfers: 0,
591        legs: Vec::new(),
592    };
593    labels.insert(request.origin.clone(), vec![start.clone()]);
594    let mut queue = VecDeque::from([start]);
595    let mut expanded = 0;
596    while let Some(current) = queue.pop_front() {
597        expanded += 1;
598        if expanded > request.policy.max_expanded_nodes {
599            return Err(RoutingError::ExpansionBudgetExceeded);
600        }
601        for connection in outgoing(index, &current.node) {
602            let Some(candidate) = extend(&current, connection, &request.policy) else {
603                continue;
604            };
605            let node_labels = labels.entry(candidate.node.clone()).or_default();
606            if node_labels.iter().any(|existing| existing == &candidate) {
607                continue;
608            }
609            node_labels.push(candidate.clone());
610            queue.push_back(candidate);
611        }
612    }
613    labels
614        .remove(&request.destination)
615        .and_then(|candidates| {
616            candidates
617                .into_iter()
618                .min_by(|left, right| left.key().cmp(&right.key()))
619        })
620        .ok_or(RoutingError::NoKnownRoute)
621}
622
623pub fn plan_route(
624    snapshot: &PlanningSnapshot,
625    request: &RoutingRequest,
626) -> Result<RoutePlan, RoutingError> {
627    snapshot.validate()?;
628    if !snapshot
629        .network
630        .endpoints
631        .iter()
632        .any(|endpoint| endpoint.id == request.origin)
633    {
634        return Err(RoutingError::UnknownOrigin);
635    }
636    if !snapshot
637        .network
638        .endpoints
639        .iter()
640        .any(|endpoint| endpoint.id == request.destination)
641    {
642        return Err(RoutingError::UnknownDestination);
643    }
644    if snapshot
645        .valid_until
646        .is_some_and(|until| request.departure_at > until)
647    {
648        return Err(RoutingError::SearchHorizonExceeded);
649    }
650    if request.origin == request.destination {
651        let mut plan = RoutePlan {
652            algorithm_version: ROUTING_ALGORITHM_VERSION.to_owned(),
653            policy_version: request.policy.version.clone(),
654            planning_snapshot_digest: snapshot.digest(),
655            origin: request.origin.clone(),
656            destination: request.destination.clone(),
657            departure_at: request.departure_at,
658            estimated_arrival_at: request.departure_at,
659            cost: RouteCost {
660                estimated_arrival_at: request.departure_at,
661                ..RouteCost::default()
662            },
663            legs: Vec::new(),
664            digest: String::new(),
665        };
666        plan.digest = canonical_digest(&plan_without_digest(&plan));
667        return Ok(plan);
668    }
669    let index = adjacency(&snapshot.network);
670    let label = match request.policy.algorithm {
671        RoutingAlgorithm::FifoDijkstraV1 => solve_dijkstra(request, &index)?,
672        RoutingAlgorithm::BoundedLabelCorrectingV1 => solve_label_correcting(request, &index)?,
673    };
674    let mut plan = RoutePlan {
675        algorithm_version: ROUTING_ALGORITHM_VERSION.to_owned(),
676        policy_version: request.policy.version.clone(),
677        planning_snapshot_digest: snapshot.digest(),
678        origin: request.origin.clone(),
679        destination: request.destination.clone(),
680        departure_at: request.departure_at,
681        estimated_arrival_at: label.arrival_at,
682        cost: RouteCost {
683            estimated_arrival_at: label.arrival_at,
684            risk_per_mille: label.risk_per_mille,
685            resource_cost: label.resource_cost,
686            transfers: label.transfers,
687        },
688        legs: label.legs,
689        digest: String::new(),
690    };
691    plan.digest = canonical_digest(&plan_without_digest(&plan));
692    Ok(plan)
693}
694
695fn plan_without_digest(
696    plan: &RoutePlan,
697) -> (
698    &str,
699    &str,
700    &str,
701    &RoutingNodeRef,
702    &RoutingNodeRef,
703    SimTime,
704    SimTime,
705    &RouteCost,
706    &Vec<RouteLeg>,
707) {
708    (
709        &plan.algorithm_version,
710        &plan.policy_version,
711        &plan.planning_snapshot_digest,
712        &plan.origin,
713        &plan.destination,
714        plan.departure_at,
715        plan.estimated_arrival_at,
716        &plan.cost,
717        &plan.legs,
718    )
719}
720
721fn canonical_digest<T: Serialize>(value: &T) -> String {
722    let bytes = serde_json::to_vec(value).expect("routing types must be serializable");
723    blake3::hash(&bytes).to_hex().to_string()
724}
725
726#[derive(Clone, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
727pub struct RoutingCache {
728    entries: BTreeMap<String, RoutePlan>,
729}
730
731impl RoutingCache {
732    #[must_use]
733    pub fn key(snapshot: &PlanningSnapshot, request: &RoutingRequest) -> String {
734        canonical_digest(&(snapshot.digest(), request))
735    }
736
737    #[must_use]
738    pub fn get(&self, key: &str) -> Option<&RoutePlan> {
739        self.entries.get(key)
740    }
741
742    pub fn insert(&mut self, key: String, plan: RoutePlan) {
743        self.entries.insert(key, plan);
744    }
745
746    pub fn clear(&mut self) {
747        self.entries.clear();
748    }
749}
750
751#[cfg(test)]
752mod tests {
753    use super::*;
754
755    fn endpoint(id: &str) -> RoutingEndpoint {
756        RoutingEndpoint {
757            id: RoutingNodeRef::new(id),
758            kind: RoutingEndpointKind::Settlement,
759        }
760    }
761
762    fn connection(id: &str, from: &str, to: &str, minutes: i64) -> RoutingConnection {
763        RoutingConnection {
764            id: RoutingConnectionRef::new(id),
765            from: RoutingNodeRef::new(from),
766            to: RoutingNodeRef::new(to),
767            mode: TransferMode::Horse,
768            traversal: TraversalModel::Fixed {
769                duration: SimDuration::minutes(minutes),
770            },
771            available_from: None,
772            available_until: None,
773            risk_per_mille: 0,
774            resource_cost: 0,
775        }
776    }
777
778    fn snapshot(connections: Vec<RoutingConnection>) -> PlanningSnapshot {
779        snapshot_with_endpoints(connections, ["a", "b", "c"])
780    }
781
782    fn snapshot_with_endpoints<const N: usize>(
783        connections: Vec<RoutingConnection>,
784        endpoint_ids: [&str; N],
785    ) -> PlanningSnapshot {
786        let network = RoutingNetwork::new(
787            "roads.v1",
788            endpoint_ids.into_iter().map(endpoint).collect(),
789            connections,
790        )
791        .unwrap();
792        PlanningSnapshot {
793            observer: "courier".to_owned(),
794            observed_at: SimTime::EPOCH,
795            valid_until: None,
796            knowledge_cut: "knowledge:1".to_owned(),
797            topology_version: "roads.v1".to_owned(),
798            timetable_version: None,
799            network,
800        }
801    }
802
803    #[test]
804    fn chooses_deterministic_earliest_arrival_across_multiple_legs() {
805        let snapshot = snapshot(vec![
806            connection("a-c", "a", "c", 90),
807            connection("a-b", "a", "b", 30),
808            connection("b-c", "b", "c", 30),
809        ]);
810        let plan = plan_route(
811            &snapshot,
812            &RoutingRequest {
813                origin: RoutingNodeRef::new("a"),
814                destination: RoutingNodeRef::new("c"),
815                departure_at: SimTime::EPOCH,
816                policy: RoutingPolicy::default(),
817            },
818        )
819        .unwrap();
820        assert_eq!(plan.estimated_arrival_at, SimTime::from_minutes(60));
821        assert_eq!(
822            plan.legs
823                .iter()
824                .map(|leg| leg.connection.clone())
825                .collect::<Vec<_>>(),
826            vec![
827                RoutingConnectionRef::new("a-b"),
828                RoutingConnectionRef::new("b-c")
829            ]
830        );
831    }
832
833    #[test]
834    fn scheduled_connections_wait_for_the_next_departure() {
835        let mut rail = connection("rail", "a", "c", 1);
836        rail.mode = TransferMode::Rail;
837        rail.traversal = TraversalModel::Departures {
838            slots: vec![DepartureSlot {
839                departure_at: SimTime::from_minutes(60),
840                duration: SimDuration::minutes(20),
841            }],
842        };
843        let snapshot = snapshot(vec![rail]);
844        let plan = plan_route(
845            &snapshot,
846            &RoutingRequest {
847                origin: RoutingNodeRef::new("a"),
848                destination: RoutingNodeRef::new("c"),
849                departure_at: SimTime::from_minutes(10),
850                policy: RoutingPolicy::default(),
851            },
852        )
853        .unwrap();
854        assert_eq!(plan.estimated_arrival_at, SimTime::from_minutes(80));
855    }
856
857    #[test]
858    fn label_correcting_handles_piecewise_duration_changes() {
859        let mut edge = connection("edge", "a", "c", 10);
860        edge.traversal = TraversalModel::Piecewise {
861            samples: vec![
862                DurationSample {
863                    from: SimTime::EPOCH,
864                    duration: SimDuration::minutes(100),
865                },
866                DurationSample {
867                    from: SimTime::from_minutes(10),
868                    duration: SimDuration::minutes(5),
869                },
870            ],
871        };
872        let snapshot = snapshot(vec![edge]);
873        let policy = RoutingPolicy {
874            algorithm: RoutingAlgorithm::BoundedLabelCorrectingV1,
875            ..RoutingPolicy::default()
876        };
877        let plan = plan_route(
878            &snapshot,
879            &RoutingRequest {
880                origin: RoutingNodeRef::new("a"),
881                destination: RoutingNodeRef::new("c"),
882                departure_at: SimTime::from_minutes(10),
883                policy,
884            },
885        )
886        .unwrap();
887        assert_eq!(plan.estimated_arrival_at, SimTime::from_minutes(15));
888    }
889
890    #[test]
891    fn label_correcting_keeps_a_later_label_for_a_faster_non_fifo_departure() {
892        let mut final_leg = connection("b-d", "b", "d", 100);
893        final_leg.traversal = TraversalModel::Piecewise {
894            samples: vec![
895                DurationSample {
896                    from: SimTime::EPOCH,
897                    duration: SimDuration::minutes(100),
898                },
899                DurationSample {
900                    from: SimTime::from_minutes(10),
901                    duration: SimDuration::minutes(1),
902                },
903            ],
904        };
905        let snapshot = snapshot_with_endpoints(
906            vec![
907                connection("a-b", "a", "b", 5),
908                connection("a-c", "a", "c", 10),
909                connection("c-b", "c", "b", 0),
910                final_leg,
911            ],
912            ["a", "b", "c", "d"],
913        );
914        let policy = RoutingPolicy {
915            algorithm: RoutingAlgorithm::BoundedLabelCorrectingV1,
916            ..RoutingPolicy::default()
917        };
918        let plan = plan_route(
919            &snapshot,
920            &RoutingRequest {
921                origin: RoutingNodeRef::new("a"),
922                destination: RoutingNodeRef::new("d"),
923                departure_at: SimTime::EPOCH,
924                policy,
925            },
926        )
927        .unwrap();
928        assert_eq!(plan.estimated_arrival_at, SimTime::from_minutes(11));
929        assert_eq!(plan.legs[1].connection, RoutingConnectionRef::new("c-b"));
930    }
931
932    #[test]
933    fn cache_key_changes_with_snapshot_or_policy() {
934        let snapshot = snapshot(vec![connection("a-c", "a", "c", 10)]);
935        let request = RoutingRequest {
936            origin: RoutingNodeRef::new("a"),
937            destination: RoutingNodeRef::new("c"),
938            departure_at: SimTime::EPOCH,
939            policy: RoutingPolicy::default(),
940        };
941        let first = RoutingCache::key(&snapshot, &request);
942        let mut changed = request.clone();
943        changed.policy.version = "policy.v2".to_owned();
944        assert_ne!(first, RoutingCache::key(&snapshot, &changed));
945    }
946}