diff --git a/crates/terminal-streamer/src/trp_decoder.rs b/crates/terminal-streamer/src/trp_decoder.rs index 751ec3d8c..e219afb33 100644 --- a/crates/terminal-streamer/src/trp_decoder.rs +++ b/crates/terminal-streamer/src/trp_decoder.rs @@ -189,3 +189,144 @@ async fn send(sender: &mut tokio::sync::mpsc::Sender>, mu sender.send(Ok(json)).await?; Ok(()) } + +/// 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, +} + +/// 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, +} + +impl TrpOutputReader { + pub fn new(reader: R) -> Self { + Self { reader, time: 0.0 } + } + + 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); + } + + 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))) + } +} + +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)] +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 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"$ ")); + 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 output = TrpOutputReader::new(trp.as_slice()) + .collect::>>() + .expect("valid recording"); + + assert_eq!( + output, + [ + TerminalOutput { + time: 0.5, + text: "$ ".to_owned() + }, + TerminalOutput { + time: 1.75, + text: "ls\r\n".to_owned() + }, + ] + ); + } + + #[test] + fn reads_one_packet_at_a_time() { + struct CountingReader<'a> { + data: &'a [u8], + read: std::rc::Rc>, + } + + 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.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.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..18ef5f47f --- /dev/null +++ b/devolutions-gateway/src/tasks/ai_log/checkpoint.rs @@ -0,0 +1,152 @@ +//! 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::{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}.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)] +#[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, chunk: &DescribedChunk) -> anyhow::Result<()> { + let partial = path.with_extension("partial"); + + 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 { + 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}"))?; + + let actions = saved + .actions + .into_iter() + .map(|saved| { + Ok(Action { + offset: Duration::try_from_secs_f64(saved.offset_seconds)?, + description: saved.description, + object: saved.object, + parameters: saved.parameters, + }) + }) + .collect::>()?; + + Ok(DescribedChunk { + actions, + model: saved.model, + usage: saved.usage.map(Usage::from), + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + 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 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 new file mode 100644 index 000000000..a3056c7df --- /dev/null +++ b/devolutions-gateway/src/tasks/ai_log/mod.rs @@ -0,0 +1,548 @@ +//! `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, 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; +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, Usage}; +use secrecy::SecretString; +use url::Url; +use uuid::Uuid; + +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; +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, + /// 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 {} + +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 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, &described).map_err(|error| workspace_error(&error))) + .await?; + } + + ctx.progress + .set(&AiLogSubstate::Describing { done: total, total }) + .await; + + let log_path = workspace.join(LOG_FILE); + let requested_model = ctx.params.model.clone(); + let merged = blocking({ + let log_path = log_path.clone(); + move || merge_checkpoints(&workspace, total, start_time, duration, &requested_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 = merged.actions, + model = merged.model, + usage = ?merged.usage, + "Session log generated" + ); + + Ok(AiLogOutput { + file_name, + model: merged.model, + usage: merged.usage.map(TokenUsage::from), + }) + } +} + +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 { + let mut parts = VecDeque::from([chunk]); + 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); + + if let Some(max_output_tokens) = max_output_tokens { + request = request.max_output_tokens(max_output_tokens); + } + + match request.send().await { + 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())); + }; + + 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(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. +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)) +} + +/// 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, + requested_model: &str, + log_path: &Utf8Path, +) -> 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 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 described = read(index)?; + described.actions.sort_by_key(|action| action.offset); + + for action in &described.actions { + log.action(action).map_err(log_error)?; + } + + actions += described.actions.len(); + usage = add_usage(usage, described.usage); + } + + 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(MergedLog { actions, model, usage }) +} + +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); + } + + #[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/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..97400c807 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,184 @@ 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.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::>(); + 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 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]) { + 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..4d6a5499a 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,289 @@ 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() +} + +/// 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 +} + +/// 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": REPORTED_MODEL, + "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"], 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-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"}"#, + ] + .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"], + expected_result("ai-analysis-1.slog", 10, 20), + "{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}"); + assert_eq!( + finished["result"], + expected_result("ai-analysis-0.slog", 30, 60), + "the cut answer counts too" + ); + + 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; +}