///|
/// Context handed to an edge callback: the mutable session-variable map for
/// the running case, bytes received from earlier steps, and the destination
/// request name. Mirrors upstream's ProtocolSession (protocol_session.py,
/// 518c139); callbacks write variables and may return replacement data.
pub(all) struct StepContext {
  variables : Map[String, Bytes]
  received : Array[Bytes]
  request : String
}

///|
pub struct SessionGraph {
  requests : Array[CompiledRequest]
  edges : Array[Array[Int]]
  names : Map[String, Int]
  edge_callbacks : Map[(Int, Int), (StepContext) -> Bytes?]
}

///|
pub fn SessionGraph::new() -> SessionGraph {
  { requests: [], edges: [], names: Map([]), edge_callbacks: Map([]), }
}

///|
pub fn SessionGraph::add(
  self : SessionGraph,
  request : CompiledRequest,
) -> Unit raise ModelError {
  guard !self.names.contains(request.name) else {
    raise Invalid("duplicate session request: " + request.name)
  }
  self.names[request.name] = self.requests.length()
  self.requests.push(request)
  self.edges.push([])
}

///|
fn SessionGraph::resolve(
  self : SessionGraph,
  name : String,
) -> Int raise ModelError {
  self.names
  .get(name)
  .unwrap_or_else(() => raise Invalid("unknown session request: " + name))
}

///|
pub fn SessionGraph::connect(
  self : SessionGraph,
  from : String,
  to : String,
  callback? : (StepContext) -> Bytes?,
) -> Unit raise ModelError {
  let source = self.resolve(from)
  let target = self.resolve(to)
  guard !self.edges[source].contains(target) else {
    raise Invalid("duplicate session edge")
  }
  let pending = [target]
  let seen = Array::make(self.requests.length(), false)
  while pending.pop() is Some(index) {
    guard index != source else { raise Invalid("session graph cycle") }
    if seen[index] {
      continue
    }
    seen[index] = true
    for next in self.edges[index] {
      pending.push(next)
    }
  }
  self.edges[source].push(target)
  match callback {
    Some(function) => self.edge_callbacks[(source, target)] = function
    None => ()
  }
}

///|
pub struct SessionPath {
  requests : Array[CompiledRequest]
  /// Callback on the edge into each request; roots have None. Sourced from
  /// SessionGraph::connect (upstream edge callbacks, session.py:754-794).
  transitions : Array[((StepContext) -> Bytes?)?]
}

///|
pub fn SessionPath::requests(self : SessionPath) -> Array[CompiledRequest] {
  self.requests.copy()
}

///|
pub fn SessionPath::names(self : SessionPath) -> Array[String] {
  self.requests.map(request => request.name)
}

///|
pub fn SessionPath::target(self : SessionPath) -> CompiledRequest {
  self.requests[self.requests.length() - 1]
}

///|
/// Only the terminal request mutates; prefixes are rendered from their defaults.
pub fn SessionPath::cases(
  self : SessionPath,
  start? : Int = 0,
  limit? : Int = 10000,
  variables? : Map[String, Bytes]? = None,
) -> CaseStream raise ModelError {
  self.target().cases(start~, limit~, vars=variables)
}

///|
pub fn SessionPath::cases_with_variables(
  self : SessionPath,
  vars : Map[String, Bytes]?,
  start? : Int = 0,
  limit? : Int = 10000,
) -> CaseStream raise ModelError {
  self.target().cases_with_variables(vars, start~, limit~)
}

///|
pub fn SessionPath::combinatorial_cases_with_variables(
  self : SessionPath,
  vars : Map[String, Bytes]?,
  start? : Int = 0,
  limit? : Int = 10000,
  max_depth? : Int,
) -> CaseStream raise ModelError {
  match max_depth {
    Some(depth) =>
      self
      .target()
      .combinatorial_cases_with_variables(vars, start~, limit~, max_depth=depth)
    None =>
      self.target().combinatorial_cases_with_variables(vars, start~, limit~)
  }
}

///|
pub fn SessionPath::prefix(self : SessionPath) -> Array[Bytes] raise ModelError {
  self.requests[:self.requests.length() - 1]
  .to_owned()
  .map(request => request.render())
}

///|
fn SessionGraph::walk(
  self : SessionGraph,
  index : Int,
  prefix : Array[CompiledRequest],
  transitions : Array[((StepContext) -> Bytes?)?],
  targets : Array[Int]?,
  reachable : Array[Bool],
  max_paths : Int,
  output : Array[SessionPath],
) -> Unit raise ModelError {
  guard prefix.length() < 256 else {
    raise Limit("session path exceeds 256 requests")
  }
  let path = prefix.copy()
  path.push(self.requests[index])
  // Transitions stay aligned with the path: one entry per request, the
  // callback on the edge into it (None for roots).
  let path_transitions = transitions.copy()
  if path_transitions.is_empty() {
    path_transitions.push(None)
  }
  if targets.map(values => values.contains(index)).unwrap_or(true) {
    guard output.length() < max_paths else {
      raise Limit("session path count limit")
    }
    output.push({ requests: path, transitions: path_transitions, })
  }
  for child in self.edges[index] {
    if reachable[child] {
      let child_transitions = path_transitions.copy()
      child_transitions.push(self.edge_callbacks.get((index, child)))
      self.walk(
        child, path, child_transitions, targets, reachable, max_paths, output,
      )
    }
  }
}

///|
/// Roots follow node insertion order; outgoing edges follow connection order.
pub fn SessionGraph::paths(
  self : SessionGraph,
  targets? : Array[String],
  max_paths? : Int = 10000,
) -> Array[SessionPath] raise ModelError {
  guard max_paths >= 0 else { raise Invalid("negative path limit") }
  let selected = targets.map(names => names.map(name => self.resolve(name)))
  let reachable = Array::make(self.requests.length(), selected is None)
  if selected is Some(indices) {
    let pending = indices.copy()
    while pending.pop() is Some(index) {
      if reachable[index] {
        continue
      }
      reachable[index] = true
      for source, edges in self.edges {
        if edges.contains(index) {
          pending.push(source)
        }
      }
    }
  }
  let incoming = Array::make(self.requests.length(), false)
  for edges in self.edges {
    for target in edges {
      incoming[target] = true
    }
  }
  let output : Array[SessionPath] = []
  for index in 0..