From 9d12369fc4332b342cca4d13f64164ae999df0e9 Mon Sep 17 00:00:00 2001 From: Junyi Ou Date: Sat, 26 Sep 2026 18:48:07 -0400 Subject: [PATCH 1/3] feat(terminal-streamer): decode a complete TRP recording The packet decoding is shared with the live stream decoder, so a finished recording can be turned into asciicast without tailing it. Co-Authored-By: Claude Opus 5.5 (1M context) --- crates/terminal-streamer/src/trp_decoder.rs | 254 ++++++++++++++------ 1 file changed, 179 insertions(+), 75 deletions(-) diff --git a/crates/terminal-streamer/src/trp_decoder.rs b/crates/terminal-streamer/src/trp_decoder.rs index 751ec3d8c..861e86186 100644 --- a/crates/terminal-streamer/src/trp_decoder.rs +++ b/crates/terminal-streamer/src/trp_decoder.rs @@ -80,107 +80,163 @@ impl AsyncRead for AsyncReadChannel { } } -async fn parse_trp_stream( - mut input_stream: impl AsyncRead + Unpin + Send + 'static, - mut tx: tokio::sync::mpsc::Sender>, -) -> anyhow::Result<()> { - let mut time = 0.0; - let mut before_setup_cache = Some(Vec::new()); - let mut header = AsciinemaHeader::default(); +/// Decodes a complete TRP recording into asciicast v2 lines, each ending with a newline. +/// +/// A truncated last packet, as left by an interrupted recording, is ignored. +pub fn decode_to_asciicast(mut input: &[u8]) -> anyhow::Result { + let mut decoder = TrpDecoder::default(); + let mut output = String::new(); - loop { - let mut packet_head_buffer = [0u8; 8]; - if let Err(e) = input_stream.read_exact(&mut packet_head_buffer).await { - if e.kind() == std::io::ErrorKind::UnexpectedEof { - tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; - continue; - } - anyhow::bail!(e); + while input.len() >= PACKET_HEADER_SIZE { + let (header, rest) = input.split_at(PACKET_HEADER_SIZE); + let header = PacketHeader::parse(header.try_into()?); + + let Some((payload, rest)) = rest.split_at_checked(usize::from(header.size)) else { + break; + }; + input = rest; + + for line in decoder.push(&header, payload)? { + output.push_str(&line); + output.push('\n'); } + } - let time_delta = u32::from_le_bytes(packet_head_buffer[0..4].try_into()?); - let event_type = u16::from_le_bytes(packet_head_buffer[4..6].try_into()?); - let size = u16::from_le_bytes(packet_head_buffer[6..8].try_into()?); + for line in decoder.finish() { + output.push_str(&line); + output.push('\n'); + } - time += f64::from(time_delta) / 1000.0; + Ok(output) +} - let mut event_payload = vec![0u8; size as usize]; - if let Err(e) = input_stream.read_exact(&mut event_payload).await { - if e.kind() == std::io::ErrorKind::UnexpectedEof { - tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; - continue; - } - anyhow::bail!(e); +const PACKET_HEADER_SIZE: usize = 8; + +struct PacketHeader { + time_delta: u32, + event_type: u16, + size: u16, +} + +impl PacketHeader { + fn parse(buffer: &[u8; PACKET_HEADER_SIZE]) -> Self { + Self { + time_delta: u32::from_le_bytes([buffer[0], buffer[1], buffer[2], buffer[3]]), + event_type: u16::from_le_bytes([buffer[4], buffer[5]]), + size: u16::from_le_bytes([buffer[6], buffer[7]]), } + } +} - match event_type { - 0 => { - // Terminal output - let event_payload = String::from_utf8_lossy(&event_payload).into_owned(); - let event = AsciinemaEvent::TerminalOutput { - payload: event_payload, - time, - }; - match before_setup_cache { - Some(ref mut cache) => { - cache.push(event); - } - None => { - send(&mut tx, event.to_json()).await?; - } - } - } - 1 => { - let event_payload = String::from_utf8_lossy(&event_payload).into_owned(); - let event = AsciinemaEvent::UserInput { - payload: event_payload, - time, +/// Turns TRP packets into asciicast lines, holding the events that come before the terminal setup. +struct TrpDecoder { + time: f64, + before_setup_cache: Option>, + header: AsciinemaHeader, +} + +impl Default for TrpDecoder { + fn default() -> Self { + Self { + time: 0.0, + before_setup_cache: Some(Vec::new()), + header: AsciinemaHeader::default(), + } + } +} + +impl TrpDecoder { + fn push(&mut self, packet: &PacketHeader, payload: &[u8]) -> anyhow::Result> { + self.time += f64::from(packet.time_delta) / 1000.0; + let time = self.time; + + let mut lines = Vec::new(); + + match packet.event_type { + 0 | 1 => { + let payload = String::from_utf8_lossy(payload).into_owned(); + let event = if packet.event_type == 0 { + AsciinemaEvent::TerminalOutput { payload, time } + } else { + AsciinemaEvent::UserInput { payload, time } }; - match before_setup_cache { - Some(ref mut cache) => { - cache.push(event); - } - None => { - send(&mut tx, event.to_json()).await?; - } + match self.before_setup_cache { + Some(ref mut cache) => cache.push(event), + None => lines.push(event.to_json()), } } 2 => { // Terminal size change. Payload is little-endian [columns, rows]. - if event_payload.len() < 4 { - anyhow::bail!( - "invalid terminal size change payload length (len={})", - event_payload.len() - ); + if payload.len() < 4 { + anyhow::bail!("invalid terminal size change payload length (len={})", payload.len()); } - header.col = u16::from_le_bytes(event_payload[0..2].try_into()?); - header.row = u16::from_le_bytes(event_payload[2..4].try_into()?); - if before_setup_cache.is_none() { + self.header.col = u16::from_le_bytes([payload[0], payload[1]]); + self.header.row = u16::from_le_bytes([payload[2], payload[3]]); + if self.before_setup_cache.is_none() { let event = AsciinemaEvent::Resize { - width: header.col, - height: header.row, + width: self.header.col, + height: self.header.row, time, }; - send(&mut tx, event.to_json()).await?; + lines.push(event.to_json()); } } 4 => { // Terminal setup - if before_setup_cache.is_some() { - let header_json = header.to_json(); - send(&mut tx, header_json).await?; - if let Some(ref mut cache) = before_setup_cache { - for event in cache.drain(..) { - send(&mut tx, event.to_json()).await?; - } - } - before_setup_cache = None; + if let Some(cache) = self.before_setup_cache.take() { + lines.push(self.header.to_json()); + lines.extend(cache.iter().map(AsciinemaEvent::to_json)); } else { warn!("Received terminal setup event but cache is empty"); } } _ => {} } + + Ok(lines) + } + + /// Flushes the cached events of a recording that never sent its terminal setup. + fn finish(self) -> Vec { + match self.before_setup_cache { + Some(cache) => core::iter::once(self.header.to_json()) + .chain(cache.iter().map(AsciinemaEvent::to_json)) + .collect(), + None => Vec::new(), + } + } +} + +async fn parse_trp_stream( + mut input_stream: impl AsyncRead + Unpin + Send + 'static, + mut tx: tokio::sync::mpsc::Sender>, +) -> anyhow::Result<()> { + let mut decoder = TrpDecoder::default(); + + loop { + let mut packet_head_buffer = [0u8; PACKET_HEADER_SIZE]; + if let Err(e) = input_stream.read_exact(&mut packet_head_buffer).await { + if e.kind() == std::io::ErrorKind::UnexpectedEof { + tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; + continue; + } + anyhow::bail!(e); + } + + let packet = PacketHeader::parse(&packet_head_buffer); + + let mut event_payload = vec![0u8; usize::from(packet.size)]; + if let Err(e) = input_stream.read_exact(&mut event_payload).await { + if e.kind() == std::io::ErrorKind::UnexpectedEof { + tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; + continue; + } + anyhow::bail!(e); + } + + for line in decoder.push(&packet, &event_payload)? { + send(&mut tx, line).await?; + } } } @@ -189,3 +245,51 @@ async fn send(sender: &mut tokio::sync::mpsc::Sender>, mu sender.send(Ok(json)).await?; Ok(()) } + +#[cfg(test)] +mod tests { + use super::*; + + fn packet(time_delta: u32, event_type: u16, payload: &[u8]) -> Vec { + let mut packet = Vec::new(); + packet.extend_from_slice(&time_delta.to_le_bytes()); + packet.extend_from_slice(&event_type.to_le_bytes()); + packet.extend_from_slice(&u16::try_from(payload.len()).expect("small payload").to_le_bytes()); + packet.extend_from_slice(payload); + packet + } + + #[test] + fn decodes_a_complete_recording_and_ignores_a_truncated_tail() { + let mut trp = Vec::new(); + trp.extend(packet(0, 2, &[100, 0, 30, 0])); + trp.extend(packet(500, 0, b"$ ")); + trp.extend(packet(0, 4, &[])); + trp.extend(packet(1000, 1, b"l")); + trp.extend(packet(250, 0, b"ls\r\n")); + trp.extend(&packet(10, 0, b"lost")[..9]); + + let cast = decode_to_asciicast(&trp).expect("valid recording"); + + assert_eq!( + cast, + concat!( + "{\"version\": 2, \"width\": 100, \"height\": 30}\n", + "[0.5,\"o\",\"$ \"]\n", + "[1.5,\"i\",\"l\"]\n", + r#"[1.75,"o","ls\u000d\u000a"]"#, + "\n", + ) + ); + } + + #[test] + fn events_of_a_recording_without_setup_are_kept() { + let cast = decode_to_asciicast(&packet(2000, 0, b"x")).expect("valid recording"); + + assert_eq!( + cast, + "{\"version\": 2, \"width\": 80, \"height\": 24}\n[2,\"o\",\"x\"]\n" + ); + } +} From c74b258217e1336a2c8ed53a8d54a65069c119e0 Mon Sep 17 00:00:00 2001 From: Junyi Ou Date: Mon, 28 Sep 2026 16:30:02 -0400 Subject: [PATCH 2/3] feat(dgw): generate a session log with AI Runs the ai-log task end to end: streams the terminal recording into chunk files in a task workspace, describes each chunk with the AI (checkpointed so a retry resumes, splitting a chunk when the answer is cut at the output limit), writes the .slog and copies it into a new `ai-analysis` artifact from the recording manager (`add_artifact`, #2003). The task reads a finished recording through the recording manager (`get_finished`). Also regenerates the ai-log substate docs. Co-Authored-By: Claude Opus 5.5 (1M context) --- crates/terminal-streamer/src/trp_decoder.rs | 321 +++++----- devolutions-gateway/openapi/doc/index.adoc | 52 +- .../dotnet-client/.openapi-generator/FILES | 2 + .../openapi/dotnet-client/README.md | 1 + .../dotnet-client/docs/AiLogSubstate.md | 2 + .../dotnet-client/docs/AiLogSubstateOneOf1.md | 13 + .../Model/AiLogSubstate.cs | 49 +- .../Model/AiLogSubstateOneOf1.cs | 132 ++++ devolutions-gateway/openapi/gateway-api.yaml | 19 + .../.openapi-generator/FILES | 1 + .../ts-angular-client/model/aiLogSubstate.ts | 3 +- .../model/aiLogSubstateOneOf1.ts | 27 + .../openapi/ts-angular-client/model/models.ts | 1 + devolutions-gateway/src/recording.rs | 118 +++- devolutions-gateway/src/tasks/ai_log.rs | 171 ----- .../src/tasks/ai_log/checkpoint.rs | 94 +++ devolutions-gateway/src/tasks/ai_log/mod.rs | 443 +++++++++++++ devolutions-gateway/src/tasks/ai_log/slog.rs | 171 +++++ .../src/tasks/ai_log/transcript.rs | 595 ++++++++++++++++++ devolutions-gateway/src/tasks/mod.rs | 87 ++- devolutions-gateway/src/tasks/tests.rs | 288 ++++++++- devolutions-gateway/tests/tasks.rs | 296 ++++++++- 22 files changed, 2529 insertions(+), 357 deletions(-) create mode 100644 devolutions-gateway/openapi/dotnet-client/docs/AiLogSubstateOneOf1.md create mode 100644 devolutions-gateway/openapi/dotnet-client/src/Devolutions.Gateway.Client/Model/AiLogSubstateOneOf1.cs create mode 100644 devolutions-gateway/openapi/ts-angular-client/model/aiLogSubstateOneOf1.ts delete mode 100644 devolutions-gateway/src/tasks/ai_log.rs create mode 100644 devolutions-gateway/src/tasks/ai_log/checkpoint.rs create mode 100644 devolutions-gateway/src/tasks/ai_log/mod.rs create mode 100644 devolutions-gateway/src/tasks/ai_log/slog.rs create mode 100644 devolutions-gateway/src/tasks/ai_log/transcript.rs diff --git a/crates/terminal-streamer/src/trp_decoder.rs b/crates/terminal-streamer/src/trp_decoder.rs index 861e86186..e219afb33 100644 --- a/crates/terminal-streamer/src/trp_decoder.rs +++ b/crates/terminal-streamer/src/trp_decoder.rs @@ -80,170 +80,185 @@ impl AsyncRead for AsyncReadChannel { } } -/// Decodes a complete TRP recording into asciicast v2 lines, each ending with a newline. -/// -/// A truncated last packet, as left by an interrupted recording, is ignored. -pub fn decode_to_asciicast(mut input: &[u8]) -> anyhow::Result { - let mut decoder = TrpDecoder::default(); - let mut output = String::new(); - - while input.len() >= PACKET_HEADER_SIZE { - let (header, rest) = input.split_at(PACKET_HEADER_SIZE); - let header = PacketHeader::parse(header.try_into()?); - - let Some((payload, rest)) = rest.split_at_checked(usize::from(header.size)) else { - break; - }; - input = rest; - - for line in decoder.push(&header, payload)? { - output.push_str(&line); - output.push('\n'); - } - } - - for line in decoder.finish() { - output.push_str(&line); - output.push('\n'); - } - - Ok(output) -} - -const PACKET_HEADER_SIZE: usize = 8; - -struct PacketHeader { - time_delta: u32, - event_type: u16, - size: u16, -} +async fn parse_trp_stream( + mut input_stream: impl AsyncRead + Unpin + Send + 'static, + mut tx: tokio::sync::mpsc::Sender>, +) -> anyhow::Result<()> { + let mut time = 0.0; + let mut before_setup_cache = Some(Vec::new()); + let mut header = AsciinemaHeader::default(); -impl PacketHeader { - fn parse(buffer: &[u8; PACKET_HEADER_SIZE]) -> Self { - Self { - time_delta: u32::from_le_bytes([buffer[0], buffer[1], buffer[2], buffer[3]]), - event_type: u16::from_le_bytes([buffer[4], buffer[5]]), - size: u16::from_le_bytes([buffer[6], buffer[7]]), + loop { + let mut packet_head_buffer = [0u8; 8]; + if let Err(e) = input_stream.read_exact(&mut packet_head_buffer).await { + if e.kind() == std::io::ErrorKind::UnexpectedEof { + tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; + continue; + } + anyhow::bail!(e); } - } -} -/// Turns TRP packets into asciicast lines, holding the events that come before the terminal setup. -struct TrpDecoder { - time: f64, - before_setup_cache: Option>, - header: AsciinemaHeader, -} - -impl Default for TrpDecoder { - fn default() -> Self { - Self { - time: 0.0, - before_setup_cache: Some(Vec::new()), - header: AsciinemaHeader::default(), - } - } -} + let time_delta = u32::from_le_bytes(packet_head_buffer[0..4].try_into()?); + let event_type = u16::from_le_bytes(packet_head_buffer[4..6].try_into()?); + let size = u16::from_le_bytes(packet_head_buffer[6..8].try_into()?); -impl TrpDecoder { - fn push(&mut self, packet: &PacketHeader, payload: &[u8]) -> anyhow::Result> { - self.time += f64::from(packet.time_delta) / 1000.0; - let time = self.time; + time += f64::from(time_delta) / 1000.0; - let mut lines = Vec::new(); + let mut event_payload = vec![0u8; size as usize]; + if let Err(e) = input_stream.read_exact(&mut event_payload).await { + if e.kind() == std::io::ErrorKind::UnexpectedEof { + tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; + continue; + } + anyhow::bail!(e); + } - match packet.event_type { - 0 | 1 => { - let payload = String::from_utf8_lossy(payload).into_owned(); - let event = if packet.event_type == 0 { - AsciinemaEvent::TerminalOutput { payload, time } - } else { - AsciinemaEvent::UserInput { payload, time } + match event_type { + 0 => { + // Terminal output + let event_payload = String::from_utf8_lossy(&event_payload).into_owned(); + let event = AsciinemaEvent::TerminalOutput { + payload: event_payload, + time, + }; + match before_setup_cache { + Some(ref mut cache) => { + cache.push(event); + } + None => { + send(&mut tx, event.to_json()).await?; + } + } + } + 1 => { + let event_payload = String::from_utf8_lossy(&event_payload).into_owned(); + let event = AsciinemaEvent::UserInput { + payload: event_payload, + time, }; - match self.before_setup_cache { - Some(ref mut cache) => cache.push(event), - None => lines.push(event.to_json()), + match before_setup_cache { + Some(ref mut cache) => { + cache.push(event); + } + None => { + send(&mut tx, event.to_json()).await?; + } } } 2 => { // Terminal size change. Payload is little-endian [columns, rows]. - if payload.len() < 4 { - anyhow::bail!("invalid terminal size change payload length (len={})", payload.len()); + if event_payload.len() < 4 { + anyhow::bail!( + "invalid terminal size change payload length (len={})", + event_payload.len() + ); } - self.header.col = u16::from_le_bytes([payload[0], payload[1]]); - self.header.row = u16::from_le_bytes([payload[2], payload[3]]); - if self.before_setup_cache.is_none() { + header.col = u16::from_le_bytes(event_payload[0..2].try_into()?); + header.row = u16::from_le_bytes(event_payload[2..4].try_into()?); + if before_setup_cache.is_none() { let event = AsciinemaEvent::Resize { - width: self.header.col, - height: self.header.row, + width: header.col, + height: header.row, time, }; - lines.push(event.to_json()); + send(&mut tx, event.to_json()).await?; } } 4 => { // Terminal setup - if let Some(cache) = self.before_setup_cache.take() { - lines.push(self.header.to_json()); - lines.extend(cache.iter().map(AsciinemaEvent::to_json)); + if before_setup_cache.is_some() { + let header_json = header.to_json(); + send(&mut tx, header_json).await?; + if let Some(ref mut cache) = before_setup_cache { + for event in cache.drain(..) { + send(&mut tx, event.to_json()).await?; + } + } + before_setup_cache = None; } else { warn!("Received terminal setup event but cache is empty"); } } _ => {} } - - Ok(lines) } +} - /// Flushes the cached events of a recording that never sent its terminal setup. - fn finish(self) -> Vec { - match self.before_setup_cache { - Some(cache) => core::iter::once(self.header.to_json()) - .chain(cache.iter().map(AsciinemaEvent::to_json)) - .collect(), - None => Vec::new(), - } - } +async fn send(sender: &mut tokio::sync::mpsc::Sender>, mut json: String) -> anyhow::Result<()> { + json.push('\n'); + sender.send(Ok(json)).await?; + Ok(()) } -async fn parse_trp_stream( - mut input_stream: impl AsyncRead + Unpin + Send + 'static, - mut tx: tokio::sync::mpsc::Sender>, -) -> anyhow::Result<()> { - let mut decoder = TrpDecoder::default(); +/// Terminal output written at `time`, in seconds since the start of the recording. +#[derive(Debug, Clone, PartialEq)] +pub struct TerminalOutput { + pub time: f64, + pub text: String, +} - loop { - let mut packet_head_buffer = [0u8; PACKET_HEADER_SIZE]; - if let Err(e) = input_stream.read_exact(&mut packet_head_buffer).await { - if e.kind() == std::io::ErrorKind::UnexpectedEof { - tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; - continue; - } - anyhow::bail!(e); - } +/// Reads the terminal output of a finished TRP recording one packet at a time. +/// +/// A truncated last packet, as left by an interrupted recording, ends the output. +pub struct TrpOutputReader { + reader: R, + time: f64, +} - let packet = PacketHeader::parse(&packet_head_buffer); +impl TrpOutputReader { + pub fn new(reader: R) -> Self { + Self { reader, time: 0.0 } + } - let mut event_payload = vec![0u8; usize::from(packet.size)]; - if let Err(e) = input_stream.read_exact(&mut event_payload).await { - if e.kind() == std::io::ErrorKind::UnexpectedEof { - tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; - continue; - } - anyhow::bail!(e); + fn read_packet(&mut self) -> std::io::Result)>> { + let mut header = [0u8; 8]; + + if let Err(error) = self.reader.read_exact(&mut header) { + return eof_as_end(error); } - for line in decoder.push(&packet, &event_payload)? { - send(&mut tx, line).await?; + let time_delta = u32::from_le_bytes([header[0], header[1], header[2], header[3]]); + let event_type = u16::from_le_bytes([header[4], header[5]]); + let size = u16::from_le_bytes([header[6], header[7]]); + + let mut payload = vec![0u8; usize::from(size)]; + + if let Err(error) = self.reader.read_exact(&mut payload) { + return eof_as_end(error); } + + self.time += f64::from(time_delta) / 1000.0; + + Ok(Some((event_type, payload))) } } -async fn send(sender: &mut tokio::sync::mpsc::Sender>, mut json: String) -> anyhow::Result<()> { - json.push('\n'); - sender.send(Ok(json)).await?; - Ok(()) +fn eof_as_end(error: std::io::Error) -> std::io::Result> { + if error.kind() == std::io::ErrorKind::UnexpectedEof { + Ok(None) + } else { + Err(error) + } +} + +impl Iterator for TrpOutputReader { + type Item = std::io::Result; + + fn next(&mut self) -> Option { + loop { + match self.read_packet() { + Ok(Some((0, payload))) => { + return Some(Ok(TerminalOutput { + time: self.time, + text: String::from_utf8_lossy(&payload).into_owned(), + })); + } + Ok(Some(_)) => {} + Ok(None) => return None, + Err(error) => return Some(Err(error)), + } + } + } } #[cfg(test)] @@ -260,7 +275,7 @@ mod tests { } #[test] - fn decodes_a_complete_recording_and_ignores_a_truncated_tail() { + fn reads_the_timed_output_and_ignores_a_truncated_tail() { let mut trp = Vec::new(); trp.extend(packet(0, 2, &[100, 0, 30, 0])); trp.extend(packet(500, 0, b"$ ")); @@ -269,27 +284,49 @@ mod tests { trp.extend(packet(250, 0, b"ls\r\n")); trp.extend(&packet(10, 0, b"lost")[..9]); - let cast = decode_to_asciicast(&trp).expect("valid recording"); + let output = TrpOutputReader::new(trp.as_slice()) + .collect::>>() + .expect("valid recording"); assert_eq!( - cast, - concat!( - "{\"version\": 2, \"width\": 100, \"height\": 30}\n", - "[0.5,\"o\",\"$ \"]\n", - "[1.5,\"i\",\"l\"]\n", - r#"[1.75,"o","ls\u000d\u000a"]"#, - "\n", - ) + output, + [ + TerminalOutput { + time: 0.5, + text: "$ ".to_owned() + }, + TerminalOutput { + time: 1.75, + text: "ls\r\n".to_owned() + }, + ] ); } #[test] - fn events_of_a_recording_without_setup_are_kept() { - let cast = decode_to_asciicast(&packet(2000, 0, b"x")).expect("valid recording"); + fn reads_one_packet_at_a_time() { + struct CountingReader<'a> { + data: &'a [u8], + read: std::rc::Rc>, + } - assert_eq!( - cast, - "{\"version\": 2, \"width\": 80, \"height\": 24}\n[2,\"o\",\"x\"]\n" - ); + impl std::io::Read for CountingReader<'_> { + fn read(&mut self, buf: &mut [u8]) -> std::io::Result { + let n = std::io::Read::read(&mut self.data, buf)?; + self.read.set(self.read.get() + n); + Ok(n) + } + } + + let trp = (0..1000).flat_map(|_| packet(1, 0, b"x")).collect::>(); + let read = std::rc::Rc::new(std::cell::Cell::new(0)); + let mut reader = TrpOutputReader::new(CountingReader { + data: &trp, + read: std::rc::Rc::clone(&read), + }); + + reader.next().expect("one packet").expect("valid packet"); + + assert_eq!(read.get(), 9); } } diff --git a/devolutions-gateway/openapi/doc/index.adoc b/devolutions-gateway/openapi/doc/index.adoc index 8e24cb2d0..3fa4b20d3 100644 --- a/devolutions-gateway/openapi/doc/index.adoc +++ b/devolutions-gateway/openapi/doc/index.adoc @@ -3836,7 +3836,21 @@ Progress of a running `ai-log` task. | | <> | -| _Enum:_ preparing, +| _Enum:_ describing, + +| done +| X +| +| Integer +| +| + +| total +| X +| +| Integer +| +| |=== @@ -3864,6 +3878,42 @@ Progress of a running `ai-log` task. +[#AiLogSubstateOneOf1] +=== _AiLogSubstateOneOf1_ + +Transcript chunks sent to the AI provider so far, out of `total`. + + +[.fields-AiLogSubstateOneOf1] +[cols="2,1,1,2,4,1"] +|=== +| Field Name| Required| Nullable | Type| Description | Format + +| done +| X +| +| Integer +| +| + +| step +| X +| +| <> +| +| _Enum:_ describing, + +| total +| X +| +| Integer +| +| + +|=== + + + [#AiProvider] === _AiProvider_ diff --git a/devolutions-gateway/openapi/dotnet-client/.openapi-generator/FILES b/devolutions-gateway/openapi/dotnet-client/.openapi-generator/FILES index bf6f2ac19..430452ecc 100644 --- a/devolutions-gateway/openapi/dotnet-client/.openapi-generator/FILES +++ b/devolutions-gateway/openapi/dotnet-client/.openapi-generator/FILES @@ -12,6 +12,7 @@ docs/AgentStatus.md docs/AiLogParams.md docs/AiLogSubstate.md docs/AiLogSubstateOneOf.md +docs/AiLogSubstateOneOf1.md docs/AiProvider.md docs/AppCredential.md docs/AppCredentialKind.md @@ -130,6 +131,7 @@ src/Devolutions.Gateway.Client/Model/AgentStatus.cs src/Devolutions.Gateway.Client/Model/AiLogParams.cs src/Devolutions.Gateway.Client/Model/AiLogSubstate.cs src/Devolutions.Gateway.Client/Model/AiLogSubstateOneOf.cs +src/Devolutions.Gateway.Client/Model/AiLogSubstateOneOf1.cs src/Devolutions.Gateway.Client/Model/AiProvider.cs src/Devolutions.Gateway.Client/Model/AppCredential.cs src/Devolutions.Gateway.Client/Model/AppCredentialKind.cs diff --git a/devolutions-gateway/openapi/dotnet-client/README.md b/devolutions-gateway/openapi/dotnet-client/README.md index f90a874e1..92dbfa7dd 100644 --- a/devolutions-gateway/openapi/dotnet-client/README.md +++ b/devolutions-gateway/openapi/dotnet-client/README.md @@ -191,6 +191,7 @@ Class | Method | HTTP request | Description - [Model.AiLogParams](docs/AiLogParams.md) - [Model.AiLogSubstate](docs/AiLogSubstate.md) - [Model.AiLogSubstateOneOf](docs/AiLogSubstateOneOf.md) + - [Model.AiLogSubstateOneOf1](docs/AiLogSubstateOneOf1.md) - [Model.AiProvider](docs/AiProvider.md) - [Model.AppCredential](docs/AppCredential.md) - [Model.AppCredentialKind](docs/AppCredentialKind.md) diff --git a/devolutions-gateway/openapi/dotnet-client/docs/AiLogSubstate.md b/devolutions-gateway/openapi/dotnet-client/docs/AiLogSubstate.md index ec54cc1b0..8d5daff71 100644 --- a/devolutions-gateway/openapi/dotnet-client/docs/AiLogSubstate.md +++ b/devolutions-gateway/openapi/dotnet-client/docs/AiLogSubstate.md @@ -6,6 +6,8 @@ Progress of a running `ai-log` task. Name | Type | Description | Notes ------------ | ------------- | ------------- | ------------- **Step** | **string** | | +**Done** | **int** | | +**Total** | **int** | | [[Back to Model list]](../README.md#documentation-for-models) [[Back to API list]](../README.md#documentation-for-api-endpoints) [[Back to README]](../README.md) diff --git a/devolutions-gateway/openapi/dotnet-client/docs/AiLogSubstateOneOf1.md b/devolutions-gateway/openapi/dotnet-client/docs/AiLogSubstateOneOf1.md new file mode 100644 index 000000000..13b367220 --- /dev/null +++ b/devolutions-gateway/openapi/dotnet-client/docs/AiLogSubstateOneOf1.md @@ -0,0 +1,13 @@ +# Devolutions.Gateway.Client.Model.AiLogSubstateOneOf1 +Transcript chunks sent to the AI provider so far, out of `total`. + +## Properties + +Name | Type | Description | Notes +------------ | ------------- | ------------- | ------------- +**Done** | **int** | | +**Step** | **string** | | +**Total** | **int** | | + +[[Back to Model list]](../README.md#documentation-for-models) [[Back to API list]](../README.md#documentation-for-api-endpoints) [[Back to README]](../README.md) + diff --git a/devolutions-gateway/openapi/dotnet-client/src/Devolutions.Gateway.Client/Model/AiLogSubstate.cs b/devolutions-gateway/openapi/dotnet-client/src/Devolutions.Gateway.Client/Model/AiLogSubstate.cs index dd005cf5b..9c3ad74c0 100644 --- a/devolutions-gateway/openapi/dotnet-client/src/Devolutions.Gateway.Client/Model/AiLogSubstate.cs +++ b/devolutions-gateway/openapi/dotnet-client/src/Devolutions.Gateway.Client/Model/AiLogSubstate.cs @@ -21,6 +21,7 @@ using Newtonsoft.Json; using Newtonsoft.Json.Converters; using Newtonsoft.Json.Linq; +using JsonSubTypes; using System.ComponentModel.DataAnnotations; using FileParameter = Devolutions.Gateway.Client.Client.FileParameter; using OpenAPIDateConverter = Devolutions.Gateway.Client.Client.OpenAPIDateConverter; @@ -47,6 +48,18 @@ public AiLogSubstate(AiLogSubstateOneOf actualInstance) this.ActualInstance = actualInstance ?? throw new ArgumentException("Invalid instance found. Must not be null."); } + /// + /// Initializes a new instance of the class + /// with the class + /// + /// An instance of AiLogSubstateOneOf1. + public AiLogSubstate(AiLogSubstateOneOf1 actualInstance) + { + this.IsNullable = false; + this.SchemaType= "oneOf"; + this.ActualInstance = actualInstance ?? throw new ArgumentException("Invalid instance found. Must not be null."); + } + private Object _actualInstance; @@ -65,9 +78,13 @@ public override Object ActualInstance { this._actualInstance = value; } + else if (value.GetType() == typeof(AiLogSubstateOneOf1) || value is AiLogSubstateOneOf1) + { + this._actualInstance = value; + } else { - throw new ArgumentException("Invalid instance found. Must be the following types: AiLogSubstateOneOf"); + throw new ArgumentException("Invalid instance found. Must be the following types: AiLogSubstateOneOf, AiLogSubstateOneOf1"); } } } @@ -82,6 +99,16 @@ public AiLogSubstateOneOf GetAiLogSubstateOneOf() return (AiLogSubstateOneOf)this.ActualInstance; } + /// + /// Get the actual instance of `AiLogSubstateOneOf1`. If the actual instance is not `AiLogSubstateOneOf1`, + /// the InvalidClassException will be thrown + /// + /// An instance of AiLogSubstateOneOf1 + public AiLogSubstateOneOf1 GetAiLogSubstateOneOf1() + { + return (AiLogSubstateOneOf1)this.ActualInstance; + } + /// /// Returns the string presentation of the object /// @@ -140,6 +167,26 @@ public static AiLogSubstate FromJson(string jsonString) System.Diagnostics.Debug.WriteLine(string.Format("Failed to deserialize `{0}` into AiLogSubstateOneOf: {1}", jsonString, exception.ToString())); } + try + { + // if it does not contains "AdditionalProperties", use SerializerSettings to deserialize + if (typeof(AiLogSubstateOneOf1).GetProperty("AdditionalProperties") == null) + { + newAiLogSubstate = new AiLogSubstate(JsonConvert.DeserializeObject(jsonString, AiLogSubstate.SerializerSettings)); + } + else + { + newAiLogSubstate = new AiLogSubstate(JsonConvert.DeserializeObject(jsonString, AiLogSubstate.AdditionalPropertiesSerializerSettings)); + } + matchedTypes.Add("AiLogSubstateOneOf1"); + match++; + } + catch (Exception exception) + { + // deserialization failed, try the next one + System.Diagnostics.Debug.WriteLine(string.Format("Failed to deserialize `{0}` into AiLogSubstateOneOf1: {1}", jsonString, exception.ToString())); + } + if (match == 0) { throw new InvalidDataException("The JSON string `" + jsonString + "` cannot be deserialized into any schema defined."); diff --git a/devolutions-gateway/openapi/dotnet-client/src/Devolutions.Gateway.Client/Model/AiLogSubstateOneOf1.cs b/devolutions-gateway/openapi/dotnet-client/src/Devolutions.Gateway.Client/Model/AiLogSubstateOneOf1.cs new file mode 100644 index 000000000..5787cbc1a --- /dev/null +++ b/devolutions-gateway/openapi/dotnet-client/src/Devolutions.Gateway.Client/Model/AiLogSubstateOneOf1.cs @@ -0,0 +1,132 @@ +/* + * devolutions-gateway + * + * Protocol-aware fine-grained relay server + * + * The version of the OpenAPI document: 2026.2.4 + * Contact: infos@devolutions.net + * Generated by: https://github.com/openapitools/openapi-generator.git + */ + + +using System; +using System.Collections; +using System.Collections.Generic; +using System.Collections.ObjectModel; +using System.Linq; +using System.IO; +using System.Runtime.Serialization; +using System.Text; +using System.Text.RegularExpressions; +using Newtonsoft.Json; +using Newtonsoft.Json.Converters; +using Newtonsoft.Json.Linq; +using System.ComponentModel.DataAnnotations; +using FileParameter = Devolutions.Gateway.Client.Client.FileParameter; +using OpenAPIDateConverter = Devolutions.Gateway.Client.Client.OpenAPIDateConverter; + +namespace Devolutions.Gateway.Client.Model +{ + /// + /// Transcript chunks sent to the AI provider so far, out of `total`. + /// + [DataContract(Name = "AiLogSubstate_oneOf_1")] + public partial class AiLogSubstateOneOf1 : IValidatableObject + { + /// + /// Defines Step + /// + [JsonConverter(typeof(StringEnumConverter))] + public enum StepEnum + { + /// + /// Enum Describing for value: describing + /// + [EnumMember(Value = "describing")] + Describing = 1 + } + + + /// + /// Gets or Sets Step + /// + [DataMember(Name = "step", IsRequired = true, EmitDefaultValue = true)] + public StepEnum Step { get; set; } + /// + /// Initializes a new instance of the class. + /// + [JsonConstructorAttribute] + protected AiLogSubstateOneOf1() { } + /// + /// Initializes a new instance of the class. + /// + /// done (required). + /// step (required). + /// total (required). + public AiLogSubstateOneOf1(int done = default(int), StepEnum step = default(StepEnum), int total = default(int)) + { + this.Done = done; + this.Step = step; + this.Total = total; + } + + /// + /// Gets or Sets Done + /// + [DataMember(Name = "done", IsRequired = true, EmitDefaultValue = true)] + public int Done { get; set; } + + /// + /// Gets or Sets Total + /// + [DataMember(Name = "total", IsRequired = true, EmitDefaultValue = true)] + public int Total { get; set; } + + /// + /// Returns the string presentation of the object + /// + /// String presentation of the object + public override string ToString() + { + StringBuilder sb = new StringBuilder(); + sb.Append("class AiLogSubstateOneOf1 {\n"); + sb.Append(" Done: ").Append(Done).Append("\n"); + sb.Append(" Step: ").Append(Step).Append("\n"); + sb.Append(" Total: ").Append(Total).Append("\n"); + sb.Append("}\n"); + return sb.ToString(); + } + + /// + /// Returns the JSON string presentation of the object + /// + /// JSON string presentation of the object + public virtual string ToJson() + { + return Newtonsoft.Json.JsonConvert.SerializeObject(this, Newtonsoft.Json.Formatting.Indented); + } + + /// + /// To validate all properties of the instance + /// + /// Validation context + /// Validation Result + IEnumerable IValidatableObject.Validate(ValidationContext validationContext) + { + // Done (int) minimum + if (this.Done < (int)0) + { + yield return new ValidationResult("Invalid value for Done, must be a value greater than or equal to 0.", new [] { "Done" }); + } + + // Total (int) minimum + if (this.Total < (int)0) + { + yield return new ValidationResult("Invalid value for Total, must be a value greater than or equal to 0.", new [] { "Total" }); + } + + yield break; + } + } + +} diff --git a/devolutions-gateway/openapi/gateway-api.yaml b/devolutions-gateway/openapi/gateway-api.yaml index 723c99bd6..7abe2f499 100644 --- a/devolutions-gateway/openapi/gateway-api.yaml +++ b/devolutions-gateway/openapi/gateway-api.yaml @@ -1535,7 +1535,26 @@ components: type: string enum: - preparing + - type: object + description: Transcript chunks sent to the AI provider so far, out of `total`. + required: + - done + - total + - step + properties: + done: + type: integer + minimum: 0 + step: + type: string + enum: + - describing + total: + type: integer + minimum: 0 description: Progress of a running `ai-log` task. + discriminator: + propertyName: step AiProvider: type: string enum: diff --git a/devolutions-gateway/openapi/ts-angular-client/.openapi-generator/FILES b/devolutions-gateway/openapi/ts-angular-client/.openapi-generator/FILES index beb588571..1be1ff8f3 100644 --- a/devolutions-gateway/openapi/ts-angular-client/.openapi-generator/FILES +++ b/devolutions-gateway/openapi/ts-angular-client/.openapi-generator/FILES @@ -30,6 +30,7 @@ model/agentStatus.ts model/aiLogParams.ts model/aiLogSubstate.ts model/aiLogSubstateOneOf.ts +model/aiLogSubstateOneOf1.ts model/aiProvider.ts model/appCredential.ts model/appCredentialKind.ts diff --git a/devolutions-gateway/openapi/ts-angular-client/model/aiLogSubstate.ts b/devolutions-gateway/openapi/ts-angular-client/model/aiLogSubstate.ts index 9e4f39ec7..13468f25b 100644 --- a/devolutions-gateway/openapi/ts-angular-client/model/aiLogSubstate.ts +++ b/devolutions-gateway/openapi/ts-angular-client/model/aiLogSubstate.ts @@ -7,6 +7,7 @@ * https://openapi-generator.tech * Do not edit the class manually. */ +import { AiLogSubstateOneOf1 } from './aiLogSubstateOneOf1'; import { AiLogSubstateOneOf } from './aiLogSubstateOneOf'; @@ -18,5 +19,5 @@ import { AiLogSubstateOneOf } from './aiLogSubstateOneOf'; * Progress of a running `ai-log` task. * @export */ -export type AiLogSubstate = AiLogSubstateOneOf; +export type AiLogSubstate = AiLogSubstateOneOf | AiLogSubstateOneOf1; diff --git a/devolutions-gateway/openapi/ts-angular-client/model/aiLogSubstateOneOf1.ts b/devolutions-gateway/openapi/ts-angular-client/model/aiLogSubstateOneOf1.ts new file mode 100644 index 000000000..6de6abae5 --- /dev/null +++ b/devolutions-gateway/openapi/ts-angular-client/model/aiLogSubstateOneOf1.ts @@ -0,0 +1,27 @@ +/** + * devolutions-gateway + * + * Contact: infos@devolutions.net + * + * NOTE: This class is auto generated by OpenAPI Generator (https://openapi-generator.tech). + * https://openapi-generator.tech + * Do not edit the class manually. + */ + + +/** + * Transcript chunks sent to the AI provider so far, out of `total`. + */ +export interface AiLogSubstateOneOf1 { + done: number; + step: AiLogSubstateOneOf1.Step; + total: number; +} +export namespace AiLogSubstateOneOf1 { + export type Step = 'describing'; + export const Step = { + Describing: 'describing' as Step + }; +} + + diff --git a/devolutions-gateway/openapi/ts-angular-client/model/models.ts b/devolutions-gateway/openapi/ts-angular-client/model/models.ts index 2a42acc64..3dfc6c2bb 100644 --- a/devolutions-gateway/openapi/ts-angular-client/model/models.ts +++ b/devolutions-gateway/openapi/ts-angular-client/model/models.ts @@ -8,6 +8,7 @@ export * from './agentStatus'; export * from './aiLogParams'; export * from './aiLogSubstate'; export * from './aiLogSubstateOneOf'; +export * from './aiLogSubstateOneOf1'; export * from './aiProvider'; export * from './appCredential'; export * from './appCredentialKind'; diff --git a/devolutions-gateway/src/recording.rs b/devolutions-gateway/src/recording.rs index f57061dfa..7ce14d600 100644 --- a/devolutions-gateway/src/recording.rs +++ b/devolutions-gateway/src/recording.rs @@ -30,15 +30,26 @@ const BUFFER_WRITER_SIZE: usize = 64 * 1024; #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] -struct JrecFile { +pub(crate) struct JrecFile { file_name: String, start_time: i64, duration: i64, } +impl JrecFile { + pub(crate) fn file_name(&self) -> &str { + &self.file_name + } + + /// Unix seconds. + pub(crate) fn start_time(&self) -> i64 { + self.start_time + } +} + #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] -struct JrecManifest { +pub(crate) struct JrecManifest { session_id: Uuid, start_time: i64, duration: i64, @@ -48,6 +59,20 @@ struct JrecManifest { } impl JrecManifest { + /// Unix seconds. + pub(crate) fn start_time(&self) -> i64 { + self.start_time + } + + /// Seconds. + pub(crate) fn duration(&self) -> i64 { + self.duration + } + + pub(crate) fn files(&self) -> &[JrecFile] { + &self.files + } + fn read_from_file(path: impl AsRef) -> anyhow::Result { let json = std::fs::read(path)?; let manifest = serde_json::from_slice(&json)?; @@ -257,6 +282,10 @@ enum RecordingManagerMessage { kind: ArtifactKind, channel: oneshot::Sender>, }, + GetFinished { + id: Uuid, + channel: oneshot::Sender>, + }, Disconnect { id: Uuid, }, @@ -281,6 +310,33 @@ enum RecordingManagerMessage { }, } +/// A session that is not recording anymore and has at least one Recording. +#[derive(Debug)] +pub(crate) struct FinishedRecording { + pub(crate) dir: Utf8PathBuf, + pub(crate) manifest: JrecManifest, +} + +#[cfg(test)] +impl FinishedRecording { + pub(crate) fn read_for_test(dir: &camino::Utf8Path) -> Self { + Self { + dir: dir.to_owned(), + manifest: JrecManifest::read_from_file(dir.join("recording.json")).expect("valid manifest"), + } + } +} + +#[derive(Debug, thiserror::Error)] +pub(crate) enum FinishedRecordingError { + #[error("session has no recording")] + NotFound, + #[error("session is still recording")] + Recording, + #[error(transparent)] + Io(anyhow::Error), +} + impl fmt::Debug for RecordingManagerMessage { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { match self { @@ -300,6 +356,9 @@ impl fmt::Debug for RecordingManagerMessage { .field("id", id) .field("kind", kind) .finish_non_exhaustive(), + RecordingManagerMessage::GetFinished { id, channel: _ } => { + f.debug_struct("GetFinished").field("id", id).finish_non_exhaustive() + } RecordingManagerMessage::Disconnect { id } => f.debug_struct("Disconnect").field("id", id).finish(), RecordingManagerMessage::GetState { id, channel: _ } => { f.debug_struct("GetState").field("id", id).finish_non_exhaustive() @@ -365,6 +424,19 @@ impl RecordingMessageSender { rx.await.context("couldn't receive AddArtifact result")? } + pub(crate) async fn get_finished( + &self, + id: Uuid, + ) -> anyhow::Result> { + let (tx, rx) = oneshot::channel(); + self.channel + .send(RecordingManagerMessage::GetFinished { id, channel: tx }) + .await + .ok() + .context("couldn't send GetFinished message")?; + rx.await.context("couldn't receive the finished recording") + } + async fn disconnect(&self, id: Uuid) -> anyhow::Result<()> { self.channel .send(RecordingManagerMessage::Disconnect { id }) @@ -727,6 +799,26 @@ impl RecordingManagerTask { Ok(artifact_path) } + fn get_finished(&self, id: Uuid) -> Result { + // The transcript is built from the whole recording, so one still being pushed is not ready. + if self.ongoing_recordings.contains_key(&id) { + return Err(FinishedRecordingError::Recording); + } + + let dir = self.recordings_path.join(id.to_string()); + let manifest_path = dir.join("recording.json"); + + if !manifest_path.exists() { + return Err(FinishedRecordingError::NotFound); + } + + let manifest = JrecManifest::read_from_file(&manifest_path) + .context("read manifest from disk") + .map_err(FinishedRecordingError::Io)?; + + Ok(FinishedRecording { dir, manifest }) + } + fn handle_remove(&mut self, id: Uuid) { if let Some(ongoing) = self.ongoing_recordings.get(&id) { let now = time::OffsetDateTime::now_utc().unix_timestamp(); @@ -857,6 +949,9 @@ async fn recording_manager_task( RecordingManagerMessage::AddArtifact { id, kind, channel } => { let _ = channel.send(manager.handle_add_artifact(id, kind).await); }, + RecordingManagerMessage::GetFinished { id, channel } => { + let _ = channel.send(manager.get_finished(id)); + } RecordingManagerMessage::Disconnect { id } => { if let Err(e) = manager.handle_disconnect(id).await { error!(error = format!("{e:#}"), "handle_disconnect"); @@ -1253,6 +1348,25 @@ mod tests { assert_eq!(harness.sender.get_count().await.expect("count"), 1); } + #[tokio::test] + async fn finished_recordings_are_read_through_the_manager() { + let harness = Harness::start(); + let id = Uuid::new_v4(); + harness.write_manifest(id, MASTER_MANIFEST); + + let finished = harness.sender.get_finished(id).await.expect("sent").expect("finished"); + assert_eq!(finished.dir, harness.recordings_path.join(id.to_string())); + assert_eq!(finished.manifest.files()[0].file_name(), "recording-0.webm"); + + let ongoing = Uuid::new_v4(); + harness.connect(ongoing, WEBM).await; + let result = harness.sender.get_finished(ongoing).await.expect("sent"); + assert!(matches!(result, Err(FinishedRecordingError::Recording)), "{result:?}"); + + let missing = harness.sender.get_finished(Uuid::new_v4()).await.expect("sent"); + assert!(matches!(missing, Err(FinishedRecordingError::NotFound)), "{missing:?}"); + } + #[test] fn manifest_without_artifacts_round_trips_byte_for_byte() { let dir = tempfile::tempdir().expect("temp dir"); diff --git a/devolutions-gateway/src/tasks/ai_log.rs b/devolutions-gateway/src/tasks/ai_log.rs deleted file mode 100644 index 58e4a80af..000000000 --- a/devolutions-gateway/src/tasks/ai_log.rs +++ /dev/null @@ -1,171 +0,0 @@ -//! `ai-log` task: describes what the user did in one session and stores the result as a new log of that session. - -use secrecy::SecretString; -use url::Url; -use uuid::Uuid; - -use super::ai::{AiProvider, AiSettings}; -use super::{EphemeralTask, RetryPolicy, SECRETS_LOST_ERROR, TaskCtx, TaskError, TaskErrorCode, TaskKind}; -use crate::DgwState; - -#[derive(Debug, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct AiLogTarget { - pub session_id: Uuid, -} - -/// AI settings used by an `ai-log` task: the body of `POST /jet/tasks` for a TASK token of kind `ai-log`. -#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))] -#[derive(Debug, Deserialize)] -#[serde(rename_all = "camelCase", deny_unknown_fields)] -pub struct AiLogParams { - pub provider: AiProvider, - /// Model identifier, passed to the provider as is. - pub model: String, - /// Kept in memory for this task only. - #[cfg_attr(feature = "openapi", schema(value_type = String))] - pub api_key: SecretString, - /// Overrides the provider default; required for `openai-compatible`. - #[cfg_attr(feature = "openapi", schema(value_type = Option))] - pub base_url: Option, - /// Upper bound of tokens in each AI answer. - pub max_output_tokens: Option, -} - -/// Progress of a running `ai-log` task. -#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))] -#[derive(Debug, Default, Serialize)] -#[serde(rename_all = "kebab-case", tag = "step")] -pub enum AiLogSubstate { - #[default] - Preparing, -} - -#[derive(Debug, Serialize)] -pub enum AiLogOutput {} - -pub enum AiLogTask {} - -impl TaskKind for AiLogTask { - const KIND: &'static str = "ai-log"; - const RETRY: RetryPolicy = RetryPolicy::JOB_QUEUE; - - type Target = AiLogTarget; - type Params = AiSettings; - type Substate = AiLogSubstate; - type Output = AiLogOutput; - - async fn run(ctx: TaskCtx) -> Result { - let Some(api_key) = ctx.secrets() else { - return Err(TaskError::Permanent(SECRETS_LOST_ERROR.to_owned())); - }; - - let _client = ctx.params.client(api_key, &ctx.state)?; - - Err(TaskError::Permanent("ai-log task not implemented yet".to_owned())) - } -} - -impl EphemeralTask for AiLogTask { - type Secrets = SecretString; - type Request = AiLogParams; - - fn prepare( - target: &AiLogTarget, - request: AiLogParams, - state: &DgwState, - ) -> Result<(AiSettings, SecretString), TaskErrorCode> { - if state.recordings.active_recordings.contains(target.session_id) { - return Err(TaskErrorCode::RecordingActive); - } - - let AiLogParams { - provider, - model, - api_key, - base_url, - max_output_tokens, - } = request; - - let settings = AiSettings { - provider, - model, - base_url, - max_output_tokens, - }; - - settings.check(&api_key, state)?; - - Ok((settings, api_key)) - } -} - -#[cfg(test)] -mod tests { - use super::*; - - const API_KEY: &str = "sk-ai-log-test-secret"; - - const CONFIG: &str = r#"{ - "ProvisionerPublicKeyData": { - "Value": "mMIIBIjANBgkqhkiG9w0BAQEFAAOCAQ8AMIIBCgKCAQEA4vuqLOkl1pWobt6su1XO9VskgCAwevEGs6kkNjJQBwkGnPKYLmNF1E/af1yCocfVn/OnPf9e4x+lXVyZ6LMDJxFxu+axdgOq3Ld392J1iAEbfvwlyRFnEXFOJNyylqg3bY6LvnWHL/XZczVdMD9xYfq2sO9bg3xjRW4s7r9EEYOFjqVT3VFznH9iWJVtcSEKukmS/3uKoO6lGhacvu0HhjXXdgq0R8zvR4XRJ9Fcnf0f9Ypoc+i6L80NVjrRCeVOH+Ld/2fA9bocpfLarcVqG3RjS+qgOtpyCc0jWVFF4zaGQ7LUDFkEIYILkICeMMn2ll29hmZNzsJzZJ9s6NocgQIDAQAB" - }, - "Listeners": [{ "InternalUrl": "http://*:7171", "ExternalUrl": "https://*:7171" }], - "Proxy": { "Mode": "Off" } - }"#; - - fn params() -> AiLogParams { - serde_json::from_value(serde_json::json!({ - "provider": "openai", - "model": "gpt-test", - "apiKey": API_KEY, - })) - .expect("valid params") - } - - fn target() -> AiLogTarget { - AiLogTarget { - session_id: Uuid::new_v4(), - } - } - - #[tokio::test] - async fn refuses_a_session_that_is_still_recording() { - let (state, _handles) = DgwState::mock(CONFIG).expect("mock state"); - let target = target(); - state.recordings.active_recordings.insert(target.session_id); - - let error = AiLogTask::prepare(&target, params(), &state).expect_err("session is busy"); - - assert_eq!(error, TaskErrorCode::RecordingActive); - } - - #[tokio::test] - async fn persisted_settings_never_hold_the_api_key() { - let (state, _handles) = DgwState::mock(CONFIG).expect("mock state"); - - let params = params(); - assert!(!format!("{params:?}").contains(API_KEY)); - - let (settings, api_key) = AiLogTask::prepare(&target(), params, &state).expect("valid task"); - - let persisted = serde_json::to_string(&settings).expect("serializable settings"); - assert_eq!( - persisted, - r#"{"provider":"openai","model":"gpt-test","baseUrl":null,"maxOutputTokens":null}"# - ); - assert!(!format!("{settings:?}").contains(API_KEY)); - assert!(!format!("{api_key:?}").contains(API_KEY)); - } - - #[tokio::test] - async fn invalid_ai_settings_are_refused_with_a_code() { - let (state, _handles) = DgwState::mock(CONFIG).expect("mock state"); - let mut params = params(); - params.model = " ".to_owned(); - - let error = AiLogTask::prepare(&target(), params, &state).expect_err("empty model"); - - assert_eq!(error, TaskErrorCode::MissingModel); - } -} diff --git a/devolutions-gateway/src/tasks/ai_log/checkpoint.rs b/devolutions-gateway/src/tasks/ai_log/checkpoint.rs new file mode 100644 index 000000000..af89bd00e --- /dev/null +++ b/devolutions-gateway/src/tasks/ai_log/checkpoint.rs @@ -0,0 +1,94 @@ +//! Actions found in one transcript chunk, saved in the task workspace so a retry does not ask the AI again. + +use std::collections::BTreeMap; +use std::io::{BufRead as _, BufReader, BufWriter, Write as _}; +use std::time::Duration; + +use anyhow::Context as _; +use camino::{Utf8Path, Utf8PathBuf}; +use devolutions_gateway_ai::session_actions::Action; + +pub(crate) fn path(workspace: &Utf8Path, index: usize) -> Utf8PathBuf { + workspace.join(format!("chunk-{index:04}.actions.jsonl")) +} + +#[derive(Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +struct SavedAction { + offset_seconds: f64, + description: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + object: Option, + #[serde(default, skip_serializing_if = "BTreeMap::is_empty")] + parameters: BTreeMap, +} + +/// Writes the checkpoint through a temporary file, so a crash never leaves a partial one. +pub(crate) fn write(path: &Utf8Path, actions: &[Action]) -> anyhow::Result<()> { + let partial = path.with_extension("partial"); + + let mut out = BufWriter::new(std::fs::File::create(&partial).with_context(|| format!("create {partial}"))?); + + for action in actions { + let saved = SavedAction { + offset_seconds: action.offset.as_secs_f64(), + description: action.description.clone(), + object: action.object.clone(), + parameters: action.parameters.clone(), + }; + serde_json::to_writer(&mut out, &saved)?; + out.write_all(b"\n")?; + } + + out.into_inner()?.sync_all()?; + std::fs::rename(&partial, path).with_context(|| format!("rename {partial}"))?; + + Ok(()) +} + +pub(crate) fn read(path: &Utf8Path) -> anyhow::Result> { + let file = BufReader::new(std::fs::File::open(path).with_context(|| format!("open {path}"))?); + + file.lines() + .map(|line| { + let saved: SavedAction = serde_json::from_str(&line?)?; + Ok(Action { + offset: Duration::try_from_secs_f64(saved.offset_seconds)?, + description: saved.description, + object: saved.object, + parameters: saved.parameters, + }) + }) + .collect() +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn actions_round_trip() { + let dir = tempfile::tempdir().expect("temp dir"); + let path = path(Utf8Path::from_path(dir.path()).expect("UTF-8"), 3); + let actions = vec![ + Action { + offset: Duration::from_millis(1500), + description: "Listed files".to_owned(), + object: Some("/var/log".to_owned()), + parameters: BTreeMap::from([("Command".to_owned(), "ls".to_owned())]), + }, + Action { + offset: Duration::from_secs(62), + description: "Closed the shell".to_owned(), + object: None, + parameters: BTreeMap::new(), + }, + ]; + + write(&path, &actions).expect("written"); + + assert!(path.as_str().ends_with("chunk-0003.actions.jsonl")); + assert_eq!(read(&path).expect("read"), actions); + assert!(!path.with_extension("partial").exists()); + } +} diff --git a/devolutions-gateway/src/tasks/ai_log/mod.rs b/devolutions-gateway/src/tasks/ai_log/mod.rs new file mode 100644 index 000000000..a545c57cb --- /dev/null +++ b/devolutions-gateway/src/tasks/ai_log/mod.rs @@ -0,0 +1,443 @@ +//! `ai-log` task: describes what the user did in one session and stores the result as a new log of that session. +//! +//! The task works in steps, keeping its files in the task workspace: +//! 1. stream the terminal recordings into transcript chunk files (`chunk-NNNN.txt`); +//! 2. ask the AI about each chunk and save its actions as a checkpoint (`chunk-NNNN.actions.jsonl`); +//! a retry skips the chunks that already have one; +//! 3. merge the checkpoints into a `.slog` file and add it to the session. + +mod checkpoint; +mod slog; +mod transcript; + +use std::collections::VecDeque; +use std::fs::File; +use std::io::BufWriter; + +use camino::{Utf8Path, Utf8PathBuf}; +use devolutions_gateway_ai::AiClient; +use devolutions_gateway_ai::session_actions::Action; +use secrecy::SecretString; +use url::Url; +use uuid::Uuid; + +use super::ai::{AiProvider, AiSettings}; +use super::{EphemeralTask, RetryPolicy, SECRETS_LOST_ERROR, TaskCtx, TaskError, TaskErrorCode, TaskKind}; +use crate::DgwState; +use crate::artifacts::ArtifactKind; +use crate::recording::{FinishedRecording, RecordingMessageSender}; + +/// Input tokens sent in one AI request, estimated at 4 characters per token. +pub const MAX_INPUT_TOKENS_PER_REQUEST: usize = 100_000; + +const CHARS_PER_TOKEN: usize = 4; + +/// Longest transcript chunk, in bytes; it bounds the memory used by the task. +const MAX_CHUNK_LEN: usize = MAX_INPUT_TOKENS_PER_REQUEST * CHARS_PER_TOKEN; + +/// A truncated AI answer for a transcript part shorter than this fails the task instead of splitting the part again. +const MIN_SPLIT_LEN: usize = 2_000; + +pub const TRUNCATED_ERROR: &str = + "AI answer was cut at the output token limit, even for a short part of the transcript"; + +/// Written once every chunk file is complete; holds the number of chunks. +const CHUNKS_DONE_FILE: &str = "chunks.done"; + +const LOG_FILE: &str = "log.slog"; + +#[derive(Debug, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct AiLogTarget { + pub session_id: Uuid, +} + +/// AI settings used by an `ai-log` task: the body of `POST /jet/tasks` for a TASK token of kind `ai-log`. +#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))] +#[derive(Debug, Deserialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +pub struct AiLogParams { + pub provider: AiProvider, + /// Model identifier, passed to the provider as is. + pub model: String, + /// Kept in memory for this task only. + #[cfg_attr(feature = "openapi", schema(value_type = String))] + pub api_key: SecretString, + /// Overrides the provider default; required for `openai-compatible`. + #[cfg_attr(feature = "openapi", schema(value_type = Option))] + pub base_url: Option, + /// Upper bound of tokens in each AI answer. + pub max_output_tokens: Option, +} + +/// Progress of a running `ai-log` task. +#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))] +#[derive(Debug, Default, Serialize)] +#[serde(rename_all = "kebab-case", tag = "step")] +pub enum AiLogSubstate { + #[default] + Preparing, + /// Transcript chunks sent to the AI provider so far, out of `total`. + Describing { done: usize, total: usize }, +} + +/// Result of a successful `ai-log` task. +#[derive(Debug, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct AiLogOutput { + /// Name of the new log in the session manifest, such as `ai-analysis-0.slog`. + pub file_name: String, +} + +pub enum AiLogTask {} + +impl TaskKind for AiLogTask { + const KIND: &'static str = "ai-log"; + const RETRY: RetryPolicy = RetryPolicy::JOB_QUEUE; + + type Target = AiLogTarget; + type Params = AiSettings; + type Substate = AiLogSubstate; + type Output = AiLogOutput; + + async fn run(ctx: TaskCtx) -> Result { + let Some(api_key) = ctx.secrets() else { + return Err(TaskError::Permanent(SECRETS_LOST_ERROR.to_owned())); + }; + + let client = ctx.params.client(api_key, &ctx.state)?; + + let session_id = ctx.target.session_id; + + let recording = match ctx.state.recordings.get_finished(session_id).await { + Ok(Ok(recording)) => recording, + Ok(Err(error)) => return Err(TaskError::Permanent(format!("{error:#}"))), + Err(error) => return Err(TaskError::Permanent(format!("failed to read the recording: {error:#}"))), + }; + + let start_time = recording.manifest.start_time(); + let duration = recording.manifest.duration(); + let workspace = ctx.workspace.clone(); + + let total = blocking(write_chunks(recording, workspace.clone())).await?; + + info!(session.id = %session_id, chunks = total, "Describe the session actions"); + + for index in 0..total { + let checkpoint = checkpoint::path(&workspace, index); + + if checkpoint.exists() { + debug!(session.id = %session_id, index, "Chunk already described"); + continue; + } + + ctx.progress + .set(&AiLogSubstate::Describing { done: index, total }) + .await; + + let chunk = tokio::fs::read_to_string(transcript::chunk_path(&workspace, index)) + .await + .map_err(|error| workspace_error(&error))?; + + let actions = describe_chunk(&client, ctx.params.max_output_tokens, &chunk).await?; + + blocking(move || checkpoint::write(&checkpoint, &actions).map_err(|error| workspace_error(&error))).await?; + } + + ctx.progress + .set(&AiLogSubstate::Describing { done: total, total }) + .await; + + let log_path = workspace.join(LOG_FILE); + let model = ctx.params.model.clone(); + let actions = blocking({ + let log_path = log_path.clone(); + move || merge_checkpoints(&workspace, total, start_time, duration, &model, &log_path) + }) + .await?; + + let file_name = add_log(&ctx.state.recordings, session_id, &log_path) + .await + .map_err(|error| TaskError::Permanent(format!("failed to add the log: {error:#}")))?; + + info!(session.id = %session_id, file_name, actions, "Session log generated"); + + Ok(AiLogOutput { file_name }) + } +} + +async fn add_log(recordings: &RecordingMessageSender, session_id: Uuid, log_path: &Utf8Path) -> anyhow::Result { + let artifact_path = recordings.add_artifact(session_id, ArtifactKind::AiAnalysis).await?; + tokio::fs::copy(log_path, &artifact_path).await?; + Ok(artifact_path + .file_name() + .expect("artifact paths end with a file name") + .to_owned()) +} + +async fn blocking( + work: impl FnOnce() -> Result + Send + 'static, +) -> Result { + tokio::task::spawn_blocking(work) + .await + .map_err(|_| TaskError::Permanent("task step panicked".to_owned()))? +} + +fn workspace_error(error: &dyn std::fmt::Display) -> TaskError { + TaskError::Permanent(format!("task workspace error: {error:#}")) +} + +/// Splits the transcript into chunk files once; a retry reuses them. +fn write_chunks(recording: FinishedRecording, workspace: Utf8PathBuf) -> impl FnOnce() -> Result { + move || { + let done = workspace.join(CHUNKS_DONE_FILE); + + if let Some(total) = std::fs::read_to_string(&done).ok().and_then(|total| total.parse().ok()) { + return Ok(total); + } + + std::fs::create_dir_all(&workspace).map_err(|error| workspace_error(&error))?; + + let total = transcript::write_chunks(&recording, &workspace, MAX_CHUNK_LEN).map_err(|error| match error { + transcript::TranscriptError::Write(error) => workspace_error(&error), + error => TaskError::Permanent(error.to_string()), + })?; + + std::fs::write(&done, total.to_string()).map_err(|error| workspace_error(&error))?; + + Ok(total) + } +} + +/// Asks the AI about one chunk; a truncated answer is asked again as two halves. +async fn describe_chunk( + client: &AiClient, + max_output_tokens: Option, + chunk: &str, +) -> Result, TaskError> { + let mut parts = VecDeque::from([chunk]); + let mut actions = Vec::new(); + + while let Some(part) = parts.pop_front() { + let mut request = client.describe_session_actions(part); + + if let Some(max_output_tokens) = max_output_tokens { + request = request.max_output_tokens(max_output_tokens); + } + + match request.send().await { + Ok(response) => actions.extend(response.output), + Err(devolutions_gateway_ai::Error::Truncated { .. }) => { + let Some((first, second)) = split_in_half(part) else { + return Err(TaskError::Permanent(TRUNCATED_ERROR.to_owned())); + }; + + debug!(part_len = part.len(), "AI answer truncated; asking again in two halves"); + parts.push_front(second); + parts.push_front(first); + } + Err(error) => return Err(error.into()), + } + } + + Ok(actions) +} + +/// Cuts `part` on the line boundary closest to its middle, unless it is too short to split. +fn split_in_half(part: &str) -> Option<(&str, &str)> { + if part.len() < MIN_SPLIT_LEN { + return None; + } + + let bytes = part.as_bytes(); + let middle = part.len() / 2; + + let before = bytes[..middle] + .iter() + .rposition(|&byte| byte == b'\n') + .map(|end| end + 1); + let after = bytes[middle..] + .iter() + .position(|&byte| byte == b'\n') + .map(|end| middle + end + 1); + + let cut = match (before, after) { + (Some(before), Some(after)) if middle - before <= after - middle => before, + (_, Some(after)) if after < part.len() => after, + (Some(before), _) => before, + _ => return None, + }; + + Some(part.split_at(cut)) +} + +/// Writes the `.slog` from the checkpoints, and returns the number of actions. +fn merge_checkpoints( + workspace: &Utf8Path, + total: usize, + start_time: i64, + duration: i64, + model: &str, + log_path: &Utf8Path, +) -> Result { + let log_error = |error: anyhow::Error| TaskError::Permanent(format!("failed to write the log: {error:#}")); + + let out = File::create(log_path).map_err(|error| workspace_error(&error))?; + let mut log = slog::SlogWriter::start(BufWriter::new(out), start_time, model).map_err(log_error)?; + let mut count = 0; + + // Chunks follow each other in time, so only the actions of one chunk need sorting. + for index in 0..total { + let mut actions = + checkpoint::read(&checkpoint::path(workspace, index)).map_err(|error| workspace_error(&error))?; + actions.sort_by_key(|action| action.offset); + + for action in &actions { + log.action(action).map_err(log_error)?; + } + + count += actions.len(); + } + + log.finish(duration) + .map_err(log_error)? + .into_inner() + .map_err(|error| workspace_error(&error.into_error()))? + .sync_all() + .map_err(|error| workspace_error(&error))?; + + Ok(count) +} + +impl EphemeralTask for AiLogTask { + type Secrets = SecretString; + type Request = AiLogParams; + + fn prepare( + target: &AiLogTarget, + request: AiLogParams, + state: &DgwState, + ) -> Result<(AiSettings, SecretString), TaskErrorCode> { + if state.recordings.active_recordings.contains(target.session_id) { + return Err(TaskErrorCode::RecordingActive); + } + + let AiLogParams { + provider, + model, + api_key, + base_url, + max_output_tokens, + } = request; + + let settings = AiSettings { + provider, + model, + base_url, + max_output_tokens, + }; + + settings.check(&api_key, state)?; + + Ok((settings, api_key)) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + const API_KEY: &str = "sk-ai-log-test-secret"; + + const CONFIG: &str = r#"{ + "ProvisionerPublicKeyData": { + "Value": "mMIIBIjANBgkqhkiG9w0BAQEFAAOCAQ8AMIIBCgKCAQEA4vuqLOkl1pWobt6su1XO9VskgCAwevEGs6kkNjJQBwkGnPKYLmNF1E/af1yCocfVn/OnPf9e4x+lXVyZ6LMDJxFxu+axdgOq3Ld392J1iAEbfvwlyRFnEXFOJNyylqg3bY6LvnWHL/XZczVdMD9xYfq2sO9bg3xjRW4s7r9EEYOFjqVT3VFznH9iWJVtcSEKukmS/3uKoO6lGhacvu0HhjXXdgq0R8zvR4XRJ9Fcnf0f9Ypoc+i6L80NVjrRCeVOH+Ld/2fA9bocpfLarcVqG3RjS+qgOtpyCc0jWVFF4zaGQ7LUDFkEIYILkICeMMn2ll29hmZNzsJzZJ9s6NocgQIDAQAB" + }, + "Listeners": [{ "InternalUrl": "http://*:7171", "ExternalUrl": "https://*:7171" }], + "Proxy": { "Mode": "Off" } + }"#; + + fn params() -> AiLogParams { + serde_json::from_value(serde_json::json!({ + "provider": "openai", + "model": "gpt-test", + "apiKey": API_KEY, + })) + .expect("valid params") + } + + fn target() -> AiLogTarget { + AiLogTarget { + session_id: Uuid::new_v4(), + } + } + + #[tokio::test] + async fn refuses_a_session_that_is_still_recording() { + let (state, _handles) = DgwState::mock(CONFIG).expect("mock state"); + let target = target(); + state.recordings.active_recordings.insert(target.session_id); + + let error = AiLogTask::prepare(&target, params(), &state).expect_err("session is busy"); + + assert_eq!(error, TaskErrorCode::RecordingActive); + } + + #[tokio::test] + async fn persisted_settings_never_hold_the_api_key() { + let (state, _handles) = DgwState::mock(CONFIG).expect("mock state"); + + let params = params(); + assert!(!format!("{params:?}").contains(API_KEY)); + + let (settings, api_key) = AiLogTask::prepare(&target(), params, &state).expect("valid task"); + + let persisted = serde_json::to_string(&settings).expect("serializable settings"); + assert_eq!( + persisted, + r#"{"provider":"openai","model":"gpt-test","baseUrl":null,"maxOutputTokens":null}"# + ); + assert!(!format!("{settings:?}").contains(API_KEY)); + assert!(!format!("{api_key:?}").contains(API_KEY)); + } + + #[test] + fn describing_substate_reports_chunk_progress() { + let substate = serde_json::to_value(AiLogSubstate::Describing { done: 1, total: 3 }).expect("serializable"); + + assert_eq!( + substate, + serde_json::json!({ "step": "describing", "done": 1, "total": 3 }) + ); + } + + #[test] + fn halves_are_cut_on_the_line_boundary_closest_to_the_middle() { + let line = format!("[1.0] {}\n", "a".repeat(94)); + let part = line.repeat(30); + + let (first, second) = split_in_half(&part).expect("long enough"); + + assert_eq!(first.len(), 15 * line.len()); + assert_eq!(second.len(), 15 * line.len()); + + let uneven = format!("{line}{}", "b".repeat(1900)); + let (first, second) = split_in_half(&uneven).expect("long enough"); + assert_eq!(first, line); + assert_eq!(second, "b".repeat(1900)); + + assert!(split_in_half(&line.repeat(3)).is_none(), "too short"); + assert!(split_in_half(&"c".repeat(MIN_SPLIT_LEN * 2)).is_none(), "one line"); + } + + #[tokio::test] + async fn invalid_ai_settings_are_refused_with_a_code() { + let (state, _handles) = DgwState::mock(CONFIG).expect("mock state"); + let mut params = params(); + params.model = " ".to_owned(); + + let error = AiLogTask::prepare(&target(), params, &state).expect_err("empty model"); + + assert_eq!(error, TaskErrorCode::MissingModel); + } +} diff --git a/devolutions-gateway/src/tasks/ai_log/slog.rs b/devolutions-gateway/src/tasks/ai_log/slog.rs new file mode 100644 index 000000000..7b7a1c2bd --- /dev/null +++ b/devolutions-gateway/src/tasks/ai_log/slog.rs @@ -0,0 +1,171 @@ +//! Writes AI actions as a Session Recording Log (`.slog`): one JSON object per line, in the shipped schema. + +use std::collections::BTreeMap; +use std::time::Duration; + +use devolutions_gateway_ai::session_actions::{Action, PROMPT_VERSION}; + +#[derive(Serialize)] +#[serde(rename_all = "camelCase")] +struct Entry<'a> { + timestamp: String, + seq: usize, + event: &'static str, + description: &'a str, + #[serde(skip_serializing_if = "Option::is_none")] + object: Option<&'a str>, + #[serde(skip_serializing_if = "Option::is_none")] + parameters: Option<&'a BTreeMap>, + #[serde(skip_serializing_if = "Option::is_none")] + source: Option<&'static str>, + #[serde(skip_serializing_if = "Option::is_none")] + model: Option<&'a str>, + #[serde(skip_serializing_if = "Option::is_none")] + prompt_version: Option<&'static str>, +} + +impl<'a> Entry<'a> { + fn new(seq: usize, timestamp: String, event: &'static str, description: &'a str) -> Self { + Self { + timestamp, + seq, + event, + description, + object: None, + parameters: None, + source: None, + model: None, + prompt_version: None, + } + } +} + +/// Writes the log of a session entry by entry, so the actions never have to be all in memory. +pub(crate) struct SlogWriter { + out: W, + start_time: i64, + seq: usize, +} + +impl SlogWriter { + /// Writes `session.start` for a session that started at `start_time` (unix seconds). + pub(crate) fn start(out: W, start_time: i64, model: &str) -> anyhow::Result { + let mut writer = Self { + out, + start_time, + seq: 0, + }; + + let timestamp = timestamp(start_time, Duration::ZERO)?; + writer.write(Entry { + source: Some("ai"), + model: Some(model), + prompt_version: Some(PROMPT_VERSION), + ..Entry::new(0, timestamp, "session.start", "Session started") + })?; + + Ok(writer) + } + + /// Actions must come in offset order. + pub(crate) fn action(&mut self, action: &Action) -> anyhow::Result<()> { + let timestamp = timestamp(self.start_time, action.offset)?; + self.write(Entry { + object: action.object.as_deref(), + parameters: Some(&action.parameters).filter(|parameters| !parameters.is_empty()), + ..Entry::new(self.seq, timestamp, "session.action", &action.description) + }) + } + + /// Writes `session.end` for a session that lasted `duration` seconds. + pub(crate) fn finish(mut self, duration: i64) -> anyhow::Result { + let end_offset = Duration::from_secs(u64::try_from(duration).unwrap_or(0)); + let timestamp = timestamp(self.start_time, end_offset)?; + self.write(Entry::new(self.seq, timestamp, "session.end", "Session ended"))?; + self.out.flush()?; + Ok(self.out) + } + + fn write(&mut self, entry: Entry<'_>) -> anyhow::Result<()> { + serde_json::to_writer(&mut self.out, &entry)?; + self.out.write_all(b"\n")?; + self.seq += 1; + Ok(()) + } +} + +/// ISO 8601 UTC time with milliseconds, like the other `.slog` writers. +fn timestamp(start_time: i64, offset: Duration) -> anyhow::Result { + let time = time::OffsetDateTime::from_unix_timestamp(start_time)? + .checked_add(time::Duration::try_from(offset)?) + .ok_or_else(|| anyhow::anyhow!("timestamp out of range"))?; + + Ok(format!( + "{:04}-{:02}-{:02}T{:02}:{:02}:{:02}.{:03}Z", + time.year(), + u8::from(time.month()), + time.day(), + time.hour(), + time.minute(), + time.second(), + time.millisecond(), + )) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn write(start_time: i64, duration: i64, actions: &[Action]) -> String { + let mut writer = SlogWriter::start(Vec::new(), start_time, "gpt-test").expect("started"); + for action in actions { + writer.action(action).expect("written"); + } + String::from_utf8(writer.finish(duration).expect("finished")).expect("UTF-8") + } + + fn action(offset_ms: u64, description: &str, object: Option<&str>, parameters: &[(&str, &str)]) -> Action { + Action { + offset: Duration::from_millis(offset_ms), + description: description.to_owned(), + object: object.map(str::to_owned), + parameters: parameters + .iter() + .map(|(key, value)| ((*key).to_owned(), (*value).to_owned())) + .collect(), + } + } + + #[test] + fn writes_the_shipped_schema_with_ai_fields_on_session_start() { + let actions = [ + action(1500, "Listed files", Some("/var/log"), &[("Command", "ls /var/log")]), + action(62_250, "Closed the shell", None, &[]), + ]; + + let slog = write(1_787_255_035, 90, &actions); + + let expected = [ + r#"{"timestamp":"2026-08-20T19:43:55.000Z","seq":0,"event":"session.start","description":"Session started","source":"ai","model":"gpt-test","promptVersion":"session-actions-1"}"#, + r#"{"timestamp":"2026-08-20T19:43:56.500Z","seq":1,"event":"session.action","description":"Listed files","object":"/var/log","parameters":{"Command":"ls /var/log"}}"#, + r#"{"timestamp":"2026-08-20T19:44:57.250Z","seq":2,"event":"session.action","description":"Closed the shell"}"#, + r#"{"timestamp":"2026-08-20T19:45:25.000Z","seq":3,"event":"session.end","description":"Session ended"}"#, + ]; + assert_eq!(slog, expected.map(|line| format!("{line}\n")).concat()); + } + + #[test] + fn a_session_without_actions_still_starts_and_ends() { + let slog = write(0, 5, &[]); + + let events: Vec = slog + .lines() + .map(|line| serde_json::from_str(line).expect("JSON line")) + .collect(); + assert_eq!(events.len(), 2); + assert_eq!(events[0]["event"], "session.start"); + assert_eq!(events[1]["event"], "session.end"); + assert_eq!(events[1]["seq"], 1); + assert_eq!(events[1]["timestamp"], "1970-01-01T00:00:05.000Z"); + } +} diff --git a/devolutions-gateway/src/tasks/ai_log/transcript.rs b/devolutions-gateway/src/tasks/ai_log/transcript.rs new file mode 100644 index 000000000..f2dbcfbc1 --- /dev/null +++ b/devolutions-gateway/src/tasks/ai_log/transcript.rs @@ -0,0 +1,595 @@ +//! Turns the terminal recordings of a session into transcript chunk files sent to the AI. +//! +//! Each transcript line reads `[] `. +//! Recordings are read as a stream, so memory use depends on the chunk size, not on the recording size. + +use std::fs::File; +use std::io::{self, BufRead, BufReader, BufWriter, Write as _}; + +use camino::{Utf8Path, Utf8PathBuf}; + +use crate::recording::FinishedRecording; +use crate::token::RecordingFileType; + +/// Longest text kept from one terminal line, in characters. +const MAX_LINE_CHARS: usize = 500; + +/// Longest asciicast event read; a longer one is skipped. +const MAX_CAST_EVENT_LEN: usize = 1024 * 1024; + +#[derive(Debug, thiserror::Error)] +pub(crate) enum TranscriptError { + #[error("unsupported recording type: {0}")] + Unsupported(&'static str), + #[error("session has no terminal recording")] + NoTerminalRecording, + #[error("failed to read {file_name}: {reason}")] + Invalid { file_name: String, reason: String }, + #[error("failed to write the transcript: {0}")] + Write(io::Error), +} + +pub(crate) fn chunk_path(workspace: &Utf8Path, index: usize) -> Utf8PathBuf { + workspace.join(format!("chunk-{index:04}.txt")) +} + +/// Writes the transcript of the terminal recordings of `recording` into chunk files of at most `max_len` bytes, +/// cut on line boundaries, and returns the number of chunks. +pub(crate) fn write_chunks( + recording: &FinishedRecording, + workspace: &Utf8Path, + max_len: usize, +) -> Result { + let manifest = &recording.manifest; + + let mut transcript = Transcript::new(ChunkWriter::new(workspace, max_len)); + let mut has_terminal_recording = false; + let mut unsupported = None; + + for file in manifest.files() { + let file_type = Utf8Path::new(file.file_name()) + .extension() + .and_then(RecordingFileType::from_extension); + + let to_transcript_error = |error: CastError| match error { + CastError::Read(reason) => TranscriptError::Invalid { + file_name: file.file_name().to_owned(), + reason, + }, + CastError::Write(error) => TranscriptError::Write(error), + }; + + let offset = u32::try_from(file.start_time().saturating_sub(manifest.start_time()).max(0)).unwrap_or(u32::MAX); + let offset = f64::from(offset); + + let open = || { + File::open(recording.dir.join(file.file_name())) + .map(BufReader::new) + .map_err(|error| to_transcript_error(CastError::Read(error.to_string()))) + }; + + match file_type { + Some(RecordingFileType::Asciicast) => transcript.add_cast(open()?, offset).map_err(to_transcript_error)?, + Some(RecordingFileType::TRP) => transcript.add_trp(open()?, offset).map_err(to_transcript_error)?, + Some(RecordingFileType::WebM) => { + unsupported.get_or_insert("webm"); + continue; + } + Some(RecordingFileType::SessionRecordingLog) | None => continue, + } + + has_terminal_recording = true; + } + + if !has_terminal_recording { + return Err(unsupported.map_or(TranscriptError::NoTerminalRecording, TranscriptError::Unsupported)); + } + + transcript.chunks.finish().map_err(TranscriptError::Write) +} + +enum CastError { + Read(String), + Write(io::Error), +} + +/// Writes lines into numbered chunk files, starting a new file when the next line would not fit. +struct ChunkWriter { + workspace: Utf8PathBuf, + max_len: usize, + count: usize, + current: Option>, + current_len: usize, +} + +impl ChunkWriter { + fn new(workspace: &Utf8Path, max_len: usize) -> Self { + Self { + workspace: workspace.to_owned(), + max_len, + count: 0, + current: None, + current_len: 0, + } + } + + fn write_line(&mut self, line: &str) -> io::Result<()> { + if self.current.is_some() && self.current_len + line.len() > self.max_len { + self.close_current()?; + } + + let current = match &mut self.current { + Some(current) => current, + None => { + let file = File::create(chunk_path(&self.workspace, self.count))?; + self.count += 1; + self.current_len = 0; + self.current.insert(BufWriter::new(file)) + } + }; + + current.write_all(line.as_bytes())?; + self.current_len += line.len(); + + Ok(()) + } + + fn close_current(&mut self) -> io::Result<()> { + if let Some(current) = self.current.take() { + current + .into_inner() + .map_err(io::IntoInnerError::into_error)? + .sync_all()?; + } + + Ok(()) + } + + fn finish(mut self) -> io::Result { + self.close_current()?; + Ok(self.count) + } +} + +struct Transcript { + chunks: ChunkWriter, + last_line: String, +} + +impl Transcript { + fn new(chunks: ChunkWriter) -> Self { + Self { + chunks, + last_line: String::new(), + } + } + + /// Adds the terminal output of an asciicast v2 or v3 recording that started `offset` seconds into the session. + fn add_cast(&mut self, mut cast: impl BufRead, offset: f64) -> Result<(), CastError> { + let mut line = Vec::new(); + let read_error = |error: io::Error| CastError::Read(error.to_string()); + + let header = loop { + if !read_event(&mut cast, &mut line).map_err(read_error)? { + return Err(CastError::Read("empty asciicast".to_owned())); + } + + if !line.trim_ascii().is_empty() { + break serde_json::from_slice::(&line) + .map_err(|error| CastError::Read(format!("invalid asciicast header: {error}")))?; + } + }; + + let relative_times = header["version"].as_u64() == Some(3); + + let mut terminal = TerminalText::default(); + let mut time = 0.0; + + while read_event(&mut cast, &mut line).map_err(read_error)? { + // Only the output is used: typed passwords are usually not echoed, but printed secrets still reach the AI. + let Ok((event_time, code, data)) = serde_json::from_slice::<(f64, String, String)>(&line) else { + continue; + }; + + time = if relative_times { time + event_time } else { event_time }; + + if code == "o" { + terminal.feed(offset + time, &data, self).map_err(CastError::Write)?; + } + } + + terminal.end_line(self).map_err(CastError::Write) + } + + /// Adds the terminal output of a TRP recording that started `offset` seconds into the session. + fn add_trp(&mut self, trp: impl io::Read, offset: f64) -> Result<(), CastError> { + let mut terminal = TerminalText::default(); + + for output in terminal_streamer::trp_decoder::TrpOutputReader::new(trp) { + let output = output.map_err(|error| CastError::Read(error.to_string()))?; + terminal + .feed(offset + output.time, &output.text, self) + .map_err(CastError::Write)?; + } + + terminal.end_line(self).map_err(CastError::Write) + } + + fn push_line(&mut self, time: f64, text: &str) -> io::Result<()> { + let text = text.trim(); + + if text.is_empty() || text == self.last_line { + return Ok(()); + } + + let text = match text.char_indices().nth(MAX_LINE_CHARS) { + Some((cut, _)) => &text[..cut], + None => text, + }; + + self.chunks.write_line(&format!("[{time:.1}] {text}\n"))?; + text.clone_into(&mut self.last_line); + + Ok(()) + } +} + +/// Reads the next asciicast event into `line`, skipping events longer than [`MAX_CAST_EVENT_LEN`]. +fn read_event(reader: &mut impl BufRead, line: &mut Vec) -> io::Result { + let limit = u64::try_from(MAX_CAST_EVENT_LEN).unwrap_or(u64::MAX); + + loop { + line.clear(); + + if io::Read::take(&mut *reader, limit).read_until(b'\n', line)? == 0 { + return Ok(false); + } + + if line.len() < MAX_CAST_EVENT_LEN || line.ends_with(b"\n") { + return Ok(true); + } + + warn!(max_len = MAX_CAST_EVENT_LEN, "Skipped an oversized asciicast event"); + skip_line(reader)?; + } +} + +fn skip_line(reader: &mut impl BufRead) -> io::Result<()> { + loop { + let buffer = reader.fill_buf()?; + + if buffer.is_empty() { + return Ok(()); + } + + if let Some(end) = buffer.iter().position(|&byte| byte == b'\n') { + reader.consume(end + 1); + return Ok(()); + } + + let len = buffer.len(); + reader.consume(len); + } +} + +#[derive(Default, Clone, Copy)] +enum EscapeState { + #[default] + Text, + Escape, + EscapeArgument, + ControlSequence, + OperatingSystemCommand, + OperatingSystemCommandEscape, +} + +/// Rebuilds the text lines of a terminal output stream, without escape sequences. +#[derive(Default)] +struct TerminalText { + state: EscapeState, + line: String, + line_start: Option, + carriage_return: bool, +} + +impl TerminalText { + fn feed(&mut self, time: f64, data: &str, transcript: &mut Transcript) -> io::Result<()> { + for c in data.chars() { + self.state = match (self.state, c) { + (EscapeState::Text, '\x1b') => EscapeState::Escape, + (EscapeState::Text, '\n') => { + self.end_line(transcript)?; + EscapeState::Text + } + (EscapeState::Text, '\r') => { + self.carriage_return = true; + EscapeState::Text + } + (EscapeState::Text, '\x08') => { + self.line.pop(); + EscapeState::Text + } + (EscapeState::Text, '\t') => { + self.put(' ', time); + EscapeState::Text + } + (EscapeState::Text, c) => { + if !c.is_control() { + self.put(c, time); + } + EscapeState::Text + } + (EscapeState::Escape, '[') => EscapeState::ControlSequence, + (EscapeState::Escape, ']') => EscapeState::OperatingSystemCommand, + (EscapeState::Escape, '(' | ')' | '*' | '+' | '#' | '%') => EscapeState::EscapeArgument, + (EscapeState::Escape | EscapeState::EscapeArgument | EscapeState::OperatingSystemCommandEscape, _) => { + EscapeState::Text + } + (EscapeState::ControlSequence, '\x40'..='\x7e') => EscapeState::Text, + (EscapeState::ControlSequence, _) => EscapeState::ControlSequence, + (EscapeState::OperatingSystemCommand, '\x07') => EscapeState::Text, + (EscapeState::OperatingSystemCommand, '\x1b') => EscapeState::OperatingSystemCommandEscape, + (EscapeState::OperatingSystemCommand, _) => EscapeState::OperatingSystemCommand, + }; + } + + Ok(()) + } + + fn put(&mut self, c: char, time: f64) { + // A carriage return not followed by a line feed means the line is redrawn. + if self.carriage_return { + self.carriage_return = false; + self.line.clear(); + self.line_start = None; + } + + self.line_start.get_or_insert(time); + + // Text past the kept length is dropped anyway, so a line that never ends cannot grow without bound. + if self.line.len() < MAX_LINE_CHARS * 4 { + self.line.push(c); + } + } + + fn end_line(&mut self, transcript: &mut Transcript) -> io::Result<()> { + self.carriage_return = false; + + if let Some(start) = self.line_start.take() { + transcript.push_line(start, &self.line)?; + } + + self.line.clear(); + + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + use std::sync::atomic::{AtomicUsize, Ordering}; + + use super::*; + + fn read_chunks(dir: &Utf8Path, count: usize) -> Vec { + (0..count) + .map(|index| std::fs::read_to_string(chunk_path(dir, index)).expect("chunk file")) + .collect() + } + + fn transcript_of(cast: &str, offset: f64) -> String { + let dir = tempfile::tempdir().expect("temp dir"); + let dir = Utf8Path::from_path(dir.path()).expect("UTF-8"); + let mut transcript = Transcript::new(ChunkWriter::new(dir, usize::MAX)); + assert!(transcript.add_cast(cast.as_bytes(), offset).is_ok()); + let count = transcript.chunks.finish().expect("written"); + read_chunks(dir, count).concat() + } + + const CAST: &str = r#"{"version": 2, "width": 80, "height": 24} +[0.3,"o","\u001b]0;user@host: ~\u0007\u001b[01;32muser@host\u001b[00m:~$ "] +[1.0,"i","l"] +[1.1,"o","l"] +[1.2,"o","s\r\n"] +[1.5,"o","file-a file-b\r\n"] +[2.0,"o","user@host:~$ sudo passwd david\r\n[sudo] password for user: "] +[3.0,"i","hunter2\r"] +[3.5,"o","\r\n"] +[4.0,"o","progress 10%\rprogress 100%\r\n"] +[5.0,"o","same\r\nsame\r\n"] +[6.0,"o","typo\b\b\u001b[Kps\r\n"] +[7.0,"r","100x30"] +"#; + + #[test] + fn cast_output_becomes_timed_plain_lines() { + assert_eq!( + transcript_of(CAST, 10.0), + concat!( + "[10.3] user@host:~$ ls\n", + "[11.5] file-a file-b\n", + "[12.0] user@host:~$ sudo passwd david\n", + "[12.0] [sudo] password for user:\n", + "[14.0] progress 100%\n", + "[15.0] same\n", + "[16.0] typs\n", + ) + ); + } + + #[test] + fn asciicast_v3_times_are_intervals() { + let cast = "{\"version\": 3, \"term\": {\"cols\": 80, \"rows\": 24}}\n[1.0,\"o\",\"a\\r\\n\"]\n[0.5,\"o\",\"b\\r\\n\"]\n"; + + assert_eq!(transcript_of(cast, 0.0), "[1.0] a\n[1.5] b\n"); + } + + #[test] + fn long_lines_are_cut() { + let cast = format!( + "{{\"version\": 2}}\n[0,\"o\",\"{}\"]\n", + "é".repeat(MAX_LINE_CHARS + 10) + ); + + let transcript = transcript_of(&cast, 0.0); + + assert_eq!(transcript.trim_end().chars().count(), "[0.0] ".len() + MAX_LINE_CHARS); + } + + #[test] + fn oversized_cast_events_are_skipped() { + let cast = format!( + "{{\"version\": 2}}\n[0,\"o\",\"{}\"]\n[1,\"o\",\"kept\\r\\n\"]\n", + "x".repeat(MAX_CAST_EVENT_LEN) + ); + + assert_eq!(transcript_of(&cast, 0.0), "[1.0] kept\n"); + } + + #[test] + fn chunks_keep_whole_lines_within_the_budget() { + let dir = tempfile::tempdir().expect("temp dir"); + let dir = Utf8Path::from_path(dir.path()).expect("UTF-8"); + let mut chunks = ChunkWriter::new(dir, 22); + for line in ["[1.0] aaaa\n", "[2.0] bbbb\n", "[3.0] cccc\n"] { + chunks.write_line(line).expect("written"); + } + + let count = chunks.finish().expect("written"); + + assert_eq!(read_chunks(dir, count), ["[1.0] aaaa\n[2.0] bbbb\n", "[3.0] cccc\n"]); + assert_eq!(ChunkWriter::new(dir, 22).finish().expect("nothing written"), 0); + } + + /// Generates a large cast lazily and records how many events were read when the first chunk file was complete. + struct LazyCast { + next_event: usize, + events: usize, + pending: Vec, + second_chunk: Utf8PathBuf, + read_when_first_chunk_done: Arc, + } + + impl io::Read for LazyCast { + fn read(&mut self, buf: &mut [u8]) -> io::Result { + if self.pending.is_empty() { + if self.next_event == self.events { + return Ok(0); + } + + if self.second_chunk.exists() { + let _ = self.read_when_first_chunk_done.compare_exchange( + 0, + self.next_event, + Ordering::Relaxed, + Ordering::Relaxed, + ); + } + + self.pending = if self.next_event == 0 { + b"{\"version\": 2}\n".to_vec() + } else { + let event = self.next_event; + format!("[{event}.0,\"o\",\"line {event} of a long session\\r\\n\"]\n").into_bytes() + }; + self.next_event += 1; + } + + let n = buf.len().min(self.pending.len()); + buf[..n].copy_from_slice(&self.pending[..n]); + self.pending.drain(..n); + Ok(n) + } + } + + #[test] + fn large_casts_are_streamed_into_bounded_chunk_files() { + let dir = tempfile::tempdir().expect("temp dir"); + let dir = Utf8Path::from_path(dir.path()).expect("UTF-8"); + let read_when_first_chunk_done = Arc::new(AtomicUsize::new(0)); + let events = 20_000; + let cast = LazyCast { + next_event: 0, + events, + pending: Vec::new(), + second_chunk: chunk_path(dir, 1), + read_when_first_chunk_done: Arc::clone(&read_when_first_chunk_done), + }; + + let mut transcript = Transcript::new(ChunkWriter::new(dir, 4096)); + assert!(transcript.add_cast(BufReader::with_capacity(64, cast), 0.0).is_ok()); + let count = transcript.chunks.finish().expect("written"); + + assert!(count > 100, "{count}"); + for chunk in read_chunks(dir, count) { + assert!(chunk.len() <= 4096); + assert!(chunk.ends_with('\n')); + } + + let read = read_when_first_chunk_done.load(Ordering::Relaxed); + assert!( + read > 0 && read < events / 50, + "first chunk done after {read} of {events} events" + ); + } + + fn session_dir(manifest: &str, files: &[(&str, &[u8])]) -> tempfile::TempDir { + let dir = tempfile::tempdir().expect("temp dir"); + std::fs::write(dir.path().join("recording.json"), manifest).expect("write manifest"); + for (name, contents) in files { + std::fs::write(dir.path().join(name), contents).expect("write recording"); + } + dir + } + + fn read(dir: &tempfile::TempDir) -> Result { + let path = Utf8Path::from_path(dir.path()).expect("utf8 path"); + let recording = FinishedRecording::read_for_test(path); + let workspace = path.join("workspace"); + std::fs::create_dir(&workspace).expect("workspace"); + let count = write_chunks(&recording, &workspace, usize::MAX)?; + Ok(read_chunks(&workspace, count).concat()) + } + + #[test] + fn recordings_are_offset_by_their_start_time() { + let manifest = r#"{"sessionId":"22fcd533-5e72-4db7-aa0f-29952dbbca9f","startTime":100,"duration":90, + "files":[{"fileName":"recording-0.cast","startTime":100,"duration":10}, + {"fileName":"recording-1.cast","startTime":160,"duration":30}], + "artifacts":{"ai-analysis":[{"fileName":"ai-analysis-0.slog"}]}}"#; + let first: &[u8] = b"{\"version\": 2}\n[1.5,\"o\",\"whoami\\r\\n\"]\n"; + let second: &[u8] = b"{\"version\": 2}\n[2.0,\"o\",\"exit\\r\\n\"]\n"; + let dir = session_dir(manifest, &[("recording-0.cast", first), ("recording-1.cast", second)]); + + assert_eq!(read(&dir).expect("transcript"), "[1.5] whoami\n[62.0] exit\n"); + } + + #[test] + fn trp_recordings_are_decoded() { + let manifest = r#"{"sessionId":"22fcd533-5e72-4db7-aa0f-29952dbbca9f","startTime":100,"duration":9, + "files":[{"fileName":"recording-0.trp","startTime":100,"duration":9}]}"#; + let mut trp = Vec::new(); + for (time_delta, event_type, payload) in [(0u32, 4u16, &b""[..]), (2500, 0, &b"uptime\r\n"[..])] { + trp.extend_from_slice(&time_delta.to_le_bytes()); + trp.extend_from_slice(&event_type.to_le_bytes()); + trp.extend_from_slice(&u16::try_from(payload.len()).expect("small").to_le_bytes()); + trp.extend_from_slice(payload); + } + let dir = session_dir(manifest, &[("recording-0.trp", &trp)]); + + assert_eq!(read(&dir).expect("transcript"), "[2.5] uptime\n"); + } + + #[test] + fn video_only_sessions_are_unsupported() { + let manifest = r#"{"sessionId":"22fcd533-5e72-4db7-aa0f-29952dbbca9f","startTime":1,"duration":5, + "files":[{"fileName":"recording-0.webm","startTime":1,"duration":5}]}"#; + let dir = session_dir(manifest, &[]); + + assert_eq!( + read(&dir).expect_err("webm only").to_string(), + "unsupported recording type: webm" + ); + } +} diff --git a/devolutions-gateway/src/tasks/mod.rs b/devolutions-gateway/src/tasks/mod.rs index 90e1b1cfb..3daae9d43 100644 --- a/devolutions-gateway/src/tasks/mod.rs +++ b/devolutions-gateway/src/tasks/mod.rs @@ -5,6 +5,8 @@ //! tasks never take a slot from the other Gateway jobs, and other jobs never delay a task. //! The job definition holds only the persisted, non-secret parameters, so a [`DurableTask`] resumes after a restart. //! The secrets of an [`EphemeralTask`] stay in memory only: when Gateway restarts, the task fails instead. +//! A task may keep intermediate files in its workspace, under the recording folder so it gets the same access rules; +//! the workspace is kept across attempts and deleted when the task finishes. //! //! The task system is unstable: it starts only when `__debug__.enable_unstable` is set. @@ -20,6 +22,7 @@ use std::time::Duration; use anyhow::Context as _; use async_trait::async_trait; +use camino::Utf8PathBuf; use devolutions_gateway_task::{ShutdownSignal, Task}; use job_queue::{DynJob, DynJobQueue, JobQueue as _, JobReader, RunnerWaker}; use job_queue_libsql::{LibSqlJobQueue, libsql}; @@ -46,6 +49,9 @@ pub const SECRETS_LOST_ERROR: &str = "gateway restarted, API key no longer avail pub const JOB_LOST_ERROR: &str = "gateway restarted, task job no longer exists"; +/// Folder of the task workspaces, in the recording folder. +pub const WORKSPACES_DIR: &str = ".provisioner-tasks"; + /// Why a run of a task failed; the message is stored in the task record and returned by the API. #[derive(Debug, Clone, PartialEq, Eq)] pub enum TaskError { @@ -182,6 +188,8 @@ pub struct TaskCtx { pub params: K::Params, pub state: DgwState, pub progress: Progress, + /// Folder of files kept across attempts; not created until the task needs it. + pub workspace: Utf8PathBuf, secrets: Option>, } @@ -215,6 +223,7 @@ struct TaskServiceInner { notify_runner: Arc, runner_waker: RunnerWaker, secrets: Mutex, + workspaces: Utf8PathBuf, timeout: Duration, } @@ -233,13 +242,17 @@ impl TaskService { return Ok(None); } - Self::open(conf.provisioner_tasks_database.as_str(), TASK_TIMEOUT) - .await - .map(Some) + Self::open( + conf.provisioner_tasks_database.as_str(), + conf.recording_path.join(WORKSPACES_DIR), + TASK_TIMEOUT, + ) + .await + .map(Some) } - /// Opens the database at `path`, then fails every unfinished task whose job is gone. - async fn open(path: &str, timeout: Duration) -> anyhow::Result { + /// Opens the database at `path`, fails every unfinished task whose job is gone, then deletes stale workspaces. + async fn open(path: &str, workspaces: Utf8PathBuf, timeout: Duration) -> anyhow::Result { let conn = libsql::Builder::new_local(path) .build() .await @@ -283,6 +296,7 @@ impl TaskService { notify_runner, runner_waker, secrets: Mutex::new(HashMap::new()), + workspaces, timeout, }), }; @@ -292,9 +306,45 @@ impl TaskService { .await .context("failed to reconcile the provisioner tasks")?; + service + .remove_stale_workspaces() + .await + .context("failed to remove the stale task workspaces")?; + Ok(service) } + fn workspace(&self, id: Uuid) -> Utf8PathBuf { + self.inner.workspaces.join(id.to_string()) + } + + /// Deletes every workspace that no unfinished task owns. + async fn remove_stale_workspaces(&self) -> anyhow::Result<()> { + let mut entries = match tokio::fs::read_dir(&self.inner.workspaces).await { + Ok(entries) => entries, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(()), + Err(error) => return Err(error.into()), + }; + + let unfinished = self.store().unfinished().await?.into_iter().collect::>(); + + while let Some(entry) = entries.next_entry().await? { + let owned = entry + .file_name() + .to_str() + .and_then(|name| Uuid::parse_str(name).ok()) + .is_some_and(|id| unfinished.contains(&id)); + + if !owned { + let path = entry.path(); + info!(path = %path.display(), "Removing a stale task workspace"); + remove_all(&path).await; + } + } + + Ok(()) + } + fn store(&self) -> &LibSqlProvisionerTaskStore { &self.inner.store } @@ -473,7 +523,7 @@ impl TaskService { let Some(attempt) = store.start_attempt(id, &substate).await? else { debug!(task.id = %id, task.kind = K::KIND, "Background task is already finished"); - self.forget_secrets(id); + self.finish(id).await; return Ok(()); }; @@ -490,6 +540,7 @@ impl TaskService { tasks: self.clone(), _substate: PhantomData, }, + workspace: self.workspace(id), secrets, }; @@ -518,24 +569,44 @@ impl TaskService { } } - self.forget_secrets(id); + self.finish(id).await; Ok(()) } async fn fail(&self, id: Uuid, error: &str) { - self.forget_secrets(id); + self.finish(id).await; if let Err(store_error) = self.store().fail(id, error).await { error!(task.id = %id, error = format!("{store_error:#}"), "Failed to record the task failure"); } } + /// Drops what a task keeps only until it finishes: its secrets and its workspace. + async fn finish(&self, id: Uuid) { + self.forget_secrets(id); + remove_all(self.workspace(id).as_std_path()).await; + } + fn forget_secrets(&self, id: Uuid) { self.inner.secrets.lock().remove(&id); } } +async fn remove_all(path: &std::path::Path) { + let removed = if path.is_dir() { + tokio::fs::remove_dir_all(path).await + } else { + tokio::fs::remove_file(path).await + }; + + match removed { + Ok(()) => {} + Err(error) if error.kind() == std::io::ErrorKind::NotFound => {} + Err(error) => warn!(path = %path.display(), %error, "Failed to remove a task workspace"), + } +} + async fn run_attempt(ctx: TaskCtx, timeout: Duration) -> Result { // The run has its own Tokio task so a panic ends as a failure instead of a task stuck in `Running`. let mut handle = tokio::spawn(K::run(ctx)); diff --git a/devolutions-gateway/src/tasks/tests.rs b/devolutions-gateway/src/tasks/tests.rs index ba420b02b..3f4e1dbe1 100644 --- a/devolutions-gateway/src/tasks/tests.rs +++ b/devolutions-gateway/src/tasks/tests.rs @@ -1,19 +1,26 @@ +use axum::http::StatusCode; +use camino::Utf8Path; use job_queue::Job as _; use provisioner_task_store_libsql::TaskState; use super::ai_log::{AiLogTarget, AiLogTask}; use super::*; use crate::MockHandles; +use crate::recording::RecordingManagerTask; const API_KEY: &str = "sk-task-unit-test-secret"; -const CONFIG: &str = r#"{ - "ProvisionerPublicKeyData": { - "Value": "mMIIBIjANBgkqhkiG9w0BAQEFAAOCAQ8AMIIBCgKCAQEA4vuqLOkl1pWobt6su1XO9VskgCAwevEGs6kkNjJQBwkGnPKYLmNF1E/af1yCocfVn/OnPf9e4x+lXVyZ6LMDJxFxu+axdgOq3Ld392J1iAEbfvwlyRFnEXFOJNyylqg3bY6LvnWHL/XZczVdMD9xYfq2sO9bg3xjRW4s7r9EEYOFjqVT3VFznH9iWJVtcSEKukmS/3uKoO6lGhacvu0HhjXXdgq0R8zvR4XRJ9Fcnf0f9Ypoc+i6L80NVjrRCeVOH+Ld/2fA9bocpfLarcVqG3RjS+qgOtpyCc0jWVFF4zaGQ7LUDFkEIYILkICeMMn2ll29hmZNzsJzZJ9s6NocgQIDAQAB" - }, - "Listeners": [{ "InternalUrl": "http://*:7171", "ExternalUrl": "https://*:7171" }], - "Proxy": { "Mode": "Off" } -}"#; +fn config(recording_path: &Utf8Path) -> String { + serde_json::json!({ + "ProvisionerPublicKeyData": { + "Value": "mMIIBIjANBgkqhkiG9w0BAQEFAAOCAQ8AMIIBCgKCAQEA4vuqLOkl1pWobt6su1XO9VskgCAwevEGs6kkNjJQBwkGnPKYLmNF1E/af1yCocfVn/OnPf9e4x+lXVyZ6LMDJxFxu+axdgOq3Ld392J1iAEbfvwlyRFnEXFOJNyylqg3bY6LvnWHL/XZczVdMD9xYfq2sO9bg3xjRW4s7r9EEYOFjqVT3VFznH9iWJVtcSEKukmS/3uKoO6lGhacvu0HhjXXdgq0R8zvR4XRJ9Fcnf0f9Ypoc+i6L80NVjrRCeVOH+Ld/2fA9bocpfLarcVqG3RjS+qgOtpyCc0jWVFF4zaGQ7LUDFkEIYILkICeMMn2ll29hmZNzsJzZJ9s6NocgQIDAQAB" + }, + "Listeners": [{ "InternalUrl": "http://*:7171", "ExternalUrl": "https://*:7171" }], + "Proxy": { "Mode": "Off" }, + "RecordingPath": recording_path, + }) + .to_string() +} #[derive(Debug, Clone, Copy, Serialize, Deserialize)] enum Outcome { @@ -98,9 +105,10 @@ impl DurableTask for LosesStore { struct Harness { state: DgwState, tasks: TaskService, - _handles: MockHandles, + _handles: Box, _dir: tempfile::TempDir, db_path: String, + recording_path: Utf8PathBuf, } impl Harness { @@ -117,18 +125,58 @@ impl Harness { .expect("UTF-8") .to_owned(); - let (state, handles) = DgwState::mock(CONFIG).expect("mock state"); - let tasks = TaskService::open(&db_path, timeout).await.expect("task service"); + let recording_path = Utf8PathBuf::from_path_buf(dir.path().join("recordings")).expect("UTF-8"); + + let (state, handles) = DgwState::mock(&config(&recording_path)).expect("mock state"); + let MockHandles { + session_manager_rx, + recording_manager_rx, + subscriber_rx, + job_queue_rx, + traffic_audit_rx, + shutdown_handle, + } = handles; + + let recording_manager = RecordingManagerTask::new( + recording_manager_rx, + recording_path.clone(), + state.sessions.clone(), + state.job_queue_handle.clone(), + ); + let (recording_shutdown_handle, recording_shutdown_signal) = devolutions_gateway_task::ShutdownHandle::new(); + tokio::spawn(recording_manager.run(recording_shutdown_signal)); + + let tasks = TaskService::open(&db_path, recording_path.join(WORKSPACES_DIR), timeout) + .await + .expect("task service"); Self { state, tasks, - _handles: handles, + _handles: Box::new(( + session_manager_rx, + subscriber_rx, + job_queue_rx, + traffic_audit_rx, + shutdown_handle, + recording_shutdown_handle, + )), _dir: dir, db_path, + recording_path, } } + async fn reopen(&self) -> TaskService { + TaskService::open(&self.db_path, self.recording_path.join(WORKSPACES_DIR), TASK_TIMEOUT) + .await + .expect("task service") + } + + fn workspace(&self, id: Uuid) -> Utf8PathBuf { + self.recording_path.join(WORKSPACES_DIR).join(id.to_string()) + } + /// Reads back the job of a task from the task job queue. async fn queued_job(&self, id: Uuid) -> TaskJob { let json = self @@ -164,17 +212,45 @@ impl Harness { } async fn start_ai_log(&self) -> TaskSnapshot { - let body = serde_json::json!({ "provider": "openai", "model": "gpt-test", "apiKey": API_KEY }).to_string(); - let target = AiLogTarget { - session_id: Uuid::new_v4(), - }; + let body = serde_json::json!({ "provider": "openai", "model": "gpt-test", "apiKey": API_KEY }); + self.start_ai_log_with(Uuid::new_v4(), body).await + } + async fn start_ai_log_with(&self, session_id: Uuid, body: serde_json::Value) -> TaskSnapshot { self.tasks - .start_ephemeral::(target, body.as_bytes(), Uuid::new_v4(), &self.state) + .start_ephemeral::( + AiLogTarget { session_id }, + body.to_string().as_bytes(), + Uuid::new_v4(), + &self.state, + ) .await .expect("task starts") } + /// Writes a finished session whose terminal output is `lines` lines of about 100 bytes each. + fn write_long_session(&self, lines: usize) -> Uuid { + let session_id = Uuid::new_v4(); + let dir = self.recording_path.join(session_id.to_string()); + std::fs::create_dir_all(&dir).expect("session dir"); + + let manifest = serde_json::json!({ + "sessionId": session_id, + "startTime": 1_787_255_035, + "duration": lines, + "files": [{ "fileName": "recording-0.cast", "startTime": 1_787_255_035, "duration": lines }], + }); + std::fs::write(dir.join("recording.json"), manifest.to_string()).expect("manifest"); + + let mut cast = String::from("{\"version\": 2}\n"); + for line in 0..lines { + cast.push_str(&format!("[{line}.0,\"o\",\"line {line} {}\\r\\n\"]\n", "x".repeat(80))); + } + std::fs::write(dir.join("recording-0.cast"), cast).expect("cast"); + + session_id + } + async fn record(&self, id: Uuid) -> TaskRecord { self.tasks.store().get(id).await.expect("read").expect("record exists") } @@ -288,9 +364,7 @@ async fn ephemeral_task_fails_without_retry_after_a_restart() { assert!(!json.contains(API_KEY), "{json}"); // A restart keeps the database but loses the secrets held in memory. - let restarted = TaskService::open(&harness.db_path, TASK_TIMEOUT) - .await - .expect("task service"); + let restarted = harness.reopen().await; let mut job = TaskJob::read_json(&json, restarted, harness.state.clone()).expect("valid job"); job.run().await.expect("no retry"); @@ -315,7 +389,7 @@ async fn secrets_are_dropped_when_the_task_finishes() { let record = harness.record(snapshot.id).await; assert_eq!(record.state, TaskState::Failed); - assert_eq!(record.error.as_deref(), Some("ai-log task not implemented yet")); + assert_eq!(record.error.as_deref(), Some("session has no recording")); assert!(!record.params.contains(API_KEY), "{}", record.params); } @@ -387,3 +461,177 @@ async fn snapshot_reflects_the_record() { assert!(harness.tasks.get(Uuid::new_v4()).await.expect("read").is_none()); } + +/// A mock OpenAI-compatible provider: `status` picks the answer of the n-th request, 0 being a valid answer. +async fn spawn_ai_provider( + status: impl Fn(usize) -> u16 + Send + Sync + 'static, +) -> (serde_json::Value, Arc>>) { + let requests = Arc::new(Mutex::new(Vec::new())); + let status = Arc::new(status); + + let app = axum::Router::new().fallback(axum::routing::post({ + let requests = Arc::clone(&requests); + move |body: String| { + let requests = Arc::clone(&requests); + let status = Arc::clone(&status); + async move { + let index = { + let mut requests = requests.lock(); + requests.push(body); + requests.len() - 1 + }; + + let answer = serde_json::json!({ + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 0, + "model": "gpt-test", + "choices": [{ + "index": 0, + "message": { "role": "assistant", "content": format!("{{\"offsetSeconds\":{index},\"description\":\"Step {index}\"}}") }, + "finish_reason": "stop" + }], + "usage": { "prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30 } + }); + + match status(index) { + 0 => (StatusCode::OK, axum::Json(answer)), + code => ( + StatusCode::from_u16(code).expect("valid status"), + axum::Json(serde_json::json!({ "error": { "message": "scripted failure" } })), + ), + } + } + } + })); + + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.expect("bind"); + let addr = listener.local_addr().expect("address"); + tokio::spawn(async move { axum::serve(listener, app).await.expect("serve") }); + + let body = serde_json::json!({ + "provider": "openai-compatible", + "model": "gpt-test", + "apiKey": API_KEY, + "baseUrl": format!("http://{addr}/v1/"), + }); + + (body, requests) +} + +fn first_line(path: &Utf8Path) -> String { + std::fs::read_to_string(path) + .expect("chunk") + .lines() + .next() + .expect("line") + .to_owned() +} + +#[tokio::test] +async fn ai_log_retry_resumes_at_the_first_chunk_without_a_checkpoint() { + let harness = Harness::new().await; + // About 1 MB of transcript: three chunks. + let session_id = harness.write_long_session(10_000); + let (body, requests) = spawn_ai_provider(|index| if index == 1 { 503 } else { 0 }).await; + let snapshot = harness.start_ai_log_with(session_id, body).await; + let workspace = harness.workspace(snapshot.id); + let mut job = harness.queued_job(snapshot.id).await; + + assert!(job.run().await.is_err(), "a provider outage is retried"); + + assert_eq!(harness.record(snapshot.id).await.state, TaskState::NotStarted); + assert!(workspace.starts_with(&harness.recording_path)); + assert_eq!( + std::fs::read_to_string(workspace.join("chunks.done")).expect("chunks"), + "3" + ); + assert!(workspace.join("chunk-0000.actions.jsonl").exists()); + assert!(!workspace.join("chunk-0001.actions.jsonl").exists()); + let chunk_starts = (0..3) + .map(|index| first_line(&workspace.join(format!("chunk-{index:04}.txt")))) + .collect::>(); + assert_eq!(requests.lock().len(), 2); + + job.run().await.expect("second attempt succeeds"); + + let record = harness.record(snapshot.id).await; + assert_eq!(record.state, TaskState::Success, "{:?}", record.error); + assert_eq!(record.attempts, 2); + + let requests = requests.lock().clone(); + assert_eq!(requests.len(), 4, "only chunks 1 and 2 are sent again"); + for (request, chunk) in requests.iter().zip([0, 1, 1, 2]) { + assert!(request.contains(&chunk_starts[chunk]), "request for chunk {chunk}"); + } + + let log = std::fs::read_to_string( + harness + .recording_path + .join(session_id.to_string()) + .join("ai-analysis-0.slog"), + ) + .expect("log"); + let descriptions = log + .lines() + .map(|line| serde_json::from_str::(line).expect("JSON")["description"].clone()) + .collect::>(); + assert_eq!( + descriptions, + ["Session started", "Step 0", "Step 2", "Step 3", "Session ended"].map(serde_json::Value::from) + ); + + assert!(!workspace.exists(), "the workspace is deleted on success"); +} + +#[tokio::test] +async fn ai_log_workspace_is_deleted_when_the_task_fails() { + let harness = Harness::new().await; + let session_id = harness.write_long_session(10); + let (body, requests) = spawn_ai_provider(|_| 401).await; + let snapshot = harness.start_ai_log_with(session_id, body).await; + + harness.queued_job(snapshot.id).await.run().await.expect("no retry"); + + let record = harness.record(snapshot.id).await; + assert_eq!(record.state, TaskState::Failed); + assert_eq!(requests.lock().len(), 1); + assert!(!harness.workspace(snapshot.id).exists()); + assert!( + !harness + .recording_path + .join(WORKSPACES_DIR) + .read_dir() + .expect("workspaces") + .any(|_| true) + ); +} + +#[tokio::test] +async fn stale_workspaces_are_removed_at_startup() { + let harness = Harness::new().await; + let unfinished = harness.start_ai_log().await; + let finished = harness.start_scripted::<5>(Outcome::Succeed).await; + harness.run_scripted::<5>(&finished).await.expect("run"); + + let workspaces = harness.recording_path.join(WORKSPACES_DIR); + for name in [ + unfinished.id.to_string(), + finished.def.task_id.to_string(), + Uuid::new_v4().to_string(), + "not-a-task".to_owned(), + ] { + std::fs::create_dir_all(workspaces.join(&name)).expect("workspace"); + std::fs::write(workspaces.join(&name).join("chunk-0000.txt"), "[0.0] x\n").expect("chunk"); + } + std::fs::write(workspaces.join("stray-file"), "x").expect("file"); + + harness.reopen().await; + + let left = workspaces + .read_dir() + .expect("workspaces") + .map(|entry| entry.expect("entry").file_name().into_string().expect("UTF-8")) + .collect::>(); + assert_eq!(left, [unfinished.id.to_string()]); +} diff --git a/devolutions-gateway/tests/tasks.rs b/devolutions-gateway/tests/tasks.rs index cc993e51a..de5cc1afa 100644 --- a/devolutions-gateway/tests/tasks.rs +++ b/devolutions-gateway/tests/tasks.rs @@ -12,6 +12,7 @@ use axum::body::Body; use axum::extract::connect_info::MockConnectInfo; use axum::http::{self, Request, StatusCode}; use base64::Engine as _; +use devolutions_gateway::recording::RecordingManagerTask; use devolutions_gateway::tasks::{SECRETS_LOST_ERROR, TaskRunnerTask, TaskService}; use devolutions_gateway::{DgwState, MockHandles}; use devolutions_gateway_task::{ChildTask, ShutdownHandle, Task as _}; @@ -37,6 +38,7 @@ fn config(dir: &Path, enable_unstable: bool) -> String { ], "Proxy": { "Mode": "Off" }, "ProvisionerTasksDatabase": tasks_db(dir), + "RecordingPath": dir.join("recordings"), "__debug__": { "disable_token_validation": true, "enable_unstable": enable_unstable @@ -78,6 +80,13 @@ impl Gateway { // The auth middleware asks the session manager about any token carrying `jet_aid`; nothing answers in the mock. drop(session_manager_rx); + let recording_manager = RecordingManagerTask::new( + recording_manager_rx, + state.conf_handle.get_conf().recording_path.clone(), + state.sessions.clone(), + state.job_queue_handle.clone(), + ); + state.tasks = TaskService::open_if_enabled(&state.conf_handle.get_conf()).await?; let (shutdown_handle, shutdown_signal) = ShutdownHandle::new(); @@ -85,9 +94,11 @@ impl Gateway { if let (Some(tasks), Jobs::Run) = (state.tasks.clone(), jobs) { let runner = TaskRunnerTask::new(tasks, state.clone()); - job_tasks.push(ChildTask::spawn(runner.run(shutdown_signal))); + job_tasks.push(ChildTask::spawn(runner.run(shutdown_signal.clone()))); } + job_tasks.push(ChildTask::spawn(recording_manager.run(shutdown_signal))); + let app = devolutions_gateway::make_http_service(state) .layer(MockConnectInfo(SocketAddr::from(([0, 0, 0, 0], 3000)))); @@ -95,13 +106,7 @@ impl Gateway { app, shutdown_handle, job_tasks, - _mock_handles: Box::new(( - recording_manager_rx, - subscriber_rx, - job_queue_rx, - traffic_audit_rx, - mock_shutdown_handle, - )), + _mock_handles: Box::new((subscriber_rx, job_queue_rx, traffic_audit_rx, mock_shutdown_handle)), }) } @@ -194,11 +199,15 @@ fn now() -> i64 { } fn task_token() -> String { + session_task_token(Uuid::new_v4()) +} + +fn session_task_token(session_id: Uuid) -> String { unsigned_jws( "TASK", &json!({ "jet_tk": "ai-log", - "jet_aid": Uuid::new_v4(), + "jet_aid": session_id, "nbf": now(), "exp": now() + 600, "jti": Uuid::new_v4(), @@ -308,7 +317,7 @@ fn capture_logs() -> (CapturedLogs, impl Sized) { } #[tokio::test] -async fn ai_log_task_is_accepted_then_fails_as_not_implemented() { +async fn ai_log_task_without_recording_fails() { let dir = tempfile::tempdir().unwrap(); let gateway = Gateway::start(dir.path(), Jobs::Run).await.unwrap(); let app = gateway.app.clone(); @@ -326,7 +335,7 @@ async fn ai_log_task_is_accepted_then_fails_as_not_implemented() { "id": id, "kind": "ai-log", "state": "failed", - "error": "ai-log task not implemented yet", + "error": "session has no recording", }) ); } @@ -548,3 +557,268 @@ async fn task_records_survive_a_restart() { after.stop().await; } + +const SESSION_START: i64 = 1_787_255_035; + +const CAST: &str = r#"{"version": 2, "width": 80, "height": 24} +[0.5,"o","user@host:~$ "] +[1.5,"o","whoami\r\n"] +[1.6,"o","user\r\n"] +"#; + +const AI_ANSWER: &str = + r#"{"offsetSeconds":0.5,"description":"Checked the current user","parameters":{"Command":"whoami"}}"#; + +/// Writes a finished session to the recording folder of the Gateway working in `dir`. +fn write_session(dir: &Path, files: &[(&str, &str)]) -> (Uuid, PathBuf) { + let session_id = Uuid::new_v4(); + let session_dir = dir.join("recordings").join(session_id.to_string()); + std::fs::create_dir_all(&session_dir).unwrap(); + + let manifest_files = files + .iter() + .map(|(name, _)| json!({ "fileName": name, "startTime": SESSION_START, "duration": 10 })) + .collect::>(); + + let manifest = json!({ + "sessionId": session_id, + "startTime": SESSION_START, + "duration": 10, + "files": manifest_files, + }); + std::fs::write(session_dir.join("recording.json"), manifest.to_string()).unwrap(); + + for (name, contents) in files { + std::fs::write(session_dir.join(name), contents).unwrap(); + } + + (session_id, session_dir) +} + +fn read_manifest(session_dir: &Path) -> Value { + serde_json::from_slice(&std::fs::read(session_dir.join("recording.json")).unwrap()).unwrap() +} + +/// A mock OpenAI-compatible provider that answers every request with the same action and keeps the request bodies. +async fn spawn_ai_provider() -> (String, Arc>>) { + spawn_truncating_ai_provider(|_| false).await +} + +/// Like [`spawn_ai_provider`], but answers as cut at the token limit when `truncated` says so for the request body. +async fn spawn_truncating_ai_provider( + truncated: impl Fn(&str) -> bool + Clone + Send + Sync + 'static, +) -> (String, Arc>>) { + let requests = Arc::new(Mutex::new(Vec::new())); + + let app = Router::new().fallback(axum::routing::post({ + let requests = Arc::clone(&requests); + move |body: String| { + let requests = Arc::clone(&requests); + let finish_reason = if truncated(&body) { "length" } else { "stop" }; + async move { + requests + .lock() + .unwrap() + .push(serde_json::from_str::(&body).unwrap()); + axum::Json(json!({ + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 0, + "model": "gpt-test", + "choices": [{ + "index": 0, + "message": { "role": "assistant", "content": AI_ANSWER }, + "finish_reason": finish_reason + }], + "usage": { "prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30 } + })) + } + } + })); + + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + + (format!("http://{addr}/v1/"), requests) +} + +async fn run_ai_log(app: &Router, session_id: Uuid, base_url: &str) -> Value { + let params = json!({ + "provider": "openai-compatible", + "model": "gpt-test", + "apiKey": API_KEY, + "baseUrl": base_url, + }); + + let (status, body) = send(app, start_request(Some(&session_task_token(session_id)), ¶ms)).await; + assert_eq!(status, StatusCode::ACCEPTED, "{body}"); + + let id = serde_json::from_str::(&body).unwrap()["id"] + .as_str() + .unwrap() + .parse::() + .unwrap(); + + wait_until_finished(app, id).await +} + +#[tokio::test] +async fn ai_log_task_appends_a_generated_log_to_the_session() { + let dir = tempfile::tempdir().unwrap(); + let (session_id, session_dir) = write_session(dir.path(), &[("recording-0.cast", CAST)]); + let files_before = read_manifest(&session_dir)["files"].clone(); + let (base_url, requests) = spawn_ai_provider().await; + let gateway = Gateway::start(dir.path(), Jobs::Run).await.unwrap(); + + let first = run_ai_log(&gateway.app, session_id, &base_url).await; + assert_eq!(first["state"], "success", "{first}"); + assert_eq!(first["result"], json!({ "fileName": "ai-analysis-0.slog" })); + + let requests_so_far = requests.lock().unwrap().clone(); + assert_eq!(requests_so_far.len(), 1); + let sent = requests_so_far[0].to_string(); + assert!(sent.contains(r"[0.5] user@host:~$ whoami\n[1.6] user\n"), "{sent}"); + + let expected = [ + r#"{"timestamp":"2026-08-20T19:43:55.000Z","seq":0,"event":"session.start","description":"Session started","source":"ai","model":"gpt-test","promptVersion":"session-actions-1"}"#, + r#"{"timestamp":"2026-08-20T19:43:55.500Z","seq":1,"event":"session.action","description":"Checked the current user","parameters":{"Command":"whoami"}}"#, + r#"{"timestamp":"2026-08-20T19:44:05.000Z","seq":2,"event":"session.end","description":"Session ended"}"#, + ] + .map(|line| format!("{line}\n")) + .concat(); + assert_eq!( + std::fs::read_to_string(session_dir.join("ai-analysis-0.slog")).unwrap(), + expected + ); + + let manifest = read_manifest(&session_dir); + assert_eq!( + manifest["artifacts"], + json!({ "ai-analysis": [{ "fileName": "ai-analysis-0.slog" }] }) + ); + assert_eq!(manifest["files"], files_before); + + let second = run_ai_log(&gateway.app, session_id, &base_url).await; + assert_eq!( + second["result"], + json!({ "fileName": "ai-analysis-1.slog" }), + "{second}" + ); + assert_eq!( + std::fs::read_to_string(session_dir.join("ai-analysis-1.slog")).unwrap(), + expected + ); + + let manifest = read_manifest(&session_dir); + assert_eq!( + manifest["artifacts"], + json!({ "ai-analysis": [{ "fileName": "ai-analysis-0.slog" }, { "fileName": "ai-analysis-1.slog" }] }) + ); + assert_eq!(manifest["files"], files_before); + + wait_for_queued_jobs(dir.path(), 0).await; + gateway.stop().await; + + assert_database_holds_settings_but_not_the_key(dir.path()); + assert_no_file_holds_the_key(dir.path()); + assert_workspaces_are_gone(dir.path()); +} + +fn assert_no_file_holds_the_key(dir: &Path) { + for entry in std::fs::read_dir(dir).unwrap() { + let path = entry.unwrap().path(); + if path.is_dir() { + assert_no_file_holds_the_key(&path); + } else { + assert!(!contains(&std::fs::read(&path).unwrap(), API_KEY), "{}", path.display()); + } + } +} + +fn assert_workspaces_are_gone(dir: &Path) { + let workspaces = dir.join("recordings").join(devolutions_gateway::tasks::WORKSPACES_DIR); + let left = std::fs::read_dir(&workspaces).map_or(0, |entries| entries.count()); + assert_eq!(left, 0, "{}", workspaces.display()); +} + +/// A finished session whose transcript is `lines` lines of about 100 bytes, marked `line `. +fn write_long_session(dir: &Path, lines: usize) -> (Uuid, PathBuf) { + let mut cast = String::from("{\"version\": 2}\n"); + for line in 0..lines { + cast.push_str(&format!("[{line}.0,\"o\",\"line {line} {}\\r\\n\"]\n", "x".repeat(80))); + } + + write_session(dir, &[("recording-0.cast", &cast)]) +} + +#[tokio::test] +async fn ai_log_task_splits_a_chunk_whose_answer_is_truncated() { + let dir = tempfile::tempdir().unwrap(); + let (session_id, session_dir) = write_long_session(dir.path(), 60); + let (base_url, requests) = + spawn_truncating_ai_provider(|body| body.contains("] line 0 ") && body.contains("] line 59 ")).await; + let gateway = Gateway::start(dir.path(), Jobs::Run).await.unwrap(); + + let finished = run_ai_log(&gateway.app, session_id, &base_url).await; + assert_eq!(finished["state"], "success", "{finished}"); + + let requests = requests + .lock() + .unwrap() + .iter() + .map(Value::to_string) + .collect::>(); + assert_eq!(requests.len(), 3, "the whole chunk, then its two halves"); + assert!(requests[1].contains("] line 0 ") && !requests[1].contains("] line 59 ")); + assert!(!requests[2].contains("] line 0 ") && requests[2].contains("] line 59 ")); + + let log = std::fs::read_to_string(session_dir.join("ai-analysis-0.slog")).unwrap(); + assert_eq!(log.lines().filter(|line| line.contains("session.action")).count(), 2); + + wait_for_queued_jobs(dir.path(), 0).await; + gateway.stop().await; + assert_workspaces_are_gone(dir.path()); +} + +#[tokio::test] +async fn ai_log_task_fails_when_even_a_short_part_is_truncated() { + let dir = tempfile::tempdir().unwrap(); + let (session_id, session_dir) = write_long_session(dir.path(), 60); + let (base_url, requests) = spawn_truncating_ai_provider(|_| true).await; + let gateway = Gateway::start(dir.path(), Jobs::Run).await.unwrap(); + + let finished = run_ai_log(&gateway.app, session_id, &base_url).await; + + assert_eq!(finished["state"], "failed"); + assert_eq!(finished["error"], devolutions_gateway::tasks::ai_log::TRUNCATED_ERROR); + assert_eq!( + requests.lock().unwrap().len(), + 3, + "6 KB, then 3 KB, then 1.5 KB, below the split floor" + ); + assert!(read_manifest(&session_dir).get("artifacts").is_none()); + + wait_for_queued_jobs(dir.path(), 0).await; + gateway.stop().await; + assert_workspaces_are_gone(dir.path()); +} + +#[tokio::test] +async fn ai_log_task_fails_for_a_video_only_session() { + let dir = tempfile::tempdir().unwrap(); + let (session_id, session_dir) = write_session(dir.path(), &[("recording-0.webm", "not a video")]); + let (base_url, requests) = spawn_ai_provider().await; + let gateway = Gateway::start(dir.path(), Jobs::Run).await.unwrap(); + + let finished = run_ai_log(&gateway.app, session_id, &base_url).await; + + assert_eq!(finished["state"], "failed"); + assert_eq!(finished["error"], "unsupported recording type: webm"); + assert!(requests.lock().unwrap().is_empty()); + assert!(read_manifest(&session_dir).get("artifacts").is_none()); + + // A permanent failure is not retried. + wait_for_queued_jobs(dir.path(), 0).await; + gateway.stop().await; +} From 8f73667ec4648ed054c4b02436aaf3948e0e469a Mon Sep 17 00:00:00 2001 From: Junyi Ou Date: Wed, 30 Sep 2026 18:37:50 -0400 Subject: [PATCH 3/3] feat(dgw): record the AI model and token usage of a session log The AI crate now reports the model that answered and the tokens each request used, but the ai-log task dropped both. Task records are kept for audit, and AI budgets will need the token counts. Each chunk checkpoint now holds the reported model and the token usage next to the actions, so a retry keeps the counts of the chunks it does not ask again. Answers cut at the output limit and asked again in two halves count too. The task result adds the model and the total usage, which is absent when the provider did not report it for every request. The .slog names the reported model, usually a dated version of the requested one, and falls back to the requested model. --- devolutions-gateway/src/tasks/ai.rs | 30 +++- .../src/tasks/ai_log/checkpoint.rs | 140 ++++++++++----- devolutions-gateway/src/tasks/ai_log/mod.rs | 161 +++++++++++++++--- devolutions-gateway/src/tasks/tests.rs | 11 +- devolutions-gateway/tests/tasks.rs | 29 +++- 5 files changed, 295 insertions(+), 76 deletions(-) diff --git a/devolutions-gateway/src/tasks/ai.rs b/devolutions-gateway/src/tasks/ai.rs index 09e50d07e..975712e47 100644 --- a/devolutions-gateway/src/tasks/ai.rs +++ b/devolutions-gateway/src/tasks/ai.rs @@ -3,7 +3,7 @@ //! The provisioner sends them with the API key in the body of `POST /jet/tasks`. //! [`AiSettings`] is the part persisted with the task; the API key stays in memory as the task secret. -use devolutions_gateway_ai::{AiClient, BuildError, Provider}; +use devolutions_gateway_ai::{AiClient, BuildError, Provider, Usage}; use secrecy::SecretString; use url::Url; @@ -111,6 +111,34 @@ impl From for TaskError { } } +/// Tokens counted by the AI provider, as stored in task checkpoints and results. +/// +/// It mirrors [`Usage`], so that a change in the AI crate never changes what the task records hold. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct TokenUsage { + pub input_tokens: u64, + pub output_tokens: u64, +} + +impl From for TokenUsage { + fn from(usage: Usage) -> Self { + Self { + input_tokens: usage.input_tokens, + output_tokens: usage.output_tokens, + } + } +} + +impl From for Usage { + fn from(usage: TokenUsage) -> Self { + Self { + input_tokens: usage.input_tokens, + output_tokens: usage.output_tokens, + } + } +} + enum ClientError { Build(BuildError), HttpClient(reqwest::Error), diff --git a/devolutions-gateway/src/tasks/ai_log/checkpoint.rs b/devolutions-gateway/src/tasks/ai_log/checkpoint.rs index af89bd00e..18ef5f47f 100644 --- a/devolutions-gateway/src/tasks/ai_log/checkpoint.rs +++ b/devolutions-gateway/src/tasks/ai_log/checkpoint.rs @@ -1,15 +1,38 @@ -//! Actions found in one transcript chunk, saved in the task workspace so a retry does not ask the AI again. +//! What the AI answered for one transcript chunk, saved in the task workspace so a retry does not ask the AI again. use std::collections::BTreeMap; -use std::io::{BufRead as _, BufReader, BufWriter, Write as _}; +use std::io::{BufReader, BufWriter}; use std::time::Duration; use anyhow::Context as _; use camino::{Utf8Path, Utf8PathBuf}; +use devolutions_gateway_ai::Usage; use devolutions_gateway_ai::session_actions::Action; +use crate::tasks::ai::TokenUsage; + pub(crate) fn path(workspace: &Utf8Path, index: usize) -> Utf8PathBuf { - workspace.join(format!("chunk-{index:04}.actions.jsonl")) + workspace.join(format!("chunk-{index:04}.json")) +} + +/// Actions found in one chunk, with what the AI provider reported for the requests about it. +#[derive(Debug, Clone, PartialEq)] +pub(crate) struct DescribedChunk { + pub(crate) actions: Vec, + /// Model that answered, as reported by the provider. + pub(crate) model: Option, + /// Tokens of every request about the chunk, answers cut and asked again included; `None` when one was not reported. + pub(crate) usage: Option, +} + +#[derive(Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +struct SavedChunk { + #[serde(default, skip_serializing_if = "Option::is_none")] + model: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + usage: Option, + actions: Vec, } #[derive(Serialize, Deserialize)] @@ -24,34 +47,40 @@ struct SavedAction { } /// Writes the checkpoint through a temporary file, so a crash never leaves a partial one. -pub(crate) fn write(path: &Utf8Path, actions: &[Action]) -> anyhow::Result<()> { +pub(crate) fn write(path: &Utf8Path, chunk: &DescribedChunk) -> anyhow::Result<()> { let partial = path.with_extension("partial"); - let mut out = BufWriter::new(std::fs::File::create(&partial).with_context(|| format!("create {partial}"))?); - - for action in actions { - let saved = SavedAction { - offset_seconds: action.offset.as_secs_f64(), - description: action.description.clone(), - object: action.object.clone(), - parameters: action.parameters.clone(), - }; - serde_json::to_writer(&mut out, &saved)?; - out.write_all(b"\n")?; - } + let saved = SavedChunk { + model: chunk.model.clone(), + usage: chunk.usage.map(TokenUsage::from), + actions: chunk + .actions + .iter() + .map(|action| SavedAction { + offset_seconds: action.offset.as_secs_f64(), + description: action.description.clone(), + object: action.object.clone(), + parameters: action.parameters.clone(), + }) + .collect(), + }; + let mut out = BufWriter::new(std::fs::File::create(&partial).with_context(|| format!("create {partial}"))?); + serde_json::to_writer(&mut out, &saved)?; out.into_inner()?.sync_all()?; std::fs::rename(&partial, path).with_context(|| format!("rename {partial}"))?; Ok(()) } -pub(crate) fn read(path: &Utf8Path) -> anyhow::Result> { +pub(crate) fn read(path: &Utf8Path) -> anyhow::Result { let file = BufReader::new(std::fs::File::open(path).with_context(|| format!("open {path}"))?); + let saved: SavedChunk = serde_json::from_reader(file).with_context(|| format!("read {path}"))?; - file.lines() - .map(|line| { - let saved: SavedAction = serde_json::from_str(&line?)?; + let actions = saved + .actions + .into_iter() + .map(|saved| { Ok(Action { offset: Duration::try_from_secs_f64(saved.offset_seconds)?, description: saved.description, @@ -59,7 +88,13 @@ pub(crate) fn read(path: &Utf8Path) -> anyhow::Result> { parameters: saved.parameters, }) }) - .collect() + .collect::>()?; + + Ok(DescribedChunk { + actions, + model: saved.model, + usage: saved.usage.map(Usage::from), + }) } #[cfg(test)] @@ -67,28 +102,51 @@ mod tests { use super::*; #[test] - fn actions_round_trip() { + fn described_chunk_round_trips() { let dir = tempfile::tempdir().expect("temp dir"); let path = path(Utf8Path::from_path(dir.path()).expect("UTF-8"), 3); - let actions = vec![ - Action { - offset: Duration::from_millis(1500), - description: "Listed files".to_owned(), - object: Some("/var/log".to_owned()), - parameters: BTreeMap::from([("Command".to_owned(), "ls".to_owned())]), - }, - Action { - offset: Duration::from_secs(62), - description: "Closed the shell".to_owned(), - object: None, - parameters: BTreeMap::new(), - }, - ]; - - write(&path, &actions).expect("written"); - - assert!(path.as_str().ends_with("chunk-0003.actions.jsonl")); - assert_eq!(read(&path).expect("read"), actions); + let chunk = DescribedChunk { + actions: vec![ + Action { + offset: Duration::from_millis(1500), + description: "Listed files".to_owned(), + object: Some("/var/log".to_owned()), + parameters: BTreeMap::from([("Command".to_owned(), "ls".to_owned())]), + }, + Action { + offset: Duration::from_secs(62), + description: "Closed the shell".to_owned(), + object: None, + parameters: BTreeMap::new(), + }, + ], + model: Some("gpt-test-2026-09-30".to_owned()), + usage: Some(Usage { + input_tokens: 10, + output_tokens: 20, + }), + }; + + write(&path, &chunk).expect("written"); + + assert!(path.as_str().ends_with("chunk-0003.json")); + assert_eq!(read(&path).expect("read"), chunk); assert!(!path.with_extension("partial").exists()); } + + #[test] + fn unreported_model_and_usage_stay_unknown() { + let dir = tempfile::tempdir().expect("temp dir"); + let path = path(Utf8Path::from_path(dir.path()).expect("UTF-8"), 0); + let chunk = DescribedChunk { + actions: Vec::new(), + model: None, + usage: None, + }; + + write(&path, &chunk).expect("written"); + + assert_eq!(std::fs::read_to_string(&path).expect("checkpoint"), r#"{"actions":[]}"#); + assert_eq!(read(&path).expect("read"), chunk); + } } diff --git a/devolutions-gateway/src/tasks/ai_log/mod.rs b/devolutions-gateway/src/tasks/ai_log/mod.rs index a545c57cb..a3056c7df 100644 --- a/devolutions-gateway/src/tasks/ai_log/mod.rs +++ b/devolutions-gateway/src/tasks/ai_log/mod.rs @@ -2,8 +2,8 @@ //! //! The task works in steps, keeping its files in the task workspace: //! 1. stream the terminal recordings into transcript chunk files (`chunk-NNNN.txt`); -//! 2. ask the AI about each chunk and save its actions as a checkpoint (`chunk-NNNN.actions.jsonl`); -//! a retry skips the chunks that already have one; +//! 2. ask the AI about each chunk and save its actions, the reported model and the token usage as a checkpoint +//! (`chunk-NNNN.json`); a retry skips the chunks that already have one; //! 3. merge the checkpoints into a `.slog` file and add it to the session. mod checkpoint; @@ -15,13 +15,13 @@ use std::fs::File; use std::io::BufWriter; use camino::{Utf8Path, Utf8PathBuf}; -use devolutions_gateway_ai::AiClient; -use devolutions_gateway_ai::session_actions::Action; +use devolutions_gateway_ai::{AiClient, Usage}; use secrecy::SecretString; use url::Url; use uuid::Uuid; -use super::ai::{AiProvider, AiSettings}; +use self::checkpoint::DescribedChunk; +use super::ai::{AiProvider, AiSettings, TokenUsage}; use super::{EphemeralTask, RetryPolicy, SECRETS_LOST_ERROR, TaskCtx, TaskError, TaskErrorCode, TaskKind}; use crate::DgwState; use crate::artifacts::ArtifactKind; @@ -87,6 +87,13 @@ pub enum AiLogSubstate { pub struct AiLogOutput { /// Name of the new log in the session manifest, such as `ai-analysis-0.slog`. pub file_name: String, + /// Model that wrote the log, as reported by the AI provider; the requested model when it reports none. + pub model: String, + /// Tokens the AI provider counted for the log, answers cut and asked again included. + /// + /// Absent when the provider did not report them for every request. + #[serde(skip_serializing_if = "Option::is_none")] + pub usage: Option, } pub enum AiLogTask {} @@ -139,9 +146,19 @@ impl TaskKind for AiLogTask { .await .map_err(|error| workspace_error(&error))?; - let actions = describe_chunk(&client, ctx.params.max_output_tokens, &chunk).await?; + let described = describe_chunk(&client, ctx.params.max_output_tokens, &chunk).await?; + + debug!( + session.id = %session_id, + index, + actions = described.actions.len(), + model = ?described.model, + usage = ?described.usage, + "Chunk described" + ); - blocking(move || checkpoint::write(&checkpoint, &actions).map_err(|error| workspace_error(&error))).await?; + blocking(move || checkpoint::write(&checkpoint, &described).map_err(|error| workspace_error(&error))) + .await?; } ctx.progress @@ -149,10 +166,10 @@ impl TaskKind for AiLogTask { .await; let log_path = workspace.join(LOG_FILE); - let model = ctx.params.model.clone(); - let actions = blocking({ + let requested_model = ctx.params.model.clone(); + let merged = blocking({ let log_path = log_path.clone(); - move || merge_checkpoints(&workspace, total, start_time, duration, &model, &log_path) + move || merge_checkpoints(&workspace, total, start_time, duration, &requested_model, &log_path) }) .await?; @@ -160,9 +177,20 @@ impl TaskKind for AiLogTask { .await .map_err(|error| TaskError::Permanent(format!("failed to add the log: {error:#}")))?; - info!(session.id = %session_id, file_name, actions, "Session log generated"); + info!( + session.id = %session_id, + file_name, + actions = merged.actions, + model = merged.model, + usage = ?merged.usage, + "Session log generated" + ); - Ok(AiLogOutput { file_name }) + Ok(AiLogOutput { + file_name, + model: merged.model, + usage: merged.usage.map(TokenUsage::from), + }) } } @@ -214,9 +242,13 @@ async fn describe_chunk( client: &AiClient, max_output_tokens: Option, chunk: &str, -) -> Result, TaskError> { +) -> Result { let mut parts = VecDeque::from([chunk]); - let mut actions = Vec::new(); + let mut described = DescribedChunk { + actions: Vec::new(), + model: None, + usage: Some(Usage::default()), + }; while let Some(part) = parts.pop_front() { let mut request = client.describe_session_actions(part); @@ -226,8 +258,14 @@ async fn describe_chunk( } match request.send().await { - Ok(response) => actions.extend(response.output), - Err(devolutions_gateway_ai::Error::Truncated { .. }) => { + Ok(response) => { + described.actions.extend(response.output); + described.model = described.model.or(response.model); + described.usage = add_usage(described.usage, response.usage); + } + Err(devolutions_gateway_ai::Error::Truncated { usage }) => { + described.usage = add_usage(described.usage, usage); + let Some((first, second)) = split_in_half(part) else { return Err(TaskError::Permanent(TRUNCATED_ERROR.to_owned())); }; @@ -240,7 +278,12 @@ async fn describe_chunk( } } - Ok(actions) + Ok(described) +} + +/// Sums token counts; the total is unknown as soon as one count is. +fn add_usage(total: Option, more: Option) -> Option { + Some(total? + more?) } /// Cuts `part` on the line boundary closest to its middle, unless it is too short to split. @@ -271,32 +314,53 @@ fn split_in_half(part: &str) -> Option<(&str, &str)> { Some(part.split_at(cut)) } -/// Writes the `.slog` from the checkpoints, and returns the number of actions. +/// What [`merge_checkpoints`] wrote. +struct MergedLog { + actions: usize, + model: String, + usage: Option, +} + +/// Writes the `.slog` from the checkpoints. +/// +/// The log names the first model reported for a chunk, or `requested_model` when the provider reported none. fn merge_checkpoints( workspace: &Utf8Path, total: usize, start_time: i64, duration: i64, - model: &str, + requested_model: &str, log_path: &Utf8Path, -) -> Result { +) -> Result { let log_error = |error: anyhow::Error| TaskError::Permanent(format!("failed to write the log: {error:#}")); + let read = |index| checkpoint::read(&checkpoint::path(workspace, index)).map_err(|error| workspace_error(&error)); + + // `session.start` names the model, so it is read before the log is written. + let mut model = None; + for index in 0..total { + model = read(index)?.model; + if model.is_some() { + break; + } + } + let model = model.unwrap_or_else(|| requested_model.to_owned()); let out = File::create(log_path).map_err(|error| workspace_error(&error))?; - let mut log = slog::SlogWriter::start(BufWriter::new(out), start_time, model).map_err(log_error)?; - let mut count = 0; + let mut log = slog::SlogWriter::start(BufWriter::new(out), start_time, &model).map_err(log_error)?; + let mut actions = 0; + let mut usage = Some(Usage::default()); // Chunks follow each other in time, so only the actions of one chunk need sorting. for index in 0..total { - let mut actions = - checkpoint::read(&checkpoint::path(workspace, index)).map_err(|error| workspace_error(&error))?; - actions.sort_by_key(|action| action.offset); + let mut described = read(index)?; + described.actions.sort_by_key(|action| action.offset); - for action in &actions { + for action in &described.actions { log.action(action).map_err(log_error)?; } - count += actions.len(); + actions += described.actions.len(); + usage = add_usage(usage, described.usage); } log.finish(duration) @@ -306,7 +370,7 @@ fn merge_checkpoints( .sync_all() .map_err(|error| workspace_error(&error))?; - Ok(count) + Ok(MergedLog { actions, model, usage }) } impl EphemeralTask for AiLogTask { @@ -440,4 +504,45 @@ mod tests { assert_eq!(error, TaskErrorCode::MissingModel); } + + #[test] + fn usage_total_is_unknown_once_a_count_is() { + let usage = |input_tokens, output_tokens| Usage { + input_tokens, + output_tokens, + }; + + assert_eq!(add_usage(Some(usage(1, 2)), Some(usage(10, 20))), Some(usage(11, 22))); + assert_eq!(add_usage(Some(usage(1, 2)), None), None); + assert_eq!(add_usage(None, Some(usage(10, 20))), None); + } + + #[test] + fn output_holds_the_model_and_the_usage_when_known() { + let output = AiLogOutput { + file_name: "ai-analysis-0.slog".to_owned(), + model: "gpt-test-2026-09-30".to_owned(), + usage: Some(TokenUsage { + input_tokens: 10, + output_tokens: 20, + }), + }; + + assert_eq!( + serde_json::to_value(&output).expect("serializable"), + serde_json::json!({ + "fileName": "ai-analysis-0.slog", + "model": "gpt-test-2026-09-30", + "usage": { "inputTokens": 10, "outputTokens": 20 } + }) + ); + + let unknown = AiLogOutput { usage: None, ..output }; + assert!( + serde_json::to_value(&unknown) + .expect("serializable") + .get("usage") + .is_none() + ); + } } diff --git a/devolutions-gateway/src/tasks/tests.rs b/devolutions-gateway/src/tasks/tests.rs index 3f4e1dbe1..97400c807 100644 --- a/devolutions-gateway/src/tasks/tests.rs +++ b/devolutions-gateway/src/tasks/tests.rs @@ -546,8 +546,8 @@ async fn ai_log_retry_resumes_at_the_first_chunk_without_a_checkpoint() { std::fs::read_to_string(workspace.join("chunks.done")).expect("chunks"), "3" ); - assert!(workspace.join("chunk-0000.actions.jsonl").exists()); - assert!(!workspace.join("chunk-0001.actions.jsonl").exists()); + assert!(workspace.join("chunk-0000.json").exists()); + assert!(!workspace.join("chunk-0001.json").exists()); let chunk_starts = (0..3) .map(|index| first_line(&workspace.join(format!("chunk-{index:04}.txt")))) .collect::>(); @@ -559,6 +559,13 @@ async fn ai_log_retry_resumes_at_the_first_chunk_without_a_checkpoint() { assert_eq!(record.state, TaskState::Success, "{:?}", record.error); assert_eq!(record.attempts, 2); + let result: serde_json::Value = serde_json::from_str(record.result.as_deref().expect("result")).expect("JSON"); + assert_eq!( + result["usage"], + serde_json::json!({ "inputTokens": 30, "outputTokens": 60 }), + "chunk 0 counts from its checkpoint, the failed request not at all" + ); + let requests = requests.lock().clone(); assert_eq!(requests.len(), 4, "only chunks 1 and 2 are sent again"); for (request, chunk) in requests.iter().zip([0, 1, 1, 2]) { diff --git a/devolutions-gateway/tests/tasks.rs b/devolutions-gateway/tests/tasks.rs index de5cc1afa..4d6a5499a 100644 --- a/devolutions-gateway/tests/tasks.rs +++ b/devolutions-gateway/tests/tasks.rs @@ -599,6 +599,18 @@ fn read_manifest(session_dir: &Path) -> Value { serde_json::from_slice(&std::fs::read(session_dir.join("recording.json")).unwrap()).unwrap() } +/// Model the mock AI provider reports, as a real one reports a dated version of the requested model. +const REPORTED_MODEL: &str = "gpt-test-2026-09-30"; + +/// Result of a successful `ai-log` task run against the mock AI provider. +fn expected_result(file_name: &str, input_tokens: u64, output_tokens: u64) -> Value { + json!({ + "fileName": file_name, + "model": REPORTED_MODEL, + "usage": { "inputTokens": input_tokens, "outputTokens": output_tokens } + }) +} + /// A mock OpenAI-compatible provider that answers every request with the same action and keeps the request bodies. async fn spawn_ai_provider() -> (String, Arc>>) { spawn_truncating_ai_provider(|_| false).await @@ -624,7 +636,7 @@ async fn spawn_truncating_ai_provider( "id": "chatcmpl-1", "object": "chat.completion", "created": 0, - "model": "gpt-test", + "model": REPORTED_MODEL, "choices": [{ "index": 0, "message": { "role": "assistant", "content": AI_ANSWER }, @@ -673,15 +685,19 @@ async fn ai_log_task_appends_a_generated_log_to_the_session() { let first = run_ai_log(&gateway.app, session_id, &base_url).await; assert_eq!(first["state"], "success", "{first}"); - assert_eq!(first["result"], json!({ "fileName": "ai-analysis-0.slog" })); + assert_eq!(first["result"], expected_result("ai-analysis-0.slog", 10, 20)); let requests_so_far = requests.lock().unwrap().clone(); assert_eq!(requests_so_far.len(), 1); let sent = requests_so_far[0].to_string(); assert!(sent.contains(r"[0.5] user@host:~$ whoami\n[1.6] user\n"), "{sent}"); + assert_eq!( + requests_so_far[0]["model"], "gpt-test", + "the requested model is sent as is" + ); let expected = [ - r#"{"timestamp":"2026-08-20T19:43:55.000Z","seq":0,"event":"session.start","description":"Session started","source":"ai","model":"gpt-test","promptVersion":"session-actions-1"}"#, + r#"{"timestamp":"2026-08-20T19:43:55.000Z","seq":0,"event":"session.start","description":"Session started","source":"ai","model":"gpt-test-2026-09-30","promptVersion":"session-actions-1"}"#, r#"{"timestamp":"2026-08-20T19:43:55.500Z","seq":1,"event":"session.action","description":"Checked the current user","parameters":{"Command":"whoami"}}"#, r#"{"timestamp":"2026-08-20T19:44:05.000Z","seq":2,"event":"session.end","description":"Session ended"}"#, ] @@ -702,7 +718,7 @@ async fn ai_log_task_appends_a_generated_log_to_the_session() { let second = run_ai_log(&gateway.app, session_id, &base_url).await; assert_eq!( second["result"], - json!({ "fileName": "ai-analysis-1.slog" }), + expected_result("ai-analysis-1.slog", 10, 20), "{second}" ); assert_eq!( @@ -762,6 +778,11 @@ async fn ai_log_task_splits_a_chunk_whose_answer_is_truncated() { let finished = run_ai_log(&gateway.app, session_id, &base_url).await; assert_eq!(finished["state"], "success", "{finished}"); + assert_eq!( + finished["result"], + expected_result("ai-analysis-0.slog", 30, 60), + "the cut answer counts too" + ); let requests = requests .lock()