Compare commits
213
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
15bc299987 | ||
|
|
87bdb0c6ed | ||
|
|
aa96bbedbc | ||
|
|
4df277c81f | ||
|
|
12646b6ca0 | ||
|
|
1e674d70c2 | ||
|
|
5ee77698db | ||
|
|
d1f5661881 | ||
|
|
532d078720 | ||
|
|
27e5df106f | ||
|
|
33d98868c3 | ||
|
|
fb13e53cb5 | ||
|
|
60a5495ccd | ||
|
|
f1dc90621c | ||
|
|
eecb116709 | ||
|
|
783d25b1c4 | ||
|
|
af06eecfd0 | ||
|
|
9bd08a3a5b | ||
|
|
3d66247e11 | ||
|
|
4390554477 | ||
|
|
74457db4eb | ||
|
|
5d61da481b | ||
|
|
89856eb7c3 | ||
|
|
64c268582d | ||
|
|
bb6558e7bf | ||
|
|
42d109cae3 | ||
|
|
9b48b1ff5d | ||
|
|
4ca8ea1694 | ||
|
|
c10d6c6914 | ||
|
|
3a94c845cf | ||
|
|
bc810beb3b | ||
|
|
eac4a0c071 | ||
|
|
c1dfb1add5 | ||
|
|
e67f9bee08 | ||
|
|
7c1d81cee9 | ||
|
|
5cec2eef60 | ||
|
|
e62c7cf4f5 | ||
|
|
0245980ea5 | ||
|
|
7abc6aca45 | ||
|
|
f1bcd41ad9 | ||
|
|
8022128993 | ||
|
|
68b1aa64e9 | ||
|
|
f3af8f21dc | ||
|
|
7c056b1db8 | ||
|
|
da14c82f71 | ||
|
|
1f32c693df | ||
|
|
74bfbe941e | ||
|
|
2884c08466 | ||
|
|
85e1ea320a | ||
|
|
7be428d8bf | ||
|
|
56798f9fb4 | ||
|
|
6b1b8a8846 | ||
|
|
9fb1b90856 | ||
|
|
30d4023475 | ||
|
|
70432f3d12 | ||
|
|
7ee6c307fc | ||
|
|
36cfbbe6d2 | ||
|
|
a2e1a3d939 | ||
|
|
9dc8d9a77a | ||
|
|
3f6bb65eb1 | ||
|
|
8a70f3cb26 | ||
|
|
6c5b8315a3 | ||
|
|
690ed0f121 | ||
|
|
fd60c2b8be | ||
|
|
e27b4feb25 | ||
|
|
09a33e7283 | ||
|
|
d87441448e | ||
|
|
f783f10f6e | ||
|
|
70bdb2d723 | ||
|
|
4c1ef04378 | ||
|
|
a595af133c | ||
|
|
f74f3cd133 | ||
|
|
bcd4848458 | ||
|
|
fc05bf9711 | ||
|
|
63ad590262 | ||
|
|
0fd1193b6b | ||
|
|
d996822957 | ||
|
|
96349721cb | ||
|
|
8344921b65 | ||
|
|
c97b3b7b77 | ||
|
|
e00e675ed1 | ||
|
|
538da1f2b2 | ||
|
|
bad37ddc7d | ||
|
|
175eda9f29 | ||
|
|
14c806d38f | ||
|
|
925100fb82 | ||
|
|
b29b003ea3 | ||
|
|
faa727965b | ||
|
|
d2ffbf2c40 | ||
|
|
5418fad7d7 | ||
|
|
510795f1c5 | ||
|
|
9e0d499987 | ||
|
|
e96fde0632 | ||
|
|
eea79dead4 | ||
|
|
c4a3f4ba1e | ||
|
|
4a4a01b730 | ||
|
|
a664e72488 | ||
|
|
21317123a4 | ||
|
|
409245cb52 | ||
|
|
1aeb6fdb35 | ||
|
|
816fa96e07 | ||
|
|
8fbe4218c6 | ||
|
|
00c8df0fc9 | ||
|
|
c52c7ead19 | ||
|
|
d1e8a827c2 | ||
|
|
12d96fb03d | ||
|
|
171a191873 | ||
|
|
47dabd8793 | ||
|
|
2bb661f1cf | ||
|
|
04e296a4ef | ||
|
|
4bba227af5 | ||
|
|
5cc78d63c6 | ||
|
|
323f5dc09c | ||
|
|
fb97edfe95 | ||
|
|
070f62ef12 | ||
|
|
981749aa3d | ||
|
|
e01b46b30a | ||
|
|
a1b659c45d | ||
|
|
37a012ef92 | ||
|
|
4927e8a843 | ||
|
|
025d6ddb47 | ||
|
|
01a4dfd5d3 | ||
|
|
1d7158a0bf | ||
|
|
9de2afbfc6 | ||
|
|
9013754a3a | ||
|
|
21eea0b104 | ||
|
|
15e8d7365c | ||
|
|
88e3bf7065 | ||
|
|
2765138bf3 | ||
|
|
6b20ceac46 | ||
|
|
d748274905 | ||
|
|
e1578217d5 | ||
|
|
3481682cb4 | ||
|
|
996b7f2468 | ||
|
|
879993b9b1 | ||
|
|
6604154e3f | ||
|
|
a9ad42a970 | ||
|
|
8b3d1302c6 | ||
|
|
ac9269d6ce | ||
|
|
8ffb716817 | ||
|
|
23f671fa48 | ||
|
|
d7cdcde443 | ||
|
|
95a81faf63 | ||
|
|
bb8bb6d099 | ||
|
|
310801a29b | ||
|
|
7d09b20445 | ||
|
|
456a06f194 | ||
|
|
62eaefb1fa | ||
|
|
10264b4019 | ||
|
|
ab9765d91d | ||
|
|
e7f4c6864f | ||
|
|
bde1dea2a5 | ||
|
|
13a021c480 | ||
|
|
a7f09fad98 | ||
|
|
10eaf4a5fb | ||
|
|
bcada300e3 | ||
|
|
9756174676 | ||
|
|
6dd8461a46 | ||
|
|
12cc2eb0e9 | ||
|
|
38dd717aa6 | ||
|
|
6640c902de | ||
|
|
583c343d08 | ||
|
|
dfc48f7a05 | ||
|
|
cd9d854595 | ||
|
|
766cbd17e5 | ||
|
|
6df95bf981 | ||
|
|
6398ca0893 | ||
|
|
62372a48cc | ||
|
|
733632509a | ||
|
|
e1c13ec314 | ||
|
|
928ff0eabe | ||
|
|
406559b13d | ||
|
|
20c16aa6fd | ||
|
|
44ba5fd6d4 | ||
|
|
745c6adbf2 | ||
|
|
a9aa09636f | ||
|
|
6945986dd1 | ||
|
|
80c1f48f0e | ||
|
|
31e18205f0 | ||
|
|
e84a9d3f9b | ||
|
|
133feb8c76 | ||
|
|
c5fd9c01e5 | ||
|
|
a7bf5ceac3 | ||
|
|
74139aeb7e | ||
|
|
0cae4fd05c | ||
|
|
4b3b4fda61 | ||
|
|
adb684a6bf | ||
|
|
8493472983 | ||
|
|
862eeb7add | ||
|
|
4d9b211d69 | ||
|
|
ebb272324c | ||
|
|
c0290512b3 | ||
|
|
4ec56fe41e | ||
|
|
f8a7c46cf9 | ||
|
|
2bab8a9bb6 | ||
|
|
89e6a6215a | ||
|
|
88be87e03e | ||
|
|
22867faa9c | ||
|
|
16c0fc704d | ||
|
|
32fdd076bf | ||
|
|
402ae0d466 | ||
|
|
acb3c6d68b | ||
|
|
40ac83e632 | ||
|
|
58da395941 | ||
|
|
0e3ef94c9e | ||
|
|
1f68dfc2b5 | ||
|
|
84977a464c | ||
|
|
8cc1dc042d | ||
|
|
651d64f34d | ||
|
|
b98d4b59f5 | ||
|
|
df6d99c07d | ||
|
|
9843510e1f | ||
|
|
83bda3dfb2 |
Generated
+53
-5
@@ -637,8 +637,9 @@ checksum = "c8d4a3bb8b1e0c1050499d1815f5ab16d04f0959b233085fb31653fbfc9d98f9"
|
||||
name = "client"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"chrono",
|
||||
"futures",
|
||||
"manifest",
|
||||
"protocol",
|
||||
"reqwest",
|
||||
"serde",
|
||||
@@ -649,7 +650,6 @@ dependencies = [
|
||||
"tokio",
|
||||
"tokio-tungstenite 0.29.0",
|
||||
"uuid",
|
||||
"workdir",
|
||||
"workspace-api",
|
||||
]
|
||||
|
||||
@@ -2629,6 +2629,7 @@ version = "0.1.0"
|
||||
dependencies = [
|
||||
"agen",
|
||||
"arc-swap",
|
||||
"decodal",
|
||||
"protocol",
|
||||
"secrets",
|
||||
"serde",
|
||||
@@ -3506,6 +3507,7 @@ dependencies = [
|
||||
"schemars",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"sha2 0.11.0",
|
||||
"tokio",
|
||||
"ts-rs",
|
||||
"uuid",
|
||||
@@ -4400,14 +4402,19 @@ dependencies = [
|
||||
"agen",
|
||||
"async-trait",
|
||||
"base64 0.22.1",
|
||||
"fs4",
|
||||
"futures",
|
||||
"protocol",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"sha2 0.11.0",
|
||||
"tempfile",
|
||||
"thiserror 2.0.18",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"unicode-normalization",
|
||||
"unicode-properties",
|
||||
"unicode-security",
|
||||
"uuid",
|
||||
]
|
||||
|
||||
@@ -4615,6 +4622,27 @@ version = "1.2.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596"
|
||||
|
||||
[[package]]
|
||||
name = "standalone"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"agen",
|
||||
"async-trait",
|
||||
"client",
|
||||
"fs4",
|
||||
"futures",
|
||||
"manifest",
|
||||
"protocol",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"session-store",
|
||||
"tempfile",
|
||||
"thiserror 2.0.18",
|
||||
"tokio",
|
||||
"uuid",
|
||||
"worker",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "static_assertions"
|
||||
version = "1.1.0"
|
||||
@@ -5301,10 +5329,10 @@ name = "tui"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"agen",
|
||||
"async-trait",
|
||||
"base64 0.22.1",
|
||||
"client",
|
||||
"crossterm 0.28.1",
|
||||
"fs4",
|
||||
"manifest",
|
||||
"protocol",
|
||||
"pulldown-cmark",
|
||||
@@ -5313,13 +5341,14 @@ dependencies = [
|
||||
"serde",
|
||||
"serde_json",
|
||||
"session-store",
|
||||
"standalone",
|
||||
"tempfile",
|
||||
"thiserror 2.0.18",
|
||||
"ticket",
|
||||
"tokio",
|
||||
"toml",
|
||||
"unicode-width",
|
||||
"uuid",
|
||||
"worker",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -5411,6 +5440,22 @@ version = "0.1.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7df058c713841ad818f1dc5d3fd88063241cc61f49f5fbea4b951e8cf5a8d71d"
|
||||
|
||||
[[package]]
|
||||
name = "unicode-script"
|
||||
version = "0.5.8"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "383ad40bb927465ec0ce7720e033cb4ca06912855fc35db31b5755d0de75b1ee"
|
||||
|
||||
[[package]]
|
||||
name = "unicode-security"
|
||||
version = "0.1.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2e4ddba1535dd35ed8b61c52166b7155d7f4e4b8847cec6f48e71dc66d8b5e50"
|
||||
dependencies = [
|
||||
"unicode-normalization",
|
||||
"unicode-script",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "unicode-segmentation"
|
||||
version = "1.13.2"
|
||||
@@ -6570,6 +6615,7 @@ dependencies = [
|
||||
"tempfile",
|
||||
"thiserror 2.0.18",
|
||||
"tokio",
|
||||
"workspace-api",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -6617,6 +6663,7 @@ dependencies = [
|
||||
"wasmtime",
|
||||
"wat",
|
||||
"workdir",
|
||||
"workspace-api",
|
||||
"yoi-plugin-pdk",
|
||||
]
|
||||
|
||||
@@ -6659,10 +6706,11 @@ dependencies = [
|
||||
name = "workspace-api"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"protocol",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"ts-rs",
|
||||
"workdir",
|
||||
"webauthn-rs-proto",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
||||
@@ -5,6 +5,7 @@ members = [
|
||||
"crates/agen",
|
||||
"crates/agen-macros",
|
||||
"crates/session-store",
|
||||
"crates/standalone",
|
||||
"crates/secrets",
|
||||
"crates/manifest",
|
||||
"crates/mcp",
|
||||
@@ -36,6 +37,7 @@ default-members = [
|
||||
"crates/agen",
|
||||
"crates/agen-macros",
|
||||
"crates/session-store",
|
||||
"crates/standalone",
|
||||
"crates/secrets",
|
||||
"crates/manifest",
|
||||
"crates/mcp",
|
||||
@@ -66,6 +68,12 @@ default-members = [
|
||||
edition = "2024"
|
||||
license = "MIT"
|
||||
|
||||
[profile.dev]
|
||||
debug = "line-tables-only"
|
||||
|
||||
[profile.dev.package."*"]
|
||||
debug = false
|
||||
|
||||
[workspace.dependencies]
|
||||
# Internal crates
|
||||
client = { path = "crates/client" }
|
||||
@@ -87,6 +95,7 @@ protocol = { path = "crates/protocol" }
|
||||
session-metrics = { path = "crates/session-metrics" }
|
||||
session-analytics = { path = "crates/session-analytics" }
|
||||
session-store = { path = "crates/session-store" }
|
||||
standalone = { path = "crates/standalone" }
|
||||
secrets = { path = "crates/secrets" }
|
||||
tools = { path = "crates/tools" }
|
||||
config-source = { path = "crates/config-source" }
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
|
||||
use agen::llm_client::scheme::{Scheme, anthropic::AnthropicScheme};
|
||||
use agen::llm_client::transport::{HttpTransport, ResolvedAuth};
|
||||
use agen::{Engine, EngineRunExit, StopReason};
|
||||
use agen::{Engine, EngineRunExit, RunInterruptionReason};
|
||||
use std::time::Duration;
|
||||
|
||||
#[tokio::main]
|
||||
@@ -51,7 +51,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
EngineRunExit::Finished => println!("✅ Task completed normally"),
|
||||
EngineRunExit::Paused => println!("⏸️ Task paused"),
|
||||
EngineRunExit::Yielded => println!("↩️ Task yielded"),
|
||||
EngineRunExit::Interrupted(StopReason::LimitReached) => {
|
||||
EngineRunExit::Interrupted(RunInterruptionReason::LimitReached) => {
|
||||
println!("🔒 Turn limit reached")
|
||||
}
|
||||
EngineRunExit::Interrupted(reason) => println!("❌ Task interrupted: {reason:?}"),
|
||||
|
||||
@@ -39,8 +39,8 @@ use tracing::info;
|
||||
use tracing_subscriber::EnvFilter;
|
||||
|
||||
use agen::{
|
||||
Engine, EngineRunExit, StopReason,
|
||||
interceptor::{Interceptor, PostToolAction, ToolResultInfo},
|
||||
Engine, EngineRunExit, RunInterruptionReason,
|
||||
interceptor::{Interceptor, InterceptorResult, PostToolAction, ToolResultInfo},
|
||||
llm_client::{
|
||||
LlmClient,
|
||||
capability::{CacheStrategy, ModelCapability, StructuredOutput, ToolCallingSupport},
|
||||
@@ -280,7 +280,10 @@ impl ToolResultPrinterPolicy {
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for ToolResultPrinterPolicy {
|
||||
async fn post_tool_call(&self, info: &mut ToolResultInfo) -> PostToolAction {
|
||||
async fn post_tool_call(
|
||||
&self,
|
||||
info: &ToolResultInfo<'_, ()>,
|
||||
) -> InterceptorResult<PostToolAction> {
|
||||
let name = self
|
||||
.call_names
|
||||
.lock()
|
||||
@@ -294,7 +297,7 @@ impl Interceptor for ToolResultPrinterPolicy {
|
||||
println!(" Result ({}): ✅ {}", name, info.result.summary);
|
||||
}
|
||||
|
||||
PostToolAction::Continue
|
||||
Ok(PostToolAction::Continue)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -478,7 +481,8 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
// One-shot mode
|
||||
if let Some(prompt) = args.prompt {
|
||||
let output = engine.run(&mut history, &prompt).await;
|
||||
if let EngineRunExit::Interrupted(StopReason::Unexpected(error)) = output.result {
|
||||
if let EngineRunExit::Interrupted(RunInterruptionReason::Unexpected(error)) = output.result
|
||||
{
|
||||
eprintln!("\n❌ Error: {error}");
|
||||
}
|
||||
|
||||
@@ -518,7 +522,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
break;
|
||||
}
|
||||
|
||||
if let EngineRunExit::Interrupted(StopReason::Unexpected(error)) =
|
||||
if let EngineRunExit::Interrupted(RunInterruptionReason::Unexpected(error)) =
|
||||
locked.run(&mut history, input).await
|
||||
{
|
||||
eprintln!("\n❌ Error: {error}");
|
||||
|
||||
+345
-113
@@ -15,8 +15,12 @@ use crate::{
|
||||
},
|
||||
handler::{ErrorKind, StatusKind, ToolUseBlockStart, UsageKind},
|
||||
interceptor::{
|
||||
DefaultInterceptor, Interceptor, PostToolAction, PreRequestAction, PreToolAction,
|
||||
PromptAction, ToolCallInfo, ToolResultInfo, TurnEndAction,
|
||||
AssistantTurnEndContext, DefaultInterceptor, Interceptor, InterceptorCallId,
|
||||
InterceptorCounter, InterceptorCounters, InterceptorError, InterceptorErrorCategory,
|
||||
InterceptorFailure, InterceptorInvocation, InterceptorPhase, InterceptorRunId,
|
||||
InterceptorTurnId, PendingHistoryAppendsContext, PostToolAction, PreLlmRequestContext,
|
||||
PreRequestAction, PreToolAction, PromptAction, PromptSubmitContext, RunExitContext,
|
||||
ToolCallInfo, ToolResultInfo, TurnEndAction,
|
||||
},
|
||||
llm_client::{
|
||||
ClientError, ConfigWarning, LlmClient, Request, RequestConfig, ResponseStream,
|
||||
@@ -58,6 +62,9 @@ pub enum EngineError {
|
||||
/// A durable-history observer rejected an item before it entered history.
|
||||
#[error("History append failed: {0}")]
|
||||
HistoryAppend(String),
|
||||
/// A trusted host interceptor callback failed.
|
||||
#[error(transparent)]
|
||||
Interceptor(#[from] InterceptorFailure),
|
||||
/// Tool terminalization lost its execution-attempt compare-and-set fence.
|
||||
#[error("Tool execution attempt fence failed: {0}")]
|
||||
ToolAttemptFence(String),
|
||||
@@ -147,12 +154,12 @@ pub enum EngineRunExit {
|
||||
Finished,
|
||||
Paused,
|
||||
Yielded,
|
||||
Interrupted(StopReason),
|
||||
Interrupted(RunInterruptionReason),
|
||||
}
|
||||
|
||||
/// A typed reason why an engine run could not finish normally.
|
||||
#[derive(Debug)]
|
||||
pub enum StopReason {
|
||||
pub enum RunInterruptionReason {
|
||||
LimitReached,
|
||||
ContextWindowExceeded,
|
||||
Cancelled,
|
||||
@@ -165,13 +172,15 @@ impl From<Result<EngineResult, EngineError>> for EngineRunExit {
|
||||
Ok(EngineResult::Finished) => Self::Finished,
|
||||
Ok(EngineResult::Paused) => Self::Paused,
|
||||
Ok(EngineResult::Yielded) => Self::Yielded,
|
||||
Ok(EngineResult::LimitReached) => Self::Interrupted(StopReason::LimitReached),
|
||||
Err(EngineError::Client(ClientError::ContextWindowExceeded)) => {
|
||||
Self::Interrupted(StopReason::ContextWindowExceeded)
|
||||
Ok(EngineResult::LimitReached) => {
|
||||
Self::Interrupted(RunInterruptionReason::LimitReached)
|
||||
}
|
||||
Err(EngineError::Cancelled) => Self::Interrupted(StopReason::Cancelled),
|
||||
Err(EngineError::Client(ClientError::ContextWindowExceeded)) => {
|
||||
Self::Interrupted(RunInterruptionReason::ContextWindowExceeded)
|
||||
}
|
||||
Err(EngineError::Cancelled) => Self::Interrupted(RunInterruptionReason::Cancelled),
|
||||
Err(EngineError::PauseRequested) => Self::Paused,
|
||||
Err(error) => Self::Interrupted(StopReason::Unexpected(error)),
|
||||
Err(error) => Self::Interrupted(RunInterruptionReason::Unexpected(error)),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -179,7 +188,7 @@ impl From<Result<EngineResult, EngineError>> for EngineRunExit {
|
||||
/// Result of [`Engine::run`] or [`Engine::resume`].
|
||||
///
|
||||
/// Contains the `Locked` Engine (ready for subsequent runs) and the outcome.
|
||||
pub struct EngineRunOutput<C: LlmClient, A = ()> {
|
||||
pub struct EngineRunOutput<C: LlmClient, A: Send + Sync = ()> {
|
||||
/// The Engine, now in Locked state.
|
||||
pub engine: Engine<C, Locked, A>,
|
||||
/// Outcome of the turn.
|
||||
@@ -303,7 +312,7 @@ enum StreamCompletion {
|
||||
Interrupted { reason: String },
|
||||
}
|
||||
|
||||
pub struct Engine<C: LlmClient, S: EngineState = Mutable, A = ()> {
|
||||
pub struct Engine<C: LlmClient, S: EngineState = Mutable, A: Send + Sync = ()> {
|
||||
/// LLM client
|
||||
client: C,
|
||||
/// Retry policy for opening an LLM response stream.
|
||||
@@ -320,7 +329,7 @@ pub struct Engine<C: LlmClient, S: EngineState = Mutable, A = ()> {
|
||||
/// Tool server handle
|
||||
tool_server: ToolServerHandle,
|
||||
/// Interceptor for control-flow decisions
|
||||
interceptor: Box<dyn Interceptor>,
|
||||
interceptor: Box<dyn Interceptor<A>>,
|
||||
/// System prompt
|
||||
system_prompt: Option<String>,
|
||||
/// History length at lock time (only meaningful in Locked state)
|
||||
@@ -339,6 +348,11 @@ pub struct Engine<C: LlmClient, S: EngineState = Mutable, A = ()> {
|
||||
/// `max_turns` is enforced against this run-scoped count rather than the
|
||||
/// cumulative `turn_count` above.
|
||||
active_run_turn_count: Option<usize>,
|
||||
/// Identity retained across pause/yield and resume.
|
||||
active_run_id: Option<InterceptorRunId>,
|
||||
next_run_id: u64,
|
||||
interceptor_invocation_count: usize,
|
||||
last_run_exit_observer_failure: Option<InterceptorFailure>,
|
||||
/// LlmCall count (per-Engine running counter, monotonic). Unlike
|
||||
/// `turn_count` this never collapses retries.
|
||||
llm_call_count: usize,
|
||||
@@ -419,21 +433,57 @@ pub struct Engine<C: LlmClient, S: EngineState = Mutable, A = ()> {
|
||||
_state: PhantomData<(S, A)>,
|
||||
}
|
||||
|
||||
impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
impl<C: LlmClient, S: EngineState, A: Send + Sync> Engine<C, S, A> {
|
||||
fn start_logical_run(&mut self) {
|
||||
self.active_run_turn_count = Some(0);
|
||||
self.active_run_id = Some(InterceptorRunId(self.next_run_id));
|
||||
self.next_run_id = self.next_run_id.wrapping_add(1).max(1);
|
||||
self.interceptor_invocation_count = 0;
|
||||
self.last_run_exit_observer_failure = None;
|
||||
}
|
||||
|
||||
fn ensure_logical_run(&mut self) {
|
||||
self.active_run_turn_count.get_or_insert(0);
|
||||
if self.active_run_id.is_none() {
|
||||
self.active_run_id = Some(InterceptorRunId(self.next_run_id));
|
||||
self.next_run_id = self.next_run_id.wrapping_add(1).max(1);
|
||||
self.interceptor_invocation_count = 0;
|
||||
}
|
||||
}
|
||||
|
||||
fn finish_logical_run(&mut self, result: &Result<EngineResult, EngineError>) {
|
||||
if !matches!(
|
||||
result,
|
||||
Ok(EngineResult::Paused | EngineResult::Yielded) | Err(EngineError::PauseRequested)
|
||||
) {
|
||||
fn interceptor_invocation(
|
||||
&mut self,
|
||||
phase: InterceptorPhase,
|
||||
turn_id: Option<usize>,
|
||||
call_id: Option<InterceptorCallId>,
|
||||
tool_call: usize,
|
||||
) -> InterceptorInvocation {
|
||||
let invocation = self.interceptor_invocation_count;
|
||||
self.interceptor_invocation_count = self.interceptor_invocation_count.saturating_add(1);
|
||||
InterceptorInvocation {
|
||||
run_id: self
|
||||
.active_run_id
|
||||
.expect("logical run identity must exist before interception"),
|
||||
turn_id: turn_id.map(|value| InterceptorTurnId(value as u64)),
|
||||
call_id,
|
||||
phase,
|
||||
counters: InterceptorCounters {
|
||||
invocation: InterceptorCounter::from_usize(invocation),
|
||||
engine_turn: InterceptorCounter::from_usize(self.turn_count),
|
||||
run_turn: InterceptorCounter::from_usize(
|
||||
self.active_run_turn_count.unwrap_or_default(),
|
||||
),
|
||||
llm_call: InterceptorCounter::from_usize(self.llm_call_count),
|
||||
tool_batch: InterceptorCounter::from_usize(self.tool_execution_batch_count),
|
||||
tool_call: InterceptorCounter::from_usize(tool_call),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn finish_logical_run(&mut self, exit: &EngineRunExit) {
|
||||
if !matches!(exit, EngineRunExit::Paused | EngineRunExit::Yielded) {
|
||||
self.active_run_turn_count = None;
|
||||
self.active_run_id = None;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -739,7 +789,7 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
/// The interceptor governs approval, skip, pause, and abort decisions
|
||||
/// at key points in the execution loop. If not set, the default
|
||||
/// interceptor is used (all Continue / Finish).
|
||||
pub fn set_interceptor(&mut self, interceptor: impl Interceptor + 'static) {
|
||||
pub fn set_interceptor(&mut self, interceptor: impl Interceptor<A> + 'static) {
|
||||
self.interceptor = Box::new(interceptor);
|
||||
}
|
||||
|
||||
@@ -840,6 +890,10 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
///
|
||||
/// `Some` is retained only while Pause or Yield permits a later
|
||||
/// [`resume`](Self::resume). Terminal outcomes return this to `None`.
|
||||
pub fn last_run_exit_observer_failure(&self) -> Option<&InterceptorFailure> {
|
||||
self.last_run_exit_observer_failure.as_ref()
|
||||
}
|
||||
|
||||
pub fn active_run_turn_count(&self) -> Option<usize> {
|
||||
self.active_run_turn_count
|
||||
}
|
||||
@@ -851,6 +905,13 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
/// [`resume`](Self::resume) starts a fresh budget.
|
||||
pub fn set_active_run_turn_count(&mut self, turn_count: Option<usize>) {
|
||||
self.active_run_turn_count = turn_count;
|
||||
if turn_count.is_none() {
|
||||
self.active_run_id = None;
|
||||
} else if self.active_run_id.is_none() {
|
||||
self.active_run_id = Some(InterceptorRunId(self.next_run_id));
|
||||
self.next_run_id = self.next_run_id.wrapping_add(1).max(1);
|
||||
self.interceptor_invocation_count = 0;
|
||||
}
|
||||
}
|
||||
|
||||
/// Get the current LlmCall count (per-Engine running counter, never
|
||||
@@ -1076,24 +1137,28 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
request
|
||||
}
|
||||
|
||||
/// Hooks: on_prompt_submit
|
||||
///
|
||||
async fn finalize_interruption<T>(
|
||||
async fn finalize_run_exit(
|
||||
&mut self,
|
||||
result: Result<T, EngineError>,
|
||||
) -> Result<T, EngineError> {
|
||||
match result {
|
||||
Ok(value) => Ok(value),
|
||||
Err(err) => {
|
||||
let reason = match &err {
|
||||
EngineError::Aborted(reason) => reason.clone(),
|
||||
EngineError::Cancelled => "Cancelled".to_string(),
|
||||
_ => err.to_string(),
|
||||
};
|
||||
self.interceptor.on_abort(&reason).await;
|
||||
Err(err)
|
||||
}
|
||||
history: &History<A>,
|
||||
result: Result<EngineResult, EngineError>,
|
||||
) -> EngineRunExit {
|
||||
let exit = EngineRunExit::from(result);
|
||||
let invocation = self.interceptor_invocation(InterceptorPhase::RunExit, None, None, 0);
|
||||
self.last_run_exit_observer_failure = None;
|
||||
if let Err(error) = self
|
||||
.interceptor
|
||||
.on_run_exit(RunExitContext {
|
||||
invocation,
|
||||
exit: &exit,
|
||||
history: history.entries(),
|
||||
})
|
||||
.await
|
||||
{
|
||||
self.last_run_exit_observer_failure =
|
||||
Some(InterceptorFailure::new(InterceptorPhase::RunExit, error));
|
||||
}
|
||||
self.finish_logical_run(&exit);
|
||||
exit
|
||||
}
|
||||
|
||||
/// Check for pending tool calls (for resuming from Pause)
|
||||
@@ -1164,21 +1229,60 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
// Phase 1: Apply pre_tool_call interceptor (determine skip/abort/synthetic result)
|
||||
let mut approved_calls = Vec::new();
|
||||
for (call_index, mut tool_call) in tool_calls.into_iter().enumerate() {
|
||||
let expected_tool_use_id = tool_call.id.clone();
|
||||
let context = ToolExecutionContext::new(&tool_call.id, &batch_id, call_index);
|
||||
if let Some((meta, tool)) = self.tool_server.get_tool(&tool_call.name) {
|
||||
let invocation = self.interceptor_invocation(
|
||||
InterceptorPhase::PreToolCall,
|
||||
Some(self.turn_count.saturating_sub(1)),
|
||||
Some(InterceptorCallId::Tool(expected_tool_use_id.clone())),
|
||||
call_index,
|
||||
);
|
||||
let mut info = ToolCallInfo {
|
||||
invocation,
|
||||
history: history.entries(),
|
||||
call: tool_call.clone(),
|
||||
meta,
|
||||
tool,
|
||||
context,
|
||||
};
|
||||
|
||||
match self.interceptor.pre_tool_call(&mut info).await {
|
||||
let pre_tool_action =
|
||||
self.interceptor
|
||||
.pre_tool_call(&mut info)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
EngineError::from(InterceptorFailure::new(
|
||||
InterceptorPhase::PreToolCall,
|
||||
error,
|
||||
))
|
||||
})?;
|
||||
if info.call.id != expected_tool_use_id {
|
||||
return Err(InterceptorFailure::new(
|
||||
InterceptorPhase::PreToolCall,
|
||||
InterceptorError::new(
|
||||
InterceptorErrorCategory::ContractViolation,
|
||||
"pre-tool interceptor changed immutable tool call identity",
|
||||
),
|
||||
)
|
||||
.into());
|
||||
}
|
||||
match pre_tool_action {
|
||||
PreToolAction::Continue => {}
|
||||
PreToolAction::Skip => {
|
||||
continue;
|
||||
}
|
||||
PreToolAction::SyntheticResult(result) => {
|
||||
if result.tool_use_id != expected_tool_use_id {
|
||||
return Err(InterceptorFailure::new(
|
||||
InterceptorPhase::PreToolCall,
|
||||
InterceptorError::new(
|
||||
InterceptorErrorCategory::ContractViolation,
|
||||
"synthetic tool result changed immutable tool call identity",
|
||||
),
|
||||
)
|
||||
.into());
|
||||
}
|
||||
let tool_call = info.call;
|
||||
let mut context = info.context;
|
||||
context.call_id = tool_call.id.clone();
|
||||
@@ -1285,20 +1389,31 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
let mut terminal_call_ids = HashSet::new();
|
||||
let mut pause_requested = false;
|
||||
let mut pause_deadline = None;
|
||||
let mut batch_error = None;
|
||||
let mut locally_enqueued_cancel = false;
|
||||
for result in synthetic_results {
|
||||
self.finalize_and_commit_tool_result(
|
||||
history,
|
||||
annotate,
|
||||
result,
|
||||
None,
|
||||
&call_info_map,
|
||||
&mut attempt_fence,
|
||||
&mut terminal_call_ids,
|
||||
)
|
||||
.await?;
|
||||
if let Err(error) = self
|
||||
.finalize_and_commit_tool_result(
|
||||
history,
|
||||
annotate,
|
||||
result,
|
||||
None,
|
||||
&call_info_map,
|
||||
&mut attempt_fence,
|
||||
&mut terminal_call_ids,
|
||||
)
|
||||
.await
|
||||
&& batch_error.is_none()
|
||||
{
|
||||
batch_error = Some(error);
|
||||
}
|
||||
}
|
||||
|
||||
let mut futures = futures;
|
||||
if batch_error.is_some() && !futures.is_empty() {
|
||||
let _ = self.cancel_tx.try_send(());
|
||||
locally_enqueued_cancel = true;
|
||||
}
|
||||
while !futures.is_empty() {
|
||||
tokio::select! {
|
||||
// If cancellation and a completed result are both ready, drain
|
||||
@@ -1308,7 +1423,7 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
result = futures.next() => {
|
||||
let (attempt_id, result) =
|
||||
result.expect("non-empty FuturesUnordered returns a result");
|
||||
self.finalize_and_commit_tool_result(
|
||||
if let Err(error) = self.finalize_and_commit_tool_result(
|
||||
history,
|
||||
annotate,
|
||||
result,
|
||||
@@ -1316,7 +1431,15 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
&call_info_map,
|
||||
&mut attempt_fence,
|
||||
&mut terminal_call_ids,
|
||||
).await?;
|
||||
).await {
|
||||
if batch_error.is_none() {
|
||||
batch_error = Some(error);
|
||||
}
|
||||
if !futures.is_empty() {
|
||||
let _ = self.cancel_tx.try_send(());
|
||||
locally_enqueued_cancel = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
pause = self.pause_rx.recv(), if !pause_requested => {
|
||||
if pause.is_some() {
|
||||
@@ -1333,6 +1456,7 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
_ = tokio::time::sleep_until(pause_deadline.unwrap_or_else(TokioInstant::now)), if pause_deadline.is_some() => {
|
||||
pause_deadline = None;
|
||||
let _ = self.cancel_tx.try_send(());
|
||||
locally_enqueued_cancel = true;
|
||||
}
|
||||
cancel = self.cancel_rx.recv() => {
|
||||
if cancel.is_some() {
|
||||
@@ -1378,7 +1502,7 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
result = futures.next() => {
|
||||
let (attempt_id, result) =
|
||||
result.expect("non-empty FuturesUnordered returns a result");
|
||||
self.finalize_and_commit_tool_result(
|
||||
if let Err(error) = self.finalize_and_commit_tool_result(
|
||||
history,
|
||||
annotate,
|
||||
result,
|
||||
@@ -1386,7 +1510,11 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
&call_info_map,
|
||||
&mut attempt_fence,
|
||||
&mut terminal_call_ids,
|
||||
).await?;
|
||||
).await
|
||||
&& batch_error.is_none()
|
||||
{
|
||||
batch_error = Some(error);
|
||||
}
|
||||
}
|
||||
_ = tokio::time::sleep_until(deadline) => break,
|
||||
}
|
||||
@@ -1400,7 +1528,7 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
if let Some(handle) = execution_handles.get(call_id) {
|
||||
handle.force_close();
|
||||
}
|
||||
self.finalize_and_commit_tool_result(
|
||||
if let Err(error) = self.finalize_and_commit_tool_result(
|
||||
history,
|
||||
annotate,
|
||||
ToolResult::outcome_unknown(call_id),
|
||||
@@ -1408,11 +1536,18 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
&call_info_map,
|
||||
&mut attempt_fence,
|
||||
&mut terminal_call_ids,
|
||||
).await?;
|
||||
).await
|
||||
&& batch_error.is_none()
|
||||
{
|
||||
batch_error = Some(error);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
self.timeline.abort_current_block();
|
||||
if let Some(error) = batch_error.take() {
|
||||
return Err(error);
|
||||
}
|
||||
if pause_requested {
|
||||
return Ok(ToolExecutionResult::Paused);
|
||||
}
|
||||
@@ -1421,6 +1556,16 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
}
|
||||
}
|
||||
|
||||
// A result-biased ready sibling can empty the batch before the local
|
||||
// cancel signal is selected. Never let that current-batch signal leak
|
||||
// into the next run or resume call.
|
||||
if locally_enqueued_cancel {
|
||||
let _ = self.cancel_rx.try_recv();
|
||||
}
|
||||
if let Some(error) = batch_error {
|
||||
self.timeline.abort_current_block();
|
||||
return Err(error);
|
||||
}
|
||||
Ok(if pause_requested {
|
||||
ToolExecutionResult::Paused
|
||||
} else {
|
||||
@@ -1464,31 +1609,13 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
}
|
||||
|
||||
let call_info = call_info_map.get(&tool_result.tool_use_id);
|
||||
let mut abort_reason = None;
|
||||
if let Some((tool_call, meta, tool, context)) = call_info {
|
||||
let mut info = ToolResultInfo {
|
||||
call: tool_call.clone(),
|
||||
result: tool_result,
|
||||
meta: meta.clone(),
|
||||
tool: tool.clone(),
|
||||
context: context.clone(),
|
||||
};
|
||||
|
||||
match self.interceptor.post_tool_call(&mut info).await {
|
||||
PostToolAction::Continue => {}
|
||||
PostToolAction::Abort(reason) => {
|
||||
abort_reason = Some(reason);
|
||||
}
|
||||
}
|
||||
tool_result = info.result;
|
||||
}
|
||||
if tool_result.is_error && tool_result.disposition.is_success() {
|
||||
tool_result.disposition = ToolResultDisposition::Error;
|
||||
}
|
||||
tool_result.is_error = !tool_result.disposition.is_success();
|
||||
|
||||
// Cap content only after post_tool_call so interceptors still observe
|
||||
// the full payload and any content they inject is bounded too.
|
||||
// Bound the terminal payload before committing it so the post-tool
|
||||
// interceptor observes exactly the model-visible durable result.
|
||||
if let (Some(limits), Some((tool_call, _, _, _)), Some(content)) = (
|
||||
self.tool_output_limits.as_ref(),
|
||||
call_info,
|
||||
@@ -1541,9 +1668,38 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
"Tool execution terminalized"
|
||||
);
|
||||
self.emit_tool_result(&tool_result);
|
||||
if let Some(reason) = abort_reason {
|
||||
return Err(EngineError::Aborted(reason));
|
||||
|
||||
if let Some((tool_call, meta, tool, context)) = call_info {
|
||||
let invocation = self.interceptor_invocation(
|
||||
InterceptorPhase::PostToolCall,
|
||||
Some(self.turn_count.saturating_sub(1)),
|
||||
Some(InterceptorCallId::Tool(tool_call.id.clone())),
|
||||
context.call_index,
|
||||
);
|
||||
let info = ToolResultInfo {
|
||||
invocation,
|
||||
history: history.entries(),
|
||||
call: tool_call.clone(),
|
||||
result: tool_result,
|
||||
meta: meta.clone(),
|
||||
tool: tool.clone(),
|
||||
context: context.clone(),
|
||||
};
|
||||
let post_tool_action =
|
||||
self.interceptor
|
||||
.post_tool_call(&info)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
EngineError::from(InterceptorFailure::new(
|
||||
InterceptorPhase::PostToolCall,
|
||||
error,
|
||||
))
|
||||
})?;
|
||||
if let PostToolAction::Abort(reason) = post_tool_action {
|
||||
return Err(EngineError::Aborted(reason));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
@@ -1606,11 +1762,25 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
// These are committed *before* the per-request clone so they
|
||||
// participate in the LLM request below and get persisted by
|
||||
// the caller that owns durable history.
|
||||
let pending_invocation = self.interceptor_invocation(
|
||||
InterceptorPhase::PendingHistoryAppends,
|
||||
Some(current_turn),
|
||||
None,
|
||||
0,
|
||||
);
|
||||
let pending = self
|
||||
.interceptor
|
||||
.pending_history_appends()
|
||||
.pending_history_appends(PendingHistoryAppendsContext {
|
||||
invocation: pending_invocation,
|
||||
history: history.entries(),
|
||||
})
|
||||
.await
|
||||
.map_err(EngineError::HistoryAppend)?;
|
||||
.map_err(|error| {
|
||||
EngineError::from(InterceptorFailure::new(
|
||||
InterceptorPhase::PendingHistoryAppends,
|
||||
error,
|
||||
))
|
||||
})?;
|
||||
if !pending.is_empty() {
|
||||
self.append_history_items(history, pending, annotate)?;
|
||||
}
|
||||
@@ -1677,7 +1847,27 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
}
|
||||
|
||||
// Interceptor: pre_llm_request
|
||||
match self.interceptor.pre_llm_request(&mut request_context).await {
|
||||
let request_invocation = self.interceptor_invocation(
|
||||
InterceptorPhase::PreLlmRequest,
|
||||
Some(current_turn),
|
||||
Some(InterceptorCallId::Llm(self.llm_call_count as u64)),
|
||||
0,
|
||||
);
|
||||
let pre_request_action = self
|
||||
.interceptor
|
||||
.pre_llm_request(PreLlmRequestContext {
|
||||
invocation: request_invocation,
|
||||
items: &mut request_context,
|
||||
history: history.entries(),
|
||||
})
|
||||
.await
|
||||
.map_err(|error| {
|
||||
EngineError::from(InterceptorFailure::new(
|
||||
InterceptorPhase::PreLlmRequest,
|
||||
error,
|
||||
))
|
||||
})?;
|
||||
match pre_request_action {
|
||||
PreRequestAction::Cancel(reason) => {
|
||||
info!(reason = %reason, "Aborted by interceptor");
|
||||
for cb in &self.turn_end_cbs {
|
||||
@@ -1789,21 +1979,45 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
let tool_calls = self.tool_call_collector.take_collected();
|
||||
let assistant_items =
|
||||
self.build_assistant_items(&reasoning_items, &text_blocks, &tool_calls);
|
||||
let assistant_start = history.len();
|
||||
self.append_history_items(history, assistant_items, annotate)?;
|
||||
|
||||
if tool_calls.is_empty() {
|
||||
let turn_end_context = history.items_cloned();
|
||||
match self.interceptor.on_turn_end(&turn_end_context).await {
|
||||
TurnEndAction::Finish => {
|
||||
return Ok(EngineResult::Finished);
|
||||
}
|
||||
TurnEndAction::ContinueWithMessages(additional) => {
|
||||
self.append_history_items(history, additional, annotate)?;
|
||||
let assistant_invocation = self.interceptor_invocation(
|
||||
InterceptorPhase::AssistantTurnEnd,
|
||||
Some(current_turn),
|
||||
Some(InterceptorCallId::Llm(
|
||||
self.llm_call_count.saturating_sub(1) as u64,
|
||||
)),
|
||||
0,
|
||||
);
|
||||
let assistant_turn_action = self
|
||||
.interceptor
|
||||
.on_assistant_turn_end(AssistantTurnEndContext {
|
||||
invocation: assistant_invocation,
|
||||
assistant_entries: &history.entries()[assistant_start..],
|
||||
history: history.entries(),
|
||||
tool_calls: &tool_calls,
|
||||
})
|
||||
.await
|
||||
.map_err(|error| {
|
||||
EngineError::from(InterceptorFailure::new(
|
||||
InterceptorPhase::AssistantTurnEnd,
|
||||
error,
|
||||
))
|
||||
})?;
|
||||
match assistant_turn_action {
|
||||
TurnEndAction::Finish if tool_calls.is_empty() => {
|
||||
return Ok(EngineResult::Finished);
|
||||
}
|
||||
TurnEndAction::Finish => {}
|
||||
TurnEndAction::ContinueWithMessages(additional) => {
|
||||
self.append_history_items(history, additional, annotate)?;
|
||||
if tool_calls.is_empty() {
|
||||
continue;
|
||||
}
|
||||
TurnEndAction::Pause => {
|
||||
return Ok(EngineResult::Paused);
|
||||
}
|
||||
}
|
||||
TurnEndAction::Pause => {
|
||||
return Ok(EngineResult::Paused);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2096,7 +2310,7 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
}
|
||||
}
|
||||
|
||||
impl<C: LlmClient, A> Engine<C, Mutable, A> {
|
||||
impl<C: LlmClient, A: Send + Sync> Engine<C, Mutable, A> {
|
||||
/// Create a new annotated Engine (in Mutable state).
|
||||
pub fn new_annotated(client: C) -> Self {
|
||||
let text_block_collector = TextBlockCollector::new();
|
||||
@@ -2124,6 +2338,10 @@ impl<C: LlmClient, A> Engine<C, Mutable, A> {
|
||||
locked_prefix_len: 0,
|
||||
turn_count: 0,
|
||||
active_run_turn_count: None,
|
||||
active_run_id: None,
|
||||
next_run_id: 1,
|
||||
interceptor_invocation_count: 0,
|
||||
last_run_exit_observer_failure: None,
|
||||
llm_call_count: 0,
|
||||
tool_execution_batch_count: 0,
|
||||
max_turns: None,
|
||||
@@ -2399,6 +2617,10 @@ impl<C: LlmClient, A> Engine<C, Mutable, A> {
|
||||
locked_prefix_len,
|
||||
turn_count: self.turn_count,
|
||||
active_run_turn_count: self.active_run_turn_count,
|
||||
active_run_id: self.active_run_id,
|
||||
next_run_id: self.next_run_id,
|
||||
interceptor_invocation_count: self.interceptor_invocation_count,
|
||||
last_run_exit_observer_failure: self.last_run_exit_observer_failure,
|
||||
llm_call_count: self.llm_call_count,
|
||||
tool_execution_batch_count: self.tool_execution_batch_count,
|
||||
max_turns: self.max_turns,
|
||||
@@ -2475,7 +2697,7 @@ impl<C: LlmClient> Engine<C, Mutable, ()> {
|
||||
}
|
||||
}
|
||||
|
||||
impl<C: LlmClient, A> Engine<C, Locked, A> {
|
||||
impl<C: LlmClient, A: Send + Sync> Engine<C, Locked, A> {
|
||||
/// Execute a turn
|
||||
///
|
||||
/// Adds a new user message to history and sends a request to the LLM.
|
||||
@@ -2486,9 +2708,10 @@ impl<C: LlmClient, A> Engine<C, Locked, A> {
|
||||
user_input: impl Into<String>,
|
||||
annotate: &mut impl FnMut(&Item) -> Result<A, String>,
|
||||
) -> EngineRunExit {
|
||||
self.run_result_with_annotation(history, user_input.into(), annotate)
|
||||
.await
|
||||
.into()
|
||||
let result = self
|
||||
.run_result_with_annotation(history, user_input.into(), annotate)
|
||||
.await;
|
||||
self.finalize_run_exit(history, result).await
|
||||
}
|
||||
|
||||
async fn run_result_with_annotation(
|
||||
@@ -2499,13 +2722,26 @@ impl<C: LlmClient, A> Engine<C, Locked, A> {
|
||||
) -> Result<EngineResult, EngineError> {
|
||||
// Supplying new user input abandons any paused/yielded logical run.
|
||||
self.active_run_turn_count = None;
|
||||
self.active_run_id = None;
|
||||
self.start_logical_run();
|
||||
let mut user_item = Item::user_message(user_input);
|
||||
let extras = match self.interceptor.on_prompt_submit(&mut user_item).await {
|
||||
PromptAction::Cancel(reason) => {
|
||||
return self
|
||||
.finalize_interruption(Err(EngineError::Aborted(reason)))
|
||||
.await;
|
||||
}
|
||||
let invocation = self.interceptor_invocation(InterceptorPhase::PromptSubmit, None, None, 0);
|
||||
let prompt_action = self
|
||||
.interceptor
|
||||
.on_prompt_submit(PromptSubmitContext {
|
||||
invocation,
|
||||
item: &mut user_item,
|
||||
history: history.entries(),
|
||||
})
|
||||
.await
|
||||
.map_err(|error| {
|
||||
EngineError::from(InterceptorFailure::new(
|
||||
InterceptorPhase::PromptSubmit,
|
||||
error,
|
||||
))
|
||||
})?;
|
||||
let extras = match prompt_action {
|
||||
PromptAction::Cancel(reason) => return Err(EngineError::Aborted(reason)),
|
||||
PromptAction::Continue => Vec::new(),
|
||||
PromptAction::ContinueWith(items) => items,
|
||||
};
|
||||
@@ -2513,14 +2749,10 @@ impl<C: LlmClient, A> Engine<C, Locked, A> {
|
||||
if !extras.is_empty() {
|
||||
self.append_history_items(history, extras, annotate)?;
|
||||
}
|
||||
self.start_logical_run();
|
||||
let result = match self.run_turn_loop(history, annotate).await {
|
||||
match self.run_turn_loop(history, annotate).await {
|
||||
Err(EngineError::PauseRequested) => Ok(EngineResult::Paused),
|
||||
other => other,
|
||||
};
|
||||
let result = self.finalize_interruption(result).await;
|
||||
self.finish_logical_run(&result);
|
||||
result
|
||||
}
|
||||
}
|
||||
|
||||
/// Resume execution (from Paused state).
|
||||
@@ -2529,9 +2761,8 @@ impl<C: LlmClient, A> Engine<C, Locked, A> {
|
||||
history: &mut History<A>,
|
||||
annotate: &mut impl FnMut(&Item) -> Result<A, String>,
|
||||
) -> EngineRunExit {
|
||||
self.resume_result_with_annotation(history, annotate)
|
||||
.await
|
||||
.into()
|
||||
let result = self.resume_result_with_annotation(history, annotate).await;
|
||||
self.finalize_run_exit(history, result).await
|
||||
}
|
||||
|
||||
async fn resume_result_with_annotation(
|
||||
@@ -2540,13 +2771,10 @@ impl<C: LlmClient, A> Engine<C, Locked, A> {
|
||||
annotate: &mut impl FnMut(&Item) -> Result<A, String>,
|
||||
) -> Result<EngineResult, EngineError> {
|
||||
self.ensure_logical_run();
|
||||
let result = match self.run_turn_loop(history, annotate).await {
|
||||
match self.run_turn_loop(history, annotate).await {
|
||||
Err(EngineError::PauseRequested) => Ok(EngineResult::Paused),
|
||||
other => other,
|
||||
};
|
||||
let result = self.finalize_interruption(result).await;
|
||||
self.finish_logical_run(&result);
|
||||
result
|
||||
}
|
||||
}
|
||||
|
||||
/// Get the prefix length at lock time
|
||||
@@ -2572,6 +2800,10 @@ impl<C: LlmClient, A> Engine<C, Locked, A> {
|
||||
locked_prefix_len: 0,
|
||||
turn_count: self.turn_count,
|
||||
active_run_turn_count: self.active_run_turn_count,
|
||||
active_run_id: self.active_run_id,
|
||||
next_run_id: self.next_run_id,
|
||||
interceptor_invocation_count: self.interceptor_invocation_count,
|
||||
last_run_exit_observer_failure: self.last_run_exit_observer_failure,
|
||||
llm_call_count: self.llm_call_count,
|
||||
tool_execution_batch_count: self.tool_execution_batch_count,
|
||||
max_turns: self.max_turns,
|
||||
|
||||
+250
-28
@@ -9,8 +9,202 @@ use std::sync::Arc;
|
||||
use async_trait::async_trait;
|
||||
|
||||
use crate::Item;
|
||||
use crate::engine::EngineRunExit;
|
||||
use crate::history::HistoryEntry;
|
||||
use crate::tool::{Tool, ToolCall, ToolExecutionContext, ToolMeta, ToolResult};
|
||||
|
||||
// =============================================================================
|
||||
// Typed lifecycle metadata and failures
|
||||
// =============================================================================
|
||||
|
||||
/// Maximum UTF-8 byte length retained for interceptor diagnostics.
|
||||
pub const MAX_INTERCEPTOR_DIAGNOSTIC_BYTES: usize = 1024;
|
||||
|
||||
/// Stable category for the source of an interceptor failure.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum InterceptorErrorCategory {
|
||||
Policy,
|
||||
Dependency,
|
||||
ContractViolation,
|
||||
Internal,
|
||||
}
|
||||
|
||||
impl std::fmt::Display for InterceptorErrorCategory {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter.write_str(match self {
|
||||
Self::Policy => "policy",
|
||||
Self::Dependency => "dependency",
|
||||
Self::ContractViolation => "contract_violation",
|
||||
Self::Internal => "internal",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// A typed, bounded failure returned by an [`Interceptor`] implementation.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
|
||||
#[error("{category}: {diagnostic}")]
|
||||
pub struct InterceptorError {
|
||||
category: InterceptorErrorCategory,
|
||||
diagnostic: String,
|
||||
}
|
||||
|
||||
impl InterceptorError {
|
||||
pub fn new(category: InterceptorErrorCategory, diagnostic: impl Into<String>) -> Self {
|
||||
let mut diagnostic = diagnostic.into();
|
||||
if diagnostic.len() > MAX_INTERCEPTOR_DIAGNOSTIC_BYTES {
|
||||
let mut end = MAX_INTERCEPTOR_DIAGNOSTIC_BYTES;
|
||||
while !diagnostic.is_char_boundary(end) {
|
||||
end -= 1;
|
||||
}
|
||||
diagnostic.truncate(end);
|
||||
}
|
||||
Self {
|
||||
category,
|
||||
diagnostic,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn category(&self) -> InterceptorErrorCategory {
|
||||
self.category
|
||||
}
|
||||
|
||||
pub fn diagnostic(&self) -> &str {
|
||||
&self.diagnostic
|
||||
}
|
||||
}
|
||||
|
||||
/// The lifecycle phase at which an interceptor callback executes.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
|
||||
pub enum InterceptorPhase {
|
||||
#[default]
|
||||
PromptSubmit,
|
||||
PendingHistoryAppends,
|
||||
PreLlmRequest,
|
||||
PreToolCall,
|
||||
PostToolCall,
|
||||
AssistantTurnEnd,
|
||||
RunExit,
|
||||
}
|
||||
|
||||
impl std::fmt::Display for InterceptorPhase {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter.write_str(match self {
|
||||
Self::PromptSubmit => "prompt_submit",
|
||||
Self::PendingHistoryAppends => "pending_history_appends",
|
||||
Self::PreLlmRequest => "pre_llm_request",
|
||||
Self::PreToolCall => "pre_tool_call",
|
||||
Self::PostToolCall => "post_tool_call",
|
||||
Self::AssistantTurnEnd => "assistant_turn_end",
|
||||
Self::RunExit => "run_exit",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
|
||||
pub struct InterceptorRunId(pub u64);
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
|
||||
pub struct InterceptorTurnId(pub u64);
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
|
||||
pub enum InterceptorCallId {
|
||||
Llm(u64),
|
||||
Tool(String),
|
||||
}
|
||||
|
||||
/// Saturating public counter used by interceptor contexts.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Default)]
|
||||
pub struct InterceptorCounter(u32);
|
||||
|
||||
impl InterceptorCounter {
|
||||
pub fn from_usize(value: usize) -> Self {
|
||||
Self(u32::try_from(value).unwrap_or(u32::MAX))
|
||||
}
|
||||
|
||||
pub fn get(self) -> u32 {
|
||||
self.0
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
|
||||
pub struct InterceptorCounters {
|
||||
pub invocation: InterceptorCounter,
|
||||
pub engine_turn: InterceptorCounter,
|
||||
pub run_turn: InterceptorCounter,
|
||||
pub llm_call: InterceptorCounter,
|
||||
pub tool_batch: InterceptorCounter,
|
||||
pub tool_call: InterceptorCounter,
|
||||
}
|
||||
|
||||
/// Identity, phase, and bounded counters common to every lifecycle callback.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Default)]
|
||||
pub struct InterceptorInvocation {
|
||||
pub run_id: InterceptorRunId,
|
||||
pub turn_id: Option<InterceptorTurnId>,
|
||||
pub call_id: Option<InterceptorCallId>,
|
||||
pub phase: InterceptorPhase,
|
||||
pub counters: InterceptorCounters,
|
||||
}
|
||||
|
||||
/// An interceptor failure bound to the exact Engine lifecycle phase that ran it.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
|
||||
#[error("{phase} interceptor failed: {error}")]
|
||||
pub struct InterceptorFailure {
|
||||
phase: InterceptorPhase,
|
||||
#[source]
|
||||
error: InterceptorError,
|
||||
}
|
||||
|
||||
impl InterceptorFailure {
|
||||
pub(crate) fn new(phase: InterceptorPhase, error: InterceptorError) -> Self {
|
||||
Self { phase, error }
|
||||
}
|
||||
|
||||
pub fn phase(&self) -> InterceptorPhase {
|
||||
self.phase
|
||||
}
|
||||
|
||||
pub fn error(&self) -> &InterceptorError {
|
||||
&self.error
|
||||
}
|
||||
}
|
||||
|
||||
pub type InterceptorResult<T> = Result<T, InterceptorError>;
|
||||
|
||||
// =============================================================================
|
||||
// Lifecycle Contexts
|
||||
// =============================================================================
|
||||
|
||||
pub struct PromptSubmitContext<'a, A = ()> {
|
||||
pub invocation: InterceptorInvocation,
|
||||
pub item: &'a mut Item,
|
||||
pub history: &'a [HistoryEntry<A>],
|
||||
}
|
||||
|
||||
pub struct PendingHistoryAppendsContext<'a, A = ()> {
|
||||
pub invocation: InterceptorInvocation,
|
||||
pub history: &'a [HistoryEntry<A>],
|
||||
}
|
||||
|
||||
pub struct PreLlmRequestContext<'a, A = ()> {
|
||||
pub invocation: InterceptorInvocation,
|
||||
pub items: &'a mut Vec<Item>,
|
||||
pub history: &'a [HistoryEntry<A>],
|
||||
}
|
||||
|
||||
pub struct AssistantTurnEndContext<'a, A = ()> {
|
||||
pub invocation: InterceptorInvocation,
|
||||
pub assistant_entries: &'a [HistoryEntry<A>],
|
||||
pub history: &'a [HistoryEntry<A>],
|
||||
pub tool_calls: &'a [ToolCall],
|
||||
}
|
||||
|
||||
pub struct RunExitContext<'a, A = ()> {
|
||||
pub invocation: InterceptorInvocation,
|
||||
pub exit: &'a EngineRunExit,
|
||||
pub history: &'a [HistoryEntry<A>],
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Action Enums
|
||||
// =============================================================================
|
||||
@@ -86,9 +280,9 @@ pub enum PostToolAction {
|
||||
/// Action at the end of a turn (when LLM produces no tool calls).
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum TurnEndAction {
|
||||
/// Turn is finished, return to caller.
|
||||
/// Accept the Engine's natural next phase: execute tools, or finish when none exist.
|
||||
Finish,
|
||||
/// Continue with additional messages injected into history.
|
||||
/// Commit additional messages, then continue through the natural next phase.
|
||||
ContinueWithMessages(Vec<Item>),
|
||||
/// Pause execution (can be resumed later).
|
||||
Pause,
|
||||
@@ -99,8 +293,9 @@ pub enum TurnEndAction {
|
||||
// =============================================================================
|
||||
|
||||
/// Context for pre-tool-call decisions.
|
||||
pub struct ToolCallInfo {
|
||||
/// Tool call information (modifiable).
|
||||
pub struct ToolCallInfo<'a, A = ()> {
|
||||
pub invocation: InterceptorInvocation,
|
||||
pub history: &'a [HistoryEntry<A>],
|
||||
pub call: ToolCall,
|
||||
/// Tool meta information.
|
||||
pub meta: ToolMeta,
|
||||
@@ -111,10 +306,11 @@ pub struct ToolCallInfo {
|
||||
}
|
||||
|
||||
/// Context for post-tool-call decisions.
|
||||
pub struct ToolResultInfo {
|
||||
/// Original tool call.
|
||||
pub struct ToolResultInfo<'a, A = ()> {
|
||||
pub invocation: InterceptorInvocation,
|
||||
pub history: &'a [HistoryEntry<A>],
|
||||
pub call: ToolCall,
|
||||
/// Tool execution result (modifiable).
|
||||
/// Committed terminal tool execution result.
|
||||
pub result: ToolResult,
|
||||
/// Tool meta information.
|
||||
pub meta: ToolMeta,
|
||||
@@ -130,14 +326,22 @@ pub struct ToolResultInfo {
|
||||
|
||||
/// Intercepts the Engine execution loop at key decision points.
|
||||
///
|
||||
/// All methods have default implementations that let the Engine
|
||||
/// proceed without intervention. Callers provide richer implementations for
|
||||
/// approval flows, permission checks, etc.
|
||||
/// Every lifecycle method is asynchronous and returns [`InterceptorResult`],
|
||||
/// keeping implementation failure separate from the method's control-flow
|
||||
/// action. The Engine reports a failure as a typed run interruption annotated
|
||||
/// with the exact [`InterceptorPhase`] that failed.
|
||||
///
|
||||
/// All methods have default implementations that let the Engine proceed
|
||||
/// without intervention. Callers provide richer implementations for approval
|
||||
/// flows, permission checks, and other trusted host adaptation.
|
||||
#[async_trait]
|
||||
pub trait Interceptor: Send + Sync {
|
||||
/// Called after receiving user input, before adding to history.
|
||||
async fn on_prompt_submit(&self, _item: &mut Item) -> PromptAction {
|
||||
PromptAction::Continue
|
||||
pub trait Interceptor<A: Send + Sync = ()>: Send + Sync {
|
||||
/// Called after receiving user input, before adding it to Engine history.
|
||||
async fn on_prompt_submit(
|
||||
&self,
|
||||
_context: PromptSubmitContext<'_, A>,
|
||||
) -> InterceptorResult<PromptAction> {
|
||||
Ok(PromptAction::Continue)
|
||||
}
|
||||
|
||||
/// Items that should be **committed to `engine.history`** just
|
||||
@@ -158,7 +362,10 @@ pub trait Interceptor: Send + Sync {
|
||||
/// reproducible per-request transformations (pruning, content
|
||||
/// trimming, cache anchors) that depend only on the existing
|
||||
/// history.
|
||||
async fn pending_history_appends(&self) -> Result<Vec<Item>, String> {
|
||||
async fn pending_history_appends(
|
||||
&self,
|
||||
_context: PendingHistoryAppendsContext<'_, A>,
|
||||
) -> InterceptorResult<Vec<Item>> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
@@ -170,27 +377,42 @@ pub trait Interceptor: Send + Sync {
|
||||
/// If an interceptor derives a human/model-visible nudge from the current
|
||||
/// request context, return [`PreRequestAction::ContinueWith`] so the Engine
|
||||
/// commits it to history before the request is sent.
|
||||
async fn pre_llm_request(&self, _context: &mut Vec<Item>) -> PreRequestAction {
|
||||
PreRequestAction::Continue
|
||||
async fn pre_llm_request(
|
||||
&self,
|
||||
_context: PreLlmRequestContext<'_, A>,
|
||||
) -> InterceptorResult<PreRequestAction> {
|
||||
Ok(PreRequestAction::Continue)
|
||||
}
|
||||
|
||||
/// Called before each tool is executed.
|
||||
async fn pre_tool_call(&self, _info: &mut ToolCallInfo) -> PreToolAction {
|
||||
PreToolAction::Continue
|
||||
async fn pre_tool_call(
|
||||
&self,
|
||||
_info: &mut ToolCallInfo<'_, A>,
|
||||
) -> InterceptorResult<PreToolAction> {
|
||||
Ok(PreToolAction::Continue)
|
||||
}
|
||||
|
||||
/// Called after each tool completes.
|
||||
async fn post_tool_call(&self, _info: &mut ToolResultInfo) -> PostToolAction {
|
||||
PostToolAction::Continue
|
||||
/// Called after each tool reaches one terminal result and that result is committed.
|
||||
async fn post_tool_call(
|
||||
&self,
|
||||
_info: &ToolResultInfo<'_, A>,
|
||||
) -> InterceptorResult<PostToolAction> {
|
||||
Ok(PostToolAction::Continue)
|
||||
}
|
||||
|
||||
/// Called when a turn ends with no tool calls.
|
||||
async fn on_turn_end(&self, _history: &[Item]) -> TurnEndAction {
|
||||
TurnEndAction::Finish
|
||||
/// Called after every terminal assistant response is committed and before
|
||||
/// the Engine decides whether to execute tools, continue, or finish.
|
||||
async fn on_assistant_turn_end(
|
||||
&self,
|
||||
_context: AssistantTurnEndContext<'_, A>,
|
||||
) -> InterceptorResult<TurnEndAction> {
|
||||
Ok(TurnEndAction::Finish)
|
||||
}
|
||||
|
||||
/// Called when execution is interrupted (abort or cancel).
|
||||
async fn on_abort(&self, _reason: &str) {}
|
||||
/// Called once for the terminal outcome of each public run or resume call.
|
||||
async fn on_run_exit(&self, _context: RunExitContext<'_, A>) -> InterceptorResult<()> {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// Default interceptor: no intervention. Engine proceeds through the loop
|
||||
@@ -198,4 +420,4 @@ pub trait Interceptor: Send + Sync {
|
||||
pub(crate) struct DefaultInterceptor;
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for DefaultInterceptor {}
|
||||
impl<A: Send + Sync> Interceptor<A> for DefaultInterceptor {}
|
||||
|
||||
@@ -22,11 +22,17 @@ pub use agen_macros::{description, tool, tool_registry};
|
||||
pub use callback::{TextBlockScope, ThinkingBlockScope, ToolUseBlockScope};
|
||||
pub use engine::{
|
||||
Engine, EngineConfig, EngineError, EngineResult, EngineRunExit, EngineRunOutput,
|
||||
LlmRetryNotice, StopReason, ToolRegistryError,
|
||||
LlmRetryNotice, RunInterruptionReason, ToolRegistryError,
|
||||
};
|
||||
pub use handler::ToolUseBlockStart;
|
||||
pub use history::{History, HistoryEntry};
|
||||
pub use interceptor::Interceptor;
|
||||
pub use interceptor::{
|
||||
AssistantTurnEndContext, Interceptor, InterceptorCallId, InterceptorCounter,
|
||||
InterceptorCounters, InterceptorError, InterceptorErrorCategory, InterceptorFailure,
|
||||
InterceptorInvocation, InterceptorPhase, InterceptorResult, InterceptorRunId,
|
||||
InterceptorTurnId, MAX_INTERCEPTOR_DIAGNOSTIC_BYTES, PendingHistoryAppendsContext,
|
||||
PreLlmRequestContext, PromptSubmitContext, RunExitContext,
|
||||
};
|
||||
pub use message::{ContentPart, Item, Message, Role};
|
||||
pub use tool::{
|
||||
ToolCall, ToolExecutionContext, ToolExecutionHandle, ToolExecutionPolicy,
|
||||
|
||||
@@ -1,8 +1,15 @@
|
||||
mod common;
|
||||
|
||||
use agen::interceptor::{
|
||||
AssistantTurnEndContext, Interceptor, InterceptorCallId, InterceptorInvocation,
|
||||
InterceptorPhase, InterceptorResult, PendingHistoryAppendsContext, PreLlmRequestContext,
|
||||
PreRequestAction, PromptAction, PromptSubmitContext, RunExitContext, TurnEndAction,
|
||||
};
|
||||
use agen::llm_client::event::{Event, ResponseStatus, StatusEvent};
|
||||
use agen::{Engine, EngineError, History, HistoryEntry, Item, Role};
|
||||
use async_trait::async_trait;
|
||||
use common::MockLlmClient;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
fn completed_text_events(text: &str) -> Vec<Event> {
|
||||
vec![
|
||||
@@ -47,6 +54,125 @@ async fn run_preserves_item_annotations_without_projecting_them() {
|
||||
assert_eq!(history.items_cloned().len(), 2);
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct AnnotationObservingInterceptor {
|
||||
observed: Arc<Mutex<Vec<(InterceptorInvocation, Vec<String>)>>>,
|
||||
}
|
||||
|
||||
impl AnnotationObservingInterceptor {
|
||||
fn record(&self, invocation: &InterceptorInvocation, history: &[HistoryEntry<String>]) {
|
||||
self.observed.lock().unwrap().push((
|
||||
invocation.clone(),
|
||||
history
|
||||
.iter()
|
||||
.map(|entry| entry.annotation.clone())
|
||||
.collect(),
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor<String> for AnnotationObservingInterceptor {
|
||||
async fn on_prompt_submit(
|
||||
&self,
|
||||
context: PromptSubmitContext<'_, String>,
|
||||
) -> InterceptorResult<PromptAction> {
|
||||
self.record(&context.invocation, context.history);
|
||||
Ok(PromptAction::Continue)
|
||||
}
|
||||
|
||||
async fn pending_history_appends(
|
||||
&self,
|
||||
context: PendingHistoryAppendsContext<'_, String>,
|
||||
) -> InterceptorResult<Vec<Item>> {
|
||||
self.record(&context.invocation, context.history);
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
async fn pre_llm_request(
|
||||
&self,
|
||||
context: PreLlmRequestContext<'_, String>,
|
||||
) -> InterceptorResult<PreRequestAction> {
|
||||
self.record(&context.invocation, context.history);
|
||||
Ok(PreRequestAction::Continue)
|
||||
}
|
||||
|
||||
async fn on_assistant_turn_end(
|
||||
&self,
|
||||
context: AssistantTurnEndContext<'_, String>,
|
||||
) -> InterceptorResult<TurnEndAction> {
|
||||
assert_eq!(context.assistant_entries.len(), 1);
|
||||
assert_eq!(context.assistant_entries[0].annotation, "2:assistant");
|
||||
self.record(&context.invocation, context.history);
|
||||
Ok(TurnEndAction::Finish)
|
||||
}
|
||||
|
||||
async fn on_run_exit(&self, context: RunExitContext<'_, String>) -> InterceptorResult<()> {
|
||||
self.record(&context.invocation, context.history);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn interceptor_contexts_preserve_annotations_and_typed_lifecycle_identity() {
|
||||
let client = MockLlmClient::new(completed_text_events("assistant reply"));
|
||||
let mut engine = Engine::<_, agen::state::Mutable, String>::new_annotated(client);
|
||||
let observed = Arc::new(Mutex::new(Vec::new()));
|
||||
engine.set_interceptor(AnnotationObservingInterceptor {
|
||||
observed: observed.clone(),
|
||||
});
|
||||
let mut history = History::<String>::new();
|
||||
let mut next = 0usize;
|
||||
let mut annotate = |item: &Item| {
|
||||
next += 1;
|
||||
let kind = if item.is_assistant_message() {
|
||||
"assistant"
|
||||
} else {
|
||||
"user"
|
||||
};
|
||||
Ok(format!("{next}:{kind}"))
|
||||
};
|
||||
|
||||
let output = engine
|
||||
.run_with_annotation(&mut history, "hello", &mut annotate)
|
||||
.await;
|
||||
assert!(matches!(output.result, agen::EngineRunExit::Finished));
|
||||
|
||||
let observed = observed.lock().unwrap();
|
||||
let phases: Vec<_> = observed
|
||||
.iter()
|
||||
.map(|(invocation, _)| invocation.phase)
|
||||
.collect();
|
||||
assert_eq!(
|
||||
phases,
|
||||
[
|
||||
InterceptorPhase::PromptSubmit,
|
||||
InterceptorPhase::PendingHistoryAppends,
|
||||
InterceptorPhase::PreLlmRequest,
|
||||
InterceptorPhase::AssistantTurnEnd,
|
||||
InterceptorPhase::RunExit,
|
||||
]
|
||||
);
|
||||
assert!(
|
||||
observed
|
||||
.iter()
|
||||
.all(|(invocation, _)| invocation.run_id == observed[0].0.run_id)
|
||||
);
|
||||
assert_eq!(
|
||||
observed
|
||||
.iter()
|
||||
.map(|(invocation, _)| invocation.counters.invocation.get())
|
||||
.collect::<Vec<_>>(),
|
||||
[0, 1, 2, 3, 4]
|
||||
);
|
||||
assert_eq!(observed[2].0.call_id, Some(InterceptorCallId::Llm(0)));
|
||||
assert_eq!(observed[3].0.call_id, Some(InterceptorCallId::Llm(0)));
|
||||
assert_eq!(observed[1].1, ["1:user"]);
|
||||
assert_eq!(observed[2].1, ["1:user"]);
|
||||
assert_eq!(observed[3].1, ["1:user", "2:assistant"]);
|
||||
assert_eq!(observed[4].1, ["1:user", "2:assistant"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn append_failure_does_not_make_item_live() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
|
||||
@@ -10,11 +10,18 @@ use std::sync::{Arc, Mutex};
|
||||
|
||||
use agen::Item;
|
||||
use agen::interceptor::{
|
||||
Interceptor, PreRequestAction, PreToolAction, ToolCallInfo, TurnEndAction,
|
||||
AssistantTurnEndContext, Interceptor, InterceptorError, InterceptorErrorCategory,
|
||||
InterceptorPhase as InterceptorPoint, InterceptorResult, MAX_INTERCEPTOR_DIAGNOSTIC_BYTES,
|
||||
PendingHistoryAppendsContext, PostToolAction, PreLlmRequestContext, PreRequestAction,
|
||||
PreToolAction, PromptAction, PromptSubmitContext, RunExitContext, ToolCallInfo, ToolResultInfo,
|
||||
TurnEndAction,
|
||||
};
|
||||
use agen::llm_client::{
|
||||
ClientError, LlmClient, Request, ResponseStream,
|
||||
event::{Event, ResponseStatus, StatusEvent},
|
||||
};
|
||||
use agen::llm_client::event::{Event, ResponseStatus, StatusEvent};
|
||||
use agen::tool::{Tool, ToolDefinition, ToolError, ToolMeta, ToolOutput};
|
||||
use agen::{Engine, EngineError, EngineRunExit, History, StopReason};
|
||||
use agen::{Engine, EngineError, EngineRunExit, History, RunInterruptionReason};
|
||||
use async_trait::async_trait;
|
||||
use common::MockLlmClient;
|
||||
|
||||
@@ -205,7 +212,7 @@ async fn history_append_failure_stops_before_tool_execution() {
|
||||
let exit = engine.run(&mut history, "use the tool").await;
|
||||
|
||||
assert!(
|
||||
matches!(exit, EngineRunExit::Interrupted(StopReason::Unexpected(EngineError::HistoryAppend(ref message))) if message == "simulated ENOSPC")
|
||||
matches!(exit, EngineRunExit::Interrupted(RunInterruptionReason::Unexpected(EngineError::HistoryAppend(ref message))) if message == "simulated ENOSPC")
|
||||
);
|
||||
assert_eq!(tool.call_count(), 0);
|
||||
assert_eq!(history.len(), 1);
|
||||
@@ -613,12 +620,15 @@ struct YieldOnce {
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for YieldOnce {
|
||||
async fn pre_llm_request(&self, _context: &mut Vec<Item>) -> PreRequestAction {
|
||||
if self.calls.fetch_add(1, Ordering::SeqCst) == 0 {
|
||||
async fn pre_llm_request(
|
||||
&self,
|
||||
_context: PreLlmRequestContext<'_, ()>,
|
||||
) -> InterceptorResult<PreRequestAction> {
|
||||
Ok(if self.calls.fetch_add(1, Ordering::SeqCst) == 0 {
|
||||
PreRequestAction::Yield
|
||||
} else {
|
||||
PreRequestAction::Continue
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -628,12 +638,15 @@ struct PauseToolOnce {
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for PauseToolOnce {
|
||||
async fn pre_tool_call(&self, _info: &mut ToolCallInfo) -> PreToolAction {
|
||||
if self.calls.fetch_add(1, Ordering::SeqCst) == 0 {
|
||||
async fn pre_tool_call(
|
||||
&self,
|
||||
_info: &mut ToolCallInfo<'_, ()>,
|
||||
) -> InterceptorResult<PreToolAction> {
|
||||
Ok(if self.calls.fetch_add(1, Ordering::SeqCst) == 0 {
|
||||
PreToolAction::Pause
|
||||
} else {
|
||||
PreToolAction::Continue
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -643,13 +656,509 @@ struct ContinueTurnOnce {
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for ContinueTurnOnce {
|
||||
async fn on_turn_end(&self, _history: &[Item]) -> TurnEndAction {
|
||||
if self.calls.fetch_add(1, Ordering::SeqCst) == 0 {
|
||||
async fn on_assistant_turn_end(
|
||||
&self,
|
||||
_context: AssistantTurnEndContext<'_, ()>,
|
||||
) -> InterceptorResult<TurnEndAction> {
|
||||
Ok(if self.calls.fetch_add(1, Ordering::SeqCst) == 0 {
|
||||
TurnEndAction::ContinueWithMessages(vec![Item::system_message("continue")])
|
||||
} else {
|
||||
TurnEndAction::Finish
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
struct FailingLifecycleInterceptor {
|
||||
failure: InterceptorPoint,
|
||||
calls: Arc<Mutex<Vec<InterceptorPoint>>>,
|
||||
}
|
||||
|
||||
impl FailingLifecycleInterceptor {
|
||||
fn new(failure: InterceptorPoint) -> Self {
|
||||
Self {
|
||||
failure,
|
||||
calls: Arc::new(Mutex::new(Vec::new())),
|
||||
}
|
||||
}
|
||||
|
||||
fn record<T>(&self, point: InterceptorPoint, action: T) -> InterceptorResult<T> {
|
||||
self.calls.lock().unwrap().push(point);
|
||||
if self.failure == point {
|
||||
Err(InterceptorError::new(
|
||||
InterceptorErrorCategory::Policy,
|
||||
format!("{point} rejected"),
|
||||
))
|
||||
} else {
|
||||
Ok(action)
|
||||
}
|
||||
}
|
||||
|
||||
fn calls(&self) -> Vec<InterceptorPoint> {
|
||||
self.calls.lock().unwrap().clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for FailingLifecycleInterceptor {
|
||||
async fn on_prompt_submit(
|
||||
&self,
|
||||
_context: PromptSubmitContext<'_, ()>,
|
||||
) -> InterceptorResult<PromptAction> {
|
||||
tokio::task::yield_now().await;
|
||||
self.record(InterceptorPoint::PromptSubmit, PromptAction::Continue)
|
||||
}
|
||||
|
||||
async fn pending_history_appends(
|
||||
&self,
|
||||
_context: PendingHistoryAppendsContext<'_, ()>,
|
||||
) -> InterceptorResult<Vec<Item>> {
|
||||
tokio::task::yield_now().await;
|
||||
self.record(InterceptorPoint::PendingHistoryAppends, Vec::new())
|
||||
}
|
||||
|
||||
async fn pre_llm_request(
|
||||
&self,
|
||||
_context: PreLlmRequestContext<'_, ()>,
|
||||
) -> InterceptorResult<PreRequestAction> {
|
||||
tokio::task::yield_now().await;
|
||||
self.record(InterceptorPoint::PreLlmRequest, PreRequestAction::Continue)
|
||||
}
|
||||
|
||||
async fn pre_tool_call(
|
||||
&self,
|
||||
_info: &mut ToolCallInfo<'_, ()>,
|
||||
) -> InterceptorResult<PreToolAction> {
|
||||
tokio::task::yield_now().await;
|
||||
self.record(InterceptorPoint::PreToolCall, PreToolAction::Continue)
|
||||
}
|
||||
|
||||
async fn post_tool_call(
|
||||
&self,
|
||||
_info: &ToolResultInfo<'_, ()>,
|
||||
) -> InterceptorResult<PostToolAction> {
|
||||
tokio::task::yield_now().await;
|
||||
self.record(InterceptorPoint::PostToolCall, PostToolAction::Continue)
|
||||
}
|
||||
|
||||
async fn on_assistant_turn_end(
|
||||
&self,
|
||||
context: AssistantTurnEndContext<'_, ()>,
|
||||
) -> InterceptorResult<TurnEndAction> {
|
||||
tokio::task::yield_now().await;
|
||||
assert!(context.history.ends_with(context.assistant_entries));
|
||||
if !context.tool_calls.is_empty() {
|
||||
assert_eq!(
|
||||
context
|
||||
.assistant_entries
|
||||
.iter()
|
||||
.filter(|entry| matches!(&entry.item, Item::ToolCall { .. }))
|
||||
.count(),
|
||||
context.tool_calls.len()
|
||||
);
|
||||
}
|
||||
self.record(InterceptorPoint::AssistantTurnEnd, TurnEndAction::Finish)
|
||||
}
|
||||
|
||||
async fn on_run_exit(&self, _context: RunExitContext<'_, ()>) -> InterceptorResult<()> {
|
||||
tokio::task::yield_now().await;
|
||||
self.record(InterceptorPoint::RunExit, ())
|
||||
}
|
||||
}
|
||||
|
||||
fn expected_interceptor_calls(failure: InterceptorPoint) -> Vec<InterceptorPoint> {
|
||||
use InterceptorPoint as Point;
|
||||
|
||||
let mut calls = match failure {
|
||||
Point::PromptSubmit => vec![Point::PromptSubmit],
|
||||
Point::PendingHistoryAppends => {
|
||||
vec![Point::PromptSubmit, Point::PendingHistoryAppends]
|
||||
}
|
||||
Point::PreLlmRequest => vec![
|
||||
Point::PromptSubmit,
|
||||
Point::PendingHistoryAppends,
|
||||
Point::PreLlmRequest,
|
||||
],
|
||||
Point::PreToolCall => vec![
|
||||
Point::PromptSubmit,
|
||||
Point::PendingHistoryAppends,
|
||||
Point::PreLlmRequest,
|
||||
Point::AssistantTurnEnd,
|
||||
Point::PreToolCall,
|
||||
],
|
||||
Point::PostToolCall => vec![
|
||||
Point::PromptSubmit,
|
||||
Point::PendingHistoryAppends,
|
||||
Point::PreLlmRequest,
|
||||
Point::AssistantTurnEnd,
|
||||
Point::PreToolCall,
|
||||
Point::PostToolCall,
|
||||
],
|
||||
Point::AssistantTurnEnd => vec![
|
||||
Point::PromptSubmit,
|
||||
Point::PendingHistoryAppends,
|
||||
Point::PreLlmRequest,
|
||||
Point::AssistantTurnEnd,
|
||||
],
|
||||
Point::RunExit => vec![
|
||||
Point::PromptSubmit,
|
||||
Point::PendingHistoryAppends,
|
||||
Point::PreLlmRequest,
|
||||
Point::AssistantTurnEnd,
|
||||
],
|
||||
};
|
||||
calls.push(Point::RunExit);
|
||||
calls
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn interceptor_failures_are_typed_and_terminal_observer_preserves_original_exit() {
|
||||
use InterceptorPoint as Point;
|
||||
|
||||
for failure_point in [
|
||||
Point::PromptSubmit,
|
||||
Point::PendingHistoryAppends,
|
||||
Point::PreLlmRequest,
|
||||
Point::PreToolCall,
|
||||
Point::PostToolCall,
|
||||
Point::AssistantTurnEnd,
|
||||
Point::RunExit,
|
||||
] {
|
||||
let interceptor = FailingLifecycleInterceptor::new(failure_point);
|
||||
let needs_tool = matches!(failure_point, Point::PreToolCall | Point::PostToolCall);
|
||||
let events = if needs_tool {
|
||||
vec![
|
||||
Event::tool_use_start(0, "call-1", "count_tool"),
|
||||
Event::tool_input_delta(0, "{}"),
|
||||
Event::tool_use_stop(0),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
]
|
||||
} else {
|
||||
completed_text_events()
|
||||
};
|
||||
let mut engine = Engine::new(MockLlmClient::new(events));
|
||||
engine.register_tool(CountingTool::new("count_tool").definition());
|
||||
engine.set_interceptor(interceptor.clone());
|
||||
let mut history = History::new();
|
||||
let mut engine = engine.lock(&history);
|
||||
|
||||
let exit = engine.run(&mut history, "test").await;
|
||||
let failure = if failure_point == Point::RunExit {
|
||||
assert!(matches!(exit, EngineRunExit::Finished));
|
||||
engine
|
||||
.last_run_exit_observer_failure()
|
||||
.expect("terminal observer diagnostic should be retained")
|
||||
} else {
|
||||
let EngineRunExit::Interrupted(RunInterruptionReason::Unexpected(
|
||||
EngineError::Interceptor(failure),
|
||||
)) = &exit
|
||||
else {
|
||||
panic!("expected typed interceptor interruption at {failure_point}, got {exit:?}");
|
||||
};
|
||||
failure
|
||||
};
|
||||
assert_eq!(failure.phase(), failure_point);
|
||||
assert_eq!(
|
||||
failure.error().diagnostic(),
|
||||
format!("{failure_point} rejected")
|
||||
);
|
||||
assert_eq!(
|
||||
interceptor.calls(),
|
||||
expected_interceptor_calls(failure_point)
|
||||
);
|
||||
if failure_point == Point::PostToolCall {
|
||||
assert!(
|
||||
history
|
||||
.items()
|
||||
.any(|item| matches!(item, Item::ToolResult { .. })),
|
||||
"post-tool failure must not precede terminal output commit"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn interceptor_error_keeps_typed_category_and_bounded_utf8_diagnostic() {
|
||||
let error = InterceptorError::new(
|
||||
InterceptorErrorCategory::Dependency,
|
||||
"界".repeat(MAX_INTERCEPTOR_DIAGNOSTIC_BYTES),
|
||||
);
|
||||
assert_eq!(error.category(), InterceptorErrorCategory::Dependency);
|
||||
assert!(error.diagnostic().len() <= MAX_INTERCEPTOR_DIAGNOSTIC_BYTES);
|
||||
assert!(
|
||||
error
|
||||
.diagnostic()
|
||||
.is_char_boundary(error.diagnostic().len())
|
||||
);
|
||||
}
|
||||
|
||||
struct FailingRunExitObserver {
|
||||
pause: bool,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for FailingRunExitObserver {
|
||||
async fn on_assistant_turn_end(
|
||||
&self,
|
||||
_context: AssistantTurnEndContext<'_, ()>,
|
||||
) -> InterceptorResult<TurnEndAction> {
|
||||
Ok(if self.pause {
|
||||
TurnEndAction::Pause
|
||||
} else {
|
||||
TurnEndAction::Finish
|
||||
})
|
||||
}
|
||||
|
||||
async fn on_run_exit(&self, _context: RunExitContext<'_, ()>) -> InterceptorResult<()> {
|
||||
Err(InterceptorError::new(
|
||||
InterceptorErrorCategory::Dependency,
|
||||
"terminal audit unavailable",
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn terminal_observer_failure_preserves_paused_and_interrupted_exits() {
|
||||
let mut paused_engine = Engine::new(MockLlmClient::new(completed_text_events()));
|
||||
paused_engine.set_interceptor(FailingRunExitObserver { pause: true });
|
||||
let mut paused_history = History::new();
|
||||
let mut paused_engine = paused_engine.lock(&paused_history);
|
||||
assert!(matches!(
|
||||
paused_engine.run(&mut paused_history, "pause").await,
|
||||
EngineRunExit::Paused
|
||||
));
|
||||
assert_eq!(
|
||||
paused_engine
|
||||
.last_run_exit_observer_failure()
|
||||
.expect("paused observer diagnostic")
|
||||
.error()
|
||||
.category(),
|
||||
InterceptorErrorCategory::Dependency
|
||||
);
|
||||
|
||||
let mut interrupted_engine = Engine::new(MockLlmClient::new(completed_text_events()));
|
||||
interrupted_engine.set_max_turns(Some(0));
|
||||
interrupted_engine.set_interceptor(FailingRunExitObserver { pause: false });
|
||||
let mut interrupted_history = History::new();
|
||||
let mut interrupted_engine = interrupted_engine.lock(&interrupted_history);
|
||||
assert!(matches!(
|
||||
interrupted_engine
|
||||
.run(&mut interrupted_history, "limit")
|
||||
.await,
|
||||
EngineRunExit::Interrupted(RunInterruptionReason::LimitReached)
|
||||
));
|
||||
assert_eq!(
|
||||
interrupted_engine
|
||||
.last_run_exit_observer_failure()
|
||||
.expect("interrupted observer diagnostic")
|
||||
.phase(),
|
||||
InterceptorPoint::RunExit
|
||||
);
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum TerminalMode {
|
||||
Finish,
|
||||
PauseOnce,
|
||||
Yield,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
struct RecordingTerminalInterceptor {
|
||||
mode: TerminalMode,
|
||||
assistant_turns: Arc<AtomicUsize>,
|
||||
exits: Arc<Mutex<Vec<&'static str>>>,
|
||||
}
|
||||
|
||||
impl RecordingTerminalInterceptor {
|
||||
fn new(mode: TerminalMode) -> Self {
|
||||
Self {
|
||||
mode,
|
||||
assistant_turns: Arc::new(AtomicUsize::new(0)),
|
||||
exits: Arc::new(Mutex::new(Vec::new())),
|
||||
}
|
||||
}
|
||||
|
||||
fn exits(&self) -> Vec<&'static str> {
|
||||
self.exits.lock().unwrap().clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for RecordingTerminalInterceptor {
|
||||
async fn pre_llm_request(
|
||||
&self,
|
||||
_context: PreLlmRequestContext<'_, ()>,
|
||||
) -> InterceptorResult<PreRequestAction> {
|
||||
Ok(if self.mode == TerminalMode::Yield {
|
||||
PreRequestAction::Yield
|
||||
} else {
|
||||
PreRequestAction::Continue
|
||||
})
|
||||
}
|
||||
|
||||
async fn on_assistant_turn_end(
|
||||
&self,
|
||||
context: AssistantTurnEndContext<'_, ()>,
|
||||
) -> InterceptorResult<TurnEndAction> {
|
||||
assert!(!context.assistant_entries.is_empty());
|
||||
assert!(
|
||||
context.history.ends_with(context.assistant_entries),
|
||||
"assistant-turn callback must observe committed terminal items"
|
||||
);
|
||||
let turn = self.assistant_turns.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(if self.mode == TerminalMode::PauseOnce && turn == 0 {
|
||||
TurnEndAction::Pause
|
||||
} else {
|
||||
TurnEndAction::Finish
|
||||
})
|
||||
}
|
||||
|
||||
async fn on_run_exit(&self, context: RunExitContext<'_, ()>) -> InterceptorResult<()> {
|
||||
let kind = match context.exit {
|
||||
EngineRunExit::Finished => "finished",
|
||||
EngineRunExit::Paused => "paused",
|
||||
EngineRunExit::Yielded => "yielded",
|
||||
EngineRunExit::Interrupted(RunInterruptionReason::LimitReached) => "limit",
|
||||
EngineRunExit::Interrupted(RunInterruptionReason::ContextWindowExceeded) => "context",
|
||||
EngineRunExit::Interrupted(RunInterruptionReason::Cancelled) => "cancelled",
|
||||
EngineRunExit::Interrupted(RunInterruptionReason::Unexpected(_)) => "unexpected",
|
||||
};
|
||||
self.exits.lock().unwrap().push(kind);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct ContextWindowClient;
|
||||
|
||||
#[async_trait]
|
||||
impl LlmClient for ContextWindowClient {
|
||||
async fn stream(&self, _request: Request) -> Result<ResponseStream, ClientError> {
|
||||
Err(ClientError::ContextWindowExceeded)
|
||||
}
|
||||
|
||||
fn clone_boxed(&self) -> Box<dyn LlmClient> {
|
||||
Box::new(self.clone())
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn terminal_observer_runs_once_for_every_exit_and_interruption_kind() {
|
||||
let finished = RecordingTerminalInterceptor::new(TerminalMode::Finish);
|
||||
let mut engine = Engine::new(MockLlmClient::new(completed_text_events()));
|
||||
engine.set_interceptor(finished.clone());
|
||||
let mut history = History::new();
|
||||
assert!(matches!(
|
||||
engine.lock(&history).run(&mut history, "finish").await,
|
||||
EngineRunExit::Finished
|
||||
));
|
||||
assert_eq!(finished.exits(), ["finished"]);
|
||||
|
||||
let yielded = RecordingTerminalInterceptor::new(TerminalMode::Yield);
|
||||
let mut engine = Engine::new(MockLlmClient::new(completed_text_events()));
|
||||
engine.set_interceptor(yielded.clone());
|
||||
let mut history = History::new();
|
||||
assert!(matches!(
|
||||
engine.lock(&history).run(&mut history, "yield").await,
|
||||
EngineRunExit::Yielded
|
||||
));
|
||||
assert_eq!(yielded.exits(), ["yielded"]);
|
||||
|
||||
let limited = RecordingTerminalInterceptor::new(TerminalMode::Finish);
|
||||
let mut engine = Engine::new(MockLlmClient::new(completed_text_events()));
|
||||
engine.set_max_turns(Some(0));
|
||||
engine.set_interceptor(limited.clone());
|
||||
let mut history = History::new();
|
||||
assert!(matches!(
|
||||
engine.lock(&history).run(&mut history, "limit").await,
|
||||
EngineRunExit::Interrupted(RunInterruptionReason::LimitReached)
|
||||
));
|
||||
assert_eq!(limited.exits(), ["limit"]);
|
||||
|
||||
let cancelled = RecordingTerminalInterceptor::new(TerminalMode::Finish);
|
||||
let mut engine = Engine::new(MockLlmClient::new(completed_text_events()));
|
||||
engine.set_interceptor(cancelled.clone());
|
||||
engine.cancel();
|
||||
let mut history = History::new();
|
||||
assert!(matches!(
|
||||
engine.lock(&history).run(&mut history, "cancel").await,
|
||||
EngineRunExit::Interrupted(RunInterruptionReason::Cancelled)
|
||||
));
|
||||
assert_eq!(cancelled.exits(), ["cancelled"]);
|
||||
|
||||
let context = RecordingTerminalInterceptor::new(TerminalMode::Finish);
|
||||
let mut engine = Engine::new(ContextWindowClient);
|
||||
engine.set_interceptor(context.clone());
|
||||
let mut history = History::new();
|
||||
assert!(matches!(
|
||||
engine.lock(&history).run(&mut history, "context").await,
|
||||
EngineRunExit::Interrupted(RunInterruptionReason::ContextWindowExceeded)
|
||||
));
|
||||
assert_eq!(context.exits(), ["context"]);
|
||||
|
||||
let unexpected = FailingLifecycleInterceptor::new(InterceptorPoint::PromptSubmit);
|
||||
let mut engine = Engine::new(MockLlmClient::new(completed_text_events()));
|
||||
engine.set_interceptor(unexpected.clone());
|
||||
let mut history = History::new();
|
||||
assert!(matches!(
|
||||
engine.lock(&history).run(&mut history, "fail").await,
|
||||
EngineRunExit::Interrupted(RunInterruptionReason::Unexpected(EngineError::Interceptor(
|
||||
_
|
||||
)))
|
||||
));
|
||||
assert_eq!(
|
||||
unexpected
|
||||
.calls()
|
||||
.iter()
|
||||
.filter(|point| **point == InterceptorPoint::RunExit)
|
||||
.count(),
|
||||
1
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn terminal_observer_does_not_duplicate_on_resume() {
|
||||
let interceptor = RecordingTerminalInterceptor::new(TerminalMode::PauseOnce);
|
||||
let first_response = vec![
|
||||
Event::tool_use_start(0, "call-1", "count_tool"),
|
||||
Event::tool_input_delta(0, "{}"),
|
||||
Event::tool_use_stop(0),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
];
|
||||
let client = MockLlmClient::with_responses(vec![first_response, completed_text_events()]);
|
||||
let tool = CountingTool::new("count_tool");
|
||||
let mut engine = Engine::new(client);
|
||||
engine.register_tool(tool.definition());
|
||||
engine.set_interceptor(interceptor.clone());
|
||||
let mut history = History::new();
|
||||
let mut engine = engine.lock(&history);
|
||||
|
||||
assert!(matches!(
|
||||
engine.run(&mut history, "pause").await,
|
||||
EngineRunExit::Paused
|
||||
));
|
||||
assert_eq!(interceptor.exits(), ["paused"]);
|
||||
assert_eq!(
|
||||
tool.call_count(),
|
||||
0,
|
||||
"pause must retain the pending tool phase"
|
||||
);
|
||||
|
||||
assert!(matches!(
|
||||
engine.resume(&mut history).await,
|
||||
EngineRunExit::Finished
|
||||
));
|
||||
assert_eq!(interceptor.exits(), ["paused", "finished"]);
|
||||
assert_eq!(
|
||||
tool.call_count(),
|
||||
1,
|
||||
"resume must execute the retained tool once"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -730,7 +1239,7 @@ async fn paused_tool_resume_does_not_reset_the_consumed_turn_budget() {
|
||||
|
||||
assert!(matches!(
|
||||
engine.resume(&mut history).await,
|
||||
EngineRunExit::Interrupted(StopReason::LimitReached)
|
||||
EngineRunExit::Interrupted(RunInterruptionReason::LimitReached)
|
||||
));
|
||||
assert_eq!(engine.turn_count(), 1);
|
||||
assert_eq!(engine.active_run_turn_count(), None);
|
||||
@@ -785,7 +1294,7 @@ async fn interceptor_continuation_consumes_the_logical_run_budget() {
|
||||
|
||||
assert!(matches!(
|
||||
engine.run(&mut history, "start").await,
|
||||
EngineRunExit::Interrupted(StopReason::LimitReached)
|
||||
EngineRunExit::Interrupted(RunInterruptionReason::LimitReached)
|
||||
));
|
||||
assert_eq!(engine.turn_count(), 1);
|
||||
assert_eq!(engine.llm_call_count(), 1);
|
||||
@@ -803,7 +1312,7 @@ async fn restored_active_run_budget_is_enforced_before_another_llm_call() {
|
||||
|
||||
assert!(matches!(
|
||||
engine.resume(&mut history).await,
|
||||
EngineRunExit::Interrupted(StopReason::LimitReached)
|
||||
EngineRunExit::Interrupted(RunInterruptionReason::LimitReached)
|
||||
));
|
||||
assert_eq!(engine.turn_count(), 7);
|
||||
assert_eq!(engine.llm_call_count(), 0);
|
||||
|
||||
@@ -6,13 +6,18 @@ use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use agen::interceptor::{Interceptor, PostToolAction, PreToolAction, ToolCallInfo, ToolResultInfo};
|
||||
use agen::interceptor::{
|
||||
Interceptor, InterceptorError, InterceptorErrorCategory, InterceptorPhase, InterceptorResult,
|
||||
PostToolAction, PreToolAction, ToolCallInfo, ToolResultInfo,
|
||||
};
|
||||
use agen::llm_client::event::{Event, ResponseStatus, StatusEvent};
|
||||
use agen::tool::{
|
||||
Tool, ToolDefinition, ToolError, ToolExecutionContext, ToolMeta, ToolOutput, ToolResult,
|
||||
ToolResultDisposition,
|
||||
};
|
||||
use agen::{Engine, History, Item, ToolExecutionPolicy};
|
||||
use agen::{
|
||||
Engine, EngineError, EngineRunExit, History, Item, RunInterruptionReason, ToolExecutionPolicy,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
|
||||
mod common;
|
||||
@@ -580,7 +585,7 @@ async fn cooperative_cancellation_commits_bounded_terminal_output() {
|
||||
);
|
||||
assert!(matches!(
|
||||
output.result,
|
||||
agen::EngineRunExit::Interrupted(agen::StopReason::Cancelled)
|
||||
agen::EngineRunExit::Interrupted(agen::RunInterruptionReason::Cancelled)
|
||||
));
|
||||
}
|
||||
|
||||
@@ -905,24 +910,30 @@ async fn test_tool_execution_context_for_skipped_and_synthetic_paths() {
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for ContextPolicy {
|
||||
async fn pre_tool_call(&self, info: &mut ToolCallInfo) -> PreToolAction {
|
||||
async fn pre_tool_call(
|
||||
&self,
|
||||
info: &mut ToolCallInfo<'_, ()>,
|
||||
) -> InterceptorResult<PreToolAction> {
|
||||
self.pre_contexts.lock().unwrap().push(info.context.clone());
|
||||
match info.call.name.as_str() {
|
||||
Ok(match info.call.name.as_str() {
|
||||
"skip_tool" => PreToolAction::Skip,
|
||||
"synthetic_tool" => PreToolAction::SyntheticResult(ToolResult::from_output(
|
||||
&info.call.id,
|
||||
ToolOutput::from("synthetic result".to_string()),
|
||||
)),
|
||||
_ => PreToolAction::Continue,
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
async fn post_tool_call(&self, info: &mut ToolResultInfo) -> PostToolAction {
|
||||
async fn post_tool_call(
|
||||
&self,
|
||||
info: &ToolResultInfo<'_, ()>,
|
||||
) -> InterceptorResult<PostToolAction> {
|
||||
self.post_contexts
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push(info.context.clone());
|
||||
PostToolAction::Continue
|
||||
Ok(PostToolAction::Continue)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -994,12 +1005,15 @@ async fn test_before_tool_call_skip() {
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for BlockingPolicy {
|
||||
async fn pre_tool_call(&self, info: &mut ToolCallInfo) -> PreToolAction {
|
||||
if info.call.name == "blocked_tool" {
|
||||
async fn pre_tool_call(
|
||||
&self,
|
||||
info: &mut ToolCallInfo<'_, ()>,
|
||||
) -> InterceptorResult<PreToolAction> {
|
||||
Ok(if info.call.name == "blocked_tool" {
|
||||
PreToolAction::Skip
|
||||
} else {
|
||||
PreToolAction::Continue
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1021,9 +1035,9 @@ async fn test_before_tool_call_skip() {
|
||||
);
|
||||
}
|
||||
|
||||
/// Hook: post_tool_call - verify that results can be modified
|
||||
/// Hook: post_tool_call - verify that the committed terminal result is observed.
|
||||
#[tokio::test]
|
||||
async fn test_post_tool_call_modification() {
|
||||
async fn test_post_tool_call_observes_committed_result() {
|
||||
// Prepare responses for multiple requests
|
||||
let client = MockLlmClient::with_responses(vec![
|
||||
// First request: tool call
|
||||
@@ -1074,40 +1088,51 @@ async fn test_post_tool_call_modification() {
|
||||
|
||||
engine.register_tool(simple_tool_definition());
|
||||
|
||||
// Policy to modify results
|
||||
struct ModifyingPolicy {
|
||||
modified_content: Arc<std::sync::Mutex<Option<String>>>,
|
||||
// Policy to observe the committed terminal result.
|
||||
struct ObservingPolicy {
|
||||
observed_content: Arc<std::sync::Mutex<Option<String>>>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for ModifyingPolicy {
|
||||
async fn post_tool_call(&self, info: &mut ToolResultInfo) -> PostToolAction {
|
||||
info.result.summary = format!("[Modified] {}", info.result.summary);
|
||||
*self.modified_content.lock().unwrap() = Some(info.result.summary.clone());
|
||||
PostToolAction::Continue
|
||||
impl Interceptor for ObservingPolicy {
|
||||
async fn post_tool_call(
|
||||
&self,
|
||||
info: &ToolResultInfo<'_, ()>,
|
||||
) -> InterceptorResult<PostToolAction> {
|
||||
assert_eq!(info.invocation.phase, InterceptorPhase::PostToolCall);
|
||||
assert_eq!(
|
||||
info.invocation.call_id,
|
||||
Some(agen::InterceptorCallId::Tool(info.call.id.clone()))
|
||||
);
|
||||
assert!(matches!(
|
||||
info.history.last().map(|entry| &entry.item),
|
||||
Some(Item::ToolResult { call_id, .. }) if call_id == &info.call.id
|
||||
));
|
||||
*self.observed_content.lock().unwrap() = Some(info.result.summary.clone());
|
||||
Ok(PostToolAction::Continue)
|
||||
}
|
||||
}
|
||||
|
||||
let modified_content = Arc::new(std::sync::Mutex::new(None));
|
||||
engine.set_interceptor(ModifyingPolicy {
|
||||
modified_content: modified_content.clone(),
|
||||
let observed_content = Arc::new(std::sync::Mutex::new(None));
|
||||
engine.set_interceptor(ObservingPolicy {
|
||||
observed_content: observed_content.clone(),
|
||||
});
|
||||
|
||||
// Mutable::run consumes self, returns (Locked, EngineResult)
|
||||
let result = engine.run(&mut history, "Test modification").await;
|
||||
let result = engine.run(&mut history, "Test observation").await;
|
||||
|
||||
assert!(
|
||||
matches!(result.result, agen::EngineRunExit::Finished),
|
||||
"Engine should complete"
|
||||
);
|
||||
|
||||
// Verify hook was called and content was modified
|
||||
let content = modified_content.lock().unwrap().clone();
|
||||
assert!(content.is_some(), "Hook should have been called");
|
||||
assert!(
|
||||
content.unwrap().contains("[Modified]"),
|
||||
"Result should be modified"
|
||||
);
|
||||
// Verify the interceptor observed the exact committed result.
|
||||
let observed = observed_content.lock().unwrap().clone();
|
||||
assert_eq!(observed.as_deref(), Some("Original Result"));
|
||||
assert!(history.items().any(|item| matches!(
|
||||
item,
|
||||
Item::ToolResult { summary, .. } if summary == "Original Result"
|
||||
)));
|
||||
}
|
||||
|
||||
/// Hook: pre_tool_call synthetic result - skipped tool gets an error result in history.
|
||||
@@ -1143,11 +1168,14 @@ async fn test_before_tool_call_synthetic_result_committed() {
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for SyntheticPolicy {
|
||||
async fn pre_tool_call(&self, info: &mut ToolCallInfo) -> PreToolAction {
|
||||
PreToolAction::SyntheticResult(ToolResult::error(
|
||||
async fn pre_tool_call(
|
||||
&self,
|
||||
info: &mut ToolCallInfo<'_, ()>,
|
||||
) -> InterceptorResult<PreToolAction> {
|
||||
Ok(PreToolAction::SyntheticResult(ToolResult::error(
|
||||
info.call.id.clone(),
|
||||
"permission denied",
|
||||
))
|
||||
)))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1167,6 +1195,80 @@ async fn test_before_tool_call_synthetic_result_committed() {
|
||||
)));
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum InvalidIdentityMode {
|
||||
ContinuedCall,
|
||||
SyntheticResult,
|
||||
}
|
||||
|
||||
struct InvalidIdentityPolicy(InvalidIdentityMode);
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for InvalidIdentityPolicy {
|
||||
async fn pre_tool_call(
|
||||
&self,
|
||||
info: &mut ToolCallInfo<'_, ()>,
|
||||
) -> InterceptorResult<PreToolAction> {
|
||||
assert_eq!(info.invocation.phase, InterceptorPhase::PreToolCall);
|
||||
assert_eq!(
|
||||
info.invocation.call_id,
|
||||
Some(agen::InterceptorCallId::Tool("call_1".to_string()))
|
||||
);
|
||||
assert!(matches!(
|
||||
info.history.last().map(|entry| &entry.item),
|
||||
Some(Item::ToolCall { call_id, .. }) if call_id == "call_1"
|
||||
));
|
||||
Ok(match self.0 {
|
||||
InvalidIdentityMode::ContinuedCall => {
|
||||
info.call.id = "different-call".to_string();
|
||||
PreToolAction::Continue
|
||||
}
|
||||
InvalidIdentityMode::SyntheticResult => PreToolAction::SyntheticResult(
|
||||
ToolResult::error("different-call", "invalid synthetic result"),
|
||||
),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn interceptor_cannot_change_tool_call_identity() {
|
||||
for mode in [
|
||||
InvalidIdentityMode::ContinuedCall,
|
||||
InvalidIdentityMode::SyntheticResult,
|
||||
] {
|
||||
let client = MockLlmClient::new(vec![
|
||||
Event::tool_use_start(0, "call_1", "echo"),
|
||||
Event::tool_input_delta(0, r#"{}"#),
|
||||
Event::tool_use_stop(0),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
]);
|
||||
let mut engine = Engine::new(client);
|
||||
engine.register_tool(SlowTool::new("echo", 1).definition());
|
||||
engine.set_interceptor(InvalidIdentityPolicy(mode));
|
||||
let mut history = History::new();
|
||||
|
||||
let result = engine.run(&mut history, "identity").await;
|
||||
let EngineRunExit::Interrupted(RunInterruptionReason::Unexpected(
|
||||
EngineError::Interceptor(failure),
|
||||
)) = result.result
|
||||
else {
|
||||
panic!("invalid tool identity must interrupt with a typed failure");
|
||||
};
|
||||
assert_eq!(failure.phase(), InterceptorPhase::PreToolCall);
|
||||
assert_eq!(
|
||||
failure.error().category(),
|
||||
InterceptorErrorCategory::ContractViolation
|
||||
);
|
||||
assert!(
|
||||
!history
|
||||
.items()
|
||||
.any(|item| matches!(item, Item::ToolResult { .. }))
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn post_tool_abort_commits_confirmed_result_before_stopping_run() {
|
||||
let client = MockLlmClient::new(vec![
|
||||
@@ -1181,16 +1283,24 @@ async fn post_tool_abort_commits_confirmed_result_before_stopping_run() {
|
||||
let tool = SlowTool::new("confirmed", 1);
|
||||
engine.register_tool(tool.definition());
|
||||
|
||||
struct AbortAfterResult;
|
||||
let observed = Arc::new(Mutex::new(Vec::<&'static str>::new()));
|
||||
struct AbortAfterResult {
|
||||
lifecycle: Arc<Mutex<Vec<&'static str>>>,
|
||||
}
|
||||
#[async_trait]
|
||||
impl Interceptor for AbortAfterResult {
|
||||
async fn post_tool_call(&self, _info: &mut ToolResultInfo) -> PostToolAction {
|
||||
PostToolAction::Abort("policy stopped the run".to_string())
|
||||
async fn post_tool_call(
|
||||
&self,
|
||||
_info: &ToolResultInfo<'_, ()>,
|
||||
) -> InterceptorResult<PostToolAction> {
|
||||
self.lifecycle.lock().unwrap().push("post_tool_call");
|
||||
Ok(PostToolAction::Abort("policy stopped the run".to_string()))
|
||||
}
|
||||
}
|
||||
engine.set_interceptor(AbortAfterResult);
|
||||
engine.set_interceptor(AbortAfterResult {
|
||||
lifecycle: observed.clone(),
|
||||
});
|
||||
|
||||
let observed = Arc::new(Mutex::new(Vec::<&'static str>::new()));
|
||||
let published = observed.clone();
|
||||
engine.on_tool_result(move |_| published.lock().unwrap().push("published"));
|
||||
let committed = observed.clone();
|
||||
@@ -1210,11 +1320,11 @@ async fn post_tool_abort_commits_confirmed_result_before_stopping_run() {
|
||||
assert_eq!(tool.call_count(), 1);
|
||||
assert_eq!(
|
||||
observed.lock().unwrap().as_slice(),
|
||||
["committed", "published", "run-returned"]
|
||||
["committed", "published", "post_tool_call", "run-returned"]
|
||||
);
|
||||
assert!(matches!(
|
||||
output.result,
|
||||
agen::EngineRunExit::Interrupted(agen::StopReason::Unexpected(
|
||||
agen::EngineRunExit::Interrupted(agen::RunInterruptionReason::Unexpected(
|
||||
agen::EngineError::Aborted(ref reason)
|
||||
)) if reason == "policy stopped the run"
|
||||
));
|
||||
@@ -1239,3 +1349,93 @@ async fn post_tool_abort_commits_confirmed_result_before_stopping_run() {
|
||||
} if call_id == "call_confirmed"
|
||||
)));
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum PostToolStopMode {
|
||||
Abort,
|
||||
Failure,
|
||||
}
|
||||
|
||||
struct StopFirstParallelResult(PostToolStopMode);
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for StopFirstParallelResult {
|
||||
async fn post_tool_call(
|
||||
&self,
|
||||
info: &ToolResultInfo<'_, ()>,
|
||||
) -> InterceptorResult<PostToolAction> {
|
||||
if info.call.id != "call_fast" {
|
||||
return Ok(PostToolAction::Continue);
|
||||
}
|
||||
tokio::time::sleep(Duration::from_millis(5)).await;
|
||||
match self.0 {
|
||||
PostToolStopMode::Abort => Ok(PostToolAction::Abort("stop parallel batch".to_string())),
|
||||
PostToolStopMode::Failure => Err(InterceptorError::new(
|
||||
InterceptorErrorCategory::Policy,
|
||||
"reject parallel batch",
|
||||
)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn post_tool_stop_terminalizes_started_parallel_siblings_before_returning() {
|
||||
for mode in [PostToolStopMode::Abort, PostToolStopMode::Failure] {
|
||||
let first_response = vec![
|
||||
Event::tool_use_start(0, "call_fast", "fast"),
|
||||
Event::tool_input_delta(0, r#"{}"#),
|
||||
Event::tool_use_stop(0),
|
||||
Event::tool_use_start(1, "call_ready", "ready"),
|
||||
Event::tool_input_delta(1, r#"{}"#),
|
||||
Event::tool_use_stop(1),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
];
|
||||
let second_response = vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "next run completed"),
|
||||
Event::text_block_stop(0, None),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
];
|
||||
let client = MockLlmClient::with_responses(vec![first_response, second_response]);
|
||||
let mut engine = Engine::new(client);
|
||||
engine.register_tool(SlowTool::new("fast", 0).definition());
|
||||
engine.register_tool(SlowTool::new("ready", 1).definition());
|
||||
engine.set_interceptor(StopFirstParallelResult(mode));
|
||||
let mut history = History::new();
|
||||
|
||||
let output = engine.run(&mut history, "parallel stop").await;
|
||||
match mode {
|
||||
PostToolStopMode::Abort => assert!(matches!(
|
||||
output.result,
|
||||
EngineRunExit::Interrupted(RunInterruptionReason::Unexpected(
|
||||
EngineError::Aborted(ref reason)
|
||||
)) if reason == "stop parallel batch"
|
||||
)),
|
||||
PostToolStopMode::Failure => assert!(matches!(
|
||||
output.result,
|
||||
EngineRunExit::Interrupted(RunInterruptionReason::Unexpected(
|
||||
EngineError::Interceptor(ref failure)
|
||||
)) if failure.phase() == InterceptorPhase::PostToolCall
|
||||
)),
|
||||
}
|
||||
|
||||
let terminal_ids: Vec<_> = history
|
||||
.iter()
|
||||
.filter_map(|entry| match &entry.item {
|
||||
Item::ToolResult { call_id, .. } => Some(call_id.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
assert_eq!(terminal_ids.len(), 2);
|
||||
assert!(terminal_ids.contains(&"call_fast"));
|
||||
assert!(terminal_ids.contains(&"call_ready"));
|
||||
|
||||
let mut engine = output.engine;
|
||||
let next = engine.run(&mut history, "next run").await;
|
||||
assert!(matches!(next, EngineRunExit::Finished));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,19 +5,19 @@ edition.workspace = true
|
||||
license.workspace = true
|
||||
|
||||
[dependencies]
|
||||
async-trait.workspace = true
|
||||
chrono = { version = "0.4", default-features = false, features = ["clock"] }
|
||||
protocol = { workspace = true }
|
||||
manifest = { workspace = true }
|
||||
ticket = { workspace = true }
|
||||
futures = { workspace = true }
|
||||
reqwest = { version = "0.13", default-features = false, features = ["blocking", "json", "native-tls"] }
|
||||
serde = { workspace = true }
|
||||
serde_json = { workspace = true }
|
||||
thiserror = { workspace = true }
|
||||
tokio = { workspace = true, features = ["rt", "macros", "net", "io-util", "sync", "time", "process", "fs"] }
|
||||
tokio = { workspace = true, features = ["rt", "macros", "net", "io-util", "sync", "time"] }
|
||||
tokio-tungstenite = { workspace = true }
|
||||
uuid = { workspace = true }
|
||||
workspace-api.workspace = true
|
||||
workdir = { workspace = true }
|
||||
|
||||
[dev-dependencies]
|
||||
tempfile = { workspace = true }
|
||||
|
||||
@@ -0,0 +1,839 @@
|
||||
use chrono::{DateTime, Utc};
|
||||
use reqwest::{Method, StatusCode, Url, redirect};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::BTreeMap;
|
||||
use std::env;
|
||||
use std::fmt;
|
||||
use std::fs::{self, OpenOptions};
|
||||
use std::io::Write as _;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
const TOKEN_FILE_NAME: &str = "backend-tokens.json";
|
||||
const MAX_REDIRECTS: usize = 10;
|
||||
|
||||
#[derive(Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
|
||||
pub struct BackendOrigin(String);
|
||||
|
||||
impl BackendOrigin {
|
||||
pub fn parse(input: &str) -> Result<Self, BackendApiClientError> {
|
||||
let url = Url::parse(input.trim()).map_err(|error| {
|
||||
BackendApiClientError::InvalidBackendOrigin(format!(
|
||||
"Backend URL is not a valid absolute URL: {error}"
|
||||
))
|
||||
})?;
|
||||
if !url.path().bytes().all(|byte| byte == b'/')
|
||||
|| url.query().is_some()
|
||||
|| url.fragment().is_some()
|
||||
{
|
||||
return Err(BackendApiClientError::InvalidBackendOrigin(
|
||||
"Backend URL must contain only an origin, without a path, query, or fragment"
|
||||
.to_string(),
|
||||
));
|
||||
}
|
||||
Self::from_url(url)
|
||||
}
|
||||
|
||||
fn from_url(mut url: Url) -> Result<Self, BackendApiClientError> {
|
||||
if !matches!(url.scheme(), "http" | "https") {
|
||||
return Err(BackendApiClientError::InvalidBackendOrigin(
|
||||
"Backend URL scheme must be http or https".to_string(),
|
||||
));
|
||||
}
|
||||
if !url.username().is_empty() || url.password().is_some() {
|
||||
return Err(BackendApiClientError::InvalidBackendOrigin(
|
||||
"Backend URL must not contain user information".to_string(),
|
||||
));
|
||||
}
|
||||
if url.host().is_none() {
|
||||
return Err(BackendApiClientError::InvalidBackendOrigin(
|
||||
"Backend URL must contain a host".to_string(),
|
||||
));
|
||||
}
|
||||
let default_port = match url.scheme() {
|
||||
"http" => 80,
|
||||
"https" => 443,
|
||||
_ => unreachable!("validated Backend URL scheme"),
|
||||
};
|
||||
if url.port() == Some(default_port) {
|
||||
url.set_port(None).map_err(|()| {
|
||||
BackendApiClientError::InvalidBackendOrigin(
|
||||
"Backend URL contains an invalid port".to_string(),
|
||||
)
|
||||
})?;
|
||||
}
|
||||
url.set_path("");
|
||||
url.set_query(None);
|
||||
url.set_fragment(None);
|
||||
let normalized = url.as_str().trim_end_matches('/').to_string();
|
||||
Ok(Self(normalized))
|
||||
}
|
||||
|
||||
pub fn as_str(&self) -> &str {
|
||||
&self.0
|
||||
}
|
||||
|
||||
fn url(&self, path_and_query: &str) -> Result<Url, BackendApiClientError> {
|
||||
if !path_and_query.starts_with('/') || path_and_query.starts_with("//") {
|
||||
return Err(BackendApiClientError::InvalidRequestPath(
|
||||
"Backend API request path must start with one `/`".to_string(),
|
||||
));
|
||||
}
|
||||
Url::parse(&format!("{}{path_and_query}", self.0)).map_err(|error| {
|
||||
BackendApiClientError::InvalidRequestPath(format!(
|
||||
"Backend API request path is invalid: {error}"
|
||||
))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for BackendOrigin {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.debug_tuple("BackendOrigin").field(&self.0).finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for BackendOrigin {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.write_str(&self.0)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct BackendAccessToken(String);
|
||||
|
||||
impl fmt::Debug for BackendAccessToken {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.write_str("BackendAccessToken([REDACTED])")
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct BackendApiClient {
|
||||
origin: BackendOrigin,
|
||||
access_token: BackendAccessToken,
|
||||
asynchronous: reqwest::Client,
|
||||
}
|
||||
|
||||
impl fmt::Debug for BackendApiClient {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.debug_struct("BackendApiClient")
|
||||
.field("origin", &self.origin)
|
||||
.field("access_token", &self.access_token)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
impl BackendApiClient {
|
||||
pub fn from_stored_token(base_url: &str) -> Result<Self, BackendApiClientError> {
|
||||
let path = backend_token_file_path()?;
|
||||
Self::from_token_file(base_url, &path)
|
||||
}
|
||||
|
||||
fn from_token_file(base_url: &str, path: &Path) -> Result<Self, BackendApiClientError> {
|
||||
let origin = BackendOrigin::parse(base_url)?;
|
||||
let token_file = read_token_file(path)?;
|
||||
let entry = token_file.tokens.get(origin.as_str()).ok_or_else(|| {
|
||||
BackendApiClientError::TokenEntryMissing {
|
||||
origin: origin.clone(),
|
||||
path: path.to_path_buf(),
|
||||
}
|
||||
})?;
|
||||
validate_token_entry(entry, &origin, path)?;
|
||||
Self::new(origin, BackendAccessToken(entry.access_token.clone()))
|
||||
}
|
||||
|
||||
fn new(
|
||||
origin: BackendOrigin,
|
||||
access_token: BackendAccessToken,
|
||||
) -> Result<Self, BackendApiClientError> {
|
||||
let asynchronous = reqwest::Client::builder()
|
||||
.redirect(redirect_policy(origin.clone()))
|
||||
.build()
|
||||
.map_err(BackendApiClientError::Http)?;
|
||||
Ok(Self {
|
||||
origin,
|
||||
access_token,
|
||||
asynchronous,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn origin(&self) -> &BackendOrigin {
|
||||
&self.origin
|
||||
}
|
||||
|
||||
pub fn request(
|
||||
&self,
|
||||
method: Method,
|
||||
path_and_query: &str,
|
||||
) -> Result<reqwest::RequestBuilder, BackendApiClientError> {
|
||||
let url = self.origin.url(path_and_query)?;
|
||||
Ok(self
|
||||
.asynchronous
|
||||
.request(method, url)
|
||||
.bearer_auth(&self.access_token.0))
|
||||
}
|
||||
|
||||
pub fn blocking_request(
|
||||
&self,
|
||||
method: Method,
|
||||
path_and_query: &str,
|
||||
) -> Result<reqwest::blocking::RequestBuilder, BackendApiClientError> {
|
||||
let url = self.origin.url(path_and_query)?;
|
||||
let client = reqwest::blocking::Client::builder()
|
||||
.redirect(redirect_policy(self.origin.clone()))
|
||||
.build()
|
||||
.map_err(BackendApiClientError::Http)?;
|
||||
Ok(client
|
||||
.request(method, url)
|
||||
.bearer_auth(&self.access_token.0))
|
||||
}
|
||||
|
||||
pub(crate) fn authorization_header_value(&self) -> String {
|
||||
format!("Bearer {}", self.access_token.0)
|
||||
}
|
||||
|
||||
pub async fn require_success(
|
||||
&self,
|
||||
response: reqwest::Response,
|
||||
) -> Result<reqwest::Response, BackendApiClientError> {
|
||||
let status = response.status();
|
||||
match status {
|
||||
StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN => {
|
||||
self.check_status(status)?;
|
||||
}
|
||||
status if !status.is_success() => {
|
||||
let detail = response
|
||||
.bytes()
|
||||
.await
|
||||
.ok()
|
||||
.and_then(|body| backend_error_detail(&body));
|
||||
return Err(BackendApiClientError::BackendResponse {
|
||||
origin: self.origin.clone(),
|
||||
status: status.as_u16(),
|
||||
detail,
|
||||
});
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
pub fn check_status(&self, status: StatusCode) -> Result<(), BackendApiClientError> {
|
||||
match status {
|
||||
StatusCode::UNAUTHORIZED => Err(BackendApiClientError::Unauthorized {
|
||||
origin: self.origin.clone(),
|
||||
}),
|
||||
StatusCode::FORBIDDEN => Err(BackendApiClientError::Forbidden {
|
||||
origin: self.origin.clone(),
|
||||
}),
|
||||
status if !status.is_success() => Err(BackendApiClientError::BackendStatus {
|
||||
origin: self.origin.clone(),
|
||||
status: status.as_u16(),
|
||||
}),
|
||||
_ => Ok(()),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn from_access_token_for_test(
|
||||
base_url: &str,
|
||||
access_token: &str,
|
||||
) -> Result<Self, BackendApiClientError> {
|
||||
Self::new(
|
||||
BackendOrigin::parse(base_url)?,
|
||||
BackendAccessToken(access_token.to_string()),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
fn redirect_policy(origin: BackendOrigin) -> redirect::Policy {
|
||||
redirect::Policy::custom(move |attempt| {
|
||||
if attempt.previous().len() >= MAX_REDIRECTS {
|
||||
return attempt.error("Backend request exceeded the redirect limit");
|
||||
}
|
||||
match BackendOrigin::from_url(attempt.url().clone()) {
|
||||
Ok(target_origin) if target_origin == origin => attempt.follow(),
|
||||
Ok(target_origin) => attempt.error(format!(
|
||||
"Backend request refused a cross-origin redirect from {origin} to {target_origin}"
|
||||
)),
|
||||
Err(error) => attempt.error(error.to_string()),
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct BackendErrorBody {
|
||||
message: String,
|
||||
}
|
||||
|
||||
fn backend_error_detail(body: &[u8]) -> Option<String> {
|
||||
serde_json::from_slice::<BackendErrorBody>(body)
|
||||
.ok()
|
||||
.map(|body| body.message)
|
||||
.filter(|message| !message.trim().is_empty())
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum BackendApiClientError {
|
||||
InvalidBackendOrigin(String),
|
||||
InvalidRequestPath(String),
|
||||
ConfigDirectoryUnavailable,
|
||||
TokenFileMissing {
|
||||
path: PathBuf,
|
||||
},
|
||||
TokenFileMalformed {
|
||||
path: PathBuf,
|
||||
message: String,
|
||||
},
|
||||
TokenEntryMissing {
|
||||
origin: BackendOrigin,
|
||||
path: PathBuf,
|
||||
},
|
||||
TokenExpired {
|
||||
origin: BackendOrigin,
|
||||
expired_at: String,
|
||||
},
|
||||
Http(reqwest::Error),
|
||||
Unauthorized {
|
||||
origin: BackendOrigin,
|
||||
},
|
||||
Forbidden {
|
||||
origin: BackendOrigin,
|
||||
},
|
||||
BackendStatus {
|
||||
origin: BackendOrigin,
|
||||
status: u16,
|
||||
},
|
||||
BackendResponse {
|
||||
origin: BackendOrigin,
|
||||
status: u16,
|
||||
detail: Option<String>,
|
||||
},
|
||||
Io {
|
||||
path: PathBuf,
|
||||
source: std::io::Error,
|
||||
},
|
||||
}
|
||||
|
||||
impl fmt::Display for BackendApiClientError {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
Self::InvalidBackendOrigin(message) | Self::InvalidRequestPath(message) => {
|
||||
f.write_str(message)
|
||||
}
|
||||
Self::ConfigDirectoryUnavailable => f.write_str(
|
||||
"cannot locate the client configuration directory for backend-tokens.json",
|
||||
),
|
||||
Self::TokenFileMissing { path } => write!(
|
||||
f,
|
||||
"Backend token file {} is missing; run `yoi login --backend <BACKEND>` first",
|
||||
path.display()
|
||||
),
|
||||
Self::TokenFileMalformed { path, message } => write!(
|
||||
f,
|
||||
"Backend token file {} is malformed: {message}; run `yoi login --backend <BACKEND>` again",
|
||||
path.display()
|
||||
),
|
||||
Self::TokenEntryMissing { origin, path } => write!(
|
||||
f,
|
||||
"no Backend token for {origin} exists in {}; login URLs are matched by normalized origin, so run `yoi login --backend {origin}`",
|
||||
path.display()
|
||||
),
|
||||
Self::TokenExpired { origin, expired_at } => write!(
|
||||
f,
|
||||
"Backend token for {origin} expired at {expired_at}; run `yoi login --backend {origin}` again"
|
||||
),
|
||||
Self::Http(error) => write!(f, "Backend request failed: {error}"),
|
||||
Self::Unauthorized { origin } => write!(
|
||||
f,
|
||||
"Backend {origin} returned HTTP 401 for the saved token; it may be expired or revoked, so run `yoi login --backend {origin}` again"
|
||||
),
|
||||
Self::Forbidden { origin } => write!(
|
||||
f,
|
||||
"Backend {origin} returned HTTP 403; the saved token is authenticated but is not authorized for this operation"
|
||||
),
|
||||
Self::BackendStatus { origin, status } => {
|
||||
write!(f, "Backend {origin} returned HTTP {status}")
|
||||
}
|
||||
Self::BackendResponse {
|
||||
origin,
|
||||
status,
|
||||
detail,
|
||||
} => {
|
||||
write!(f, "Backend {origin} returned HTTP {status}")?;
|
||||
if let Some(detail) = detail {
|
||||
write!(f, ": {detail}")?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
Self::Io { path, source } => {
|
||||
write!(f, "failed to access {}: {source}", path.display())
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for BackendApiClientError {
|
||||
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
|
||||
match self {
|
||||
Self::Http(error) => Some(error),
|
||||
Self::Io { source, .. } => Some(source),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Deserialize, Serialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct BackendTokenFile {
|
||||
tokens: BTreeMap<String, BackendTokenEntry>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Deserialize, Serialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct BackendTokenEntry {
|
||||
token_type: String,
|
||||
access_token: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
expires_at: Option<String>,
|
||||
}
|
||||
|
||||
pub fn save_backend_token(
|
||||
base_url: &str,
|
||||
token_type: &str,
|
||||
access_token: &str,
|
||||
) -> Result<PathBuf, BackendApiClientError> {
|
||||
save_backend_token_with_expiry(base_url, token_type, access_token, None)
|
||||
}
|
||||
|
||||
fn save_backend_token_with_expiry(
|
||||
base_url: &str,
|
||||
token_type: &str,
|
||||
access_token: &str,
|
||||
expires_at: Option<String>,
|
||||
) -> Result<PathBuf, BackendApiClientError> {
|
||||
let path = backend_token_file_path()?;
|
||||
save_backend_token_to_file(base_url, token_type, access_token, expires_at, &path)?;
|
||||
Ok(path)
|
||||
}
|
||||
|
||||
fn save_backend_token_to_file(
|
||||
base_url: &str,
|
||||
token_type: &str,
|
||||
access_token: &str,
|
||||
expires_at: Option<String>,
|
||||
path: &Path,
|
||||
) -> Result<(), BackendApiClientError> {
|
||||
let origin = BackendOrigin::parse(base_url)?;
|
||||
let mut token_file = if path.exists() {
|
||||
read_token_file(&path)?
|
||||
} else {
|
||||
BackendTokenFile {
|
||||
tokens: BTreeMap::new(),
|
||||
}
|
||||
};
|
||||
let entry = BackendTokenEntry {
|
||||
token_type: token_type.to_string(),
|
||||
access_token: access_token.to_string(),
|
||||
expires_at,
|
||||
};
|
||||
validate_token_entry(&entry, &origin, path)?;
|
||||
token_file.tokens.insert(origin.to_string(), entry);
|
||||
write_token_file(path, &token_file)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn backend_token_file_path() -> Result<PathBuf, BackendApiClientError> {
|
||||
if let Some(config_home) = env::var_os("XDG_CONFIG_HOME") {
|
||||
return Ok(PathBuf::from(config_home).join("yoi").join(TOKEN_FILE_NAME));
|
||||
}
|
||||
let Some(home) = env::var_os("HOME") else {
|
||||
return Err(BackendApiClientError::ConfigDirectoryUnavailable);
|
||||
};
|
||||
Ok(PathBuf::from(home)
|
||||
.join(".config")
|
||||
.join("yoi")
|
||||
.join(TOKEN_FILE_NAME))
|
||||
}
|
||||
|
||||
fn read_token_file(path: &Path) -> Result<BackendTokenFile, BackendApiClientError> {
|
||||
let bytes = fs::read(path).map_err(|source| {
|
||||
if source.kind() == std::io::ErrorKind::NotFound {
|
||||
BackendApiClientError::TokenFileMissing {
|
||||
path: path.to_path_buf(),
|
||||
}
|
||||
} else {
|
||||
BackendApiClientError::Io {
|
||||
path: path.to_path_buf(),
|
||||
source,
|
||||
}
|
||||
}
|
||||
})?;
|
||||
let raw: BackendTokenFile = serde_json::from_slice(&bytes).map_err(|error| {
|
||||
BackendApiClientError::TokenFileMalformed {
|
||||
path: path.to_path_buf(),
|
||||
message: error.to_string(),
|
||||
}
|
||||
})?;
|
||||
normalize_token_file(raw, path)
|
||||
}
|
||||
|
||||
fn normalize_token_file(
|
||||
token_file: BackendTokenFile,
|
||||
path: &Path,
|
||||
) -> Result<BackendTokenFile, BackendApiClientError> {
|
||||
let mut normalized = BTreeMap::new();
|
||||
for (raw_origin, entry) in token_file.tokens {
|
||||
let origin = BackendOrigin::parse(&raw_origin).map_err(|error| {
|
||||
BackendApiClientError::TokenFileMalformed {
|
||||
path: path.to_path_buf(),
|
||||
message: format!("token key `{raw_origin}` is invalid: {error}"),
|
||||
}
|
||||
})?;
|
||||
if normalized.insert(origin.to_string(), entry).is_some() {
|
||||
return Err(BackendApiClientError::TokenFileMalformed {
|
||||
path: path.to_path_buf(),
|
||||
message: format!("more than one token entry normalizes to `{origin}`"),
|
||||
});
|
||||
}
|
||||
}
|
||||
Ok(BackendTokenFile { tokens: normalized })
|
||||
}
|
||||
|
||||
fn validate_token_entry(
|
||||
entry: &BackendTokenEntry,
|
||||
origin: &BackendOrigin,
|
||||
path: &Path,
|
||||
) -> Result<(), BackendApiClientError> {
|
||||
if !entry.token_type.eq_ignore_ascii_case("Bearer") {
|
||||
return Err(BackendApiClientError::TokenFileMalformed {
|
||||
path: path.to_path_buf(),
|
||||
message: format!("token for `{origin}` does not use the Bearer token type"),
|
||||
});
|
||||
}
|
||||
if entry.access_token.trim().is_empty()
|
||||
|| entry.access_token.contains('\r')
|
||||
|| entry.access_token.contains('\n')
|
||||
{
|
||||
return Err(BackendApiClientError::TokenFileMalformed {
|
||||
path: path.to_path_buf(),
|
||||
message: format!("token for `{origin}` is empty or contains an invalid line break"),
|
||||
});
|
||||
}
|
||||
if reqwest::header::HeaderValue::from_str(&format!("Bearer {}", entry.access_token)).is_err() {
|
||||
return Err(BackendApiClientError::TokenFileMalformed {
|
||||
path: path.to_path_buf(),
|
||||
message: format!("token for `{origin}` cannot be represented as an HTTP header"),
|
||||
});
|
||||
}
|
||||
if let Some(expires_at) = entry.expires_at.as_deref() {
|
||||
let expiration = DateTime::parse_from_rfc3339(expires_at).map_err(|error| {
|
||||
BackendApiClientError::TokenFileMalformed {
|
||||
path: path.to_path_buf(),
|
||||
message: format!("token for `{origin}` has invalid expires_at: {error}"),
|
||||
}
|
||||
})?;
|
||||
if expiration <= Utc::now() {
|
||||
return Err(BackendApiClientError::TokenExpired {
|
||||
origin: origin.clone(),
|
||||
expired_at: expires_at.to_string(),
|
||||
});
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn write_token_file(
|
||||
path: &Path,
|
||||
token_file: &BackendTokenFile,
|
||||
) -> Result<(), BackendApiClientError> {
|
||||
let parent = path
|
||||
.parent()
|
||||
.ok_or(BackendApiClientError::ConfigDirectoryUnavailable)?;
|
||||
fs::create_dir_all(parent).map_err(|source| BackendApiClientError::Io {
|
||||
path: parent.to_path_buf(),
|
||||
source,
|
||||
})?;
|
||||
let payload = serde_json::to_vec_pretty(token_file).map_err(|error| {
|
||||
BackendApiClientError::TokenFileMalformed {
|
||||
path: path.to_path_buf(),
|
||||
message: error.to_string(),
|
||||
}
|
||||
})?;
|
||||
let temp_path = parent.join(format!(".{TOKEN_FILE_NAME}.tmp-{}", std::process::id()));
|
||||
let mut options = OpenOptions::new();
|
||||
options.write(true).create(true).truncate(true);
|
||||
#[cfg(unix)]
|
||||
{
|
||||
use std::os::unix::fs::OpenOptionsExt;
|
||||
options.mode(0o600);
|
||||
}
|
||||
let mut file = options
|
||||
.open(&temp_path)
|
||||
.map_err(|source| BackendApiClientError::Io {
|
||||
path: temp_path.clone(),
|
||||
source,
|
||||
})?;
|
||||
file.write_all(&payload)
|
||||
.and_then(|()| file.write_all(b"\n"))
|
||||
.and_then(|()| file.sync_all())
|
||||
.map_err(|source| BackendApiClientError::Io {
|
||||
path: temp_path.clone(),
|
||||
source,
|
||||
})?;
|
||||
fs::rename(&temp_path, path).map_err(|source| BackendApiClientError::Io {
|
||||
path: path.to_path_buf(),
|
||||
source,
|
||||
})?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::io::{Read, Write};
|
||||
use std::net::TcpListener;
|
||||
use std::thread;
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
|
||||
fn temp_path(label: &str) -> PathBuf {
|
||||
let nonce = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_nanos();
|
||||
env::temp_dir().join(format!(
|
||||
"yoi-client-{label}-{}-{nonce}.json",
|
||||
std::process::id()
|
||||
))
|
||||
}
|
||||
|
||||
fn write_fixture(path: &Path, value: serde_json::Value) {
|
||||
fs::write(path, serde_json::to_vec(&value).unwrap()).unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backend_origin_normalizes_safe_equivalents() {
|
||||
let variants = [
|
||||
"HTTP://Example.COM",
|
||||
"http://example.com/",
|
||||
"http://example.com:80////",
|
||||
];
|
||||
for variant in variants {
|
||||
assert_eq!(
|
||||
BackendOrigin::parse(variant).unwrap().as_str(),
|
||||
"http://example.com"
|
||||
);
|
||||
}
|
||||
assert_eq!(
|
||||
BackendOrigin::parse("https://EXAMPLE.com:443/")
|
||||
.unwrap()
|
||||
.as_str(),
|
||||
"https://example.com"
|
||||
);
|
||||
assert_eq!(
|
||||
BackendOrigin::parse("https://example.com:8443/")
|
||||
.unwrap()
|
||||
.as_str(),
|
||||
"https://example.com:8443"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backend_error_detail_preserves_public_server_message() {
|
||||
let detail = backend_error_detail(
|
||||
br#"{"error":"Bad Request","message":"working_directory_runtime_mismatch: Working directory is owned by a different Runtime","diagnostics":[{"code":"working_directory_runtime_mismatch"}]}"#,
|
||||
);
|
||||
let error = BackendApiClientError::BackendResponse {
|
||||
origin: BackendOrigin::parse("http://127.0.0.1:8787").unwrap(),
|
||||
status: 400,
|
||||
detail,
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
error.to_string(),
|
||||
"Backend http://127.0.0.1:8787 returned HTTP 400: working_directory_runtime_mismatch: Working directory is owned by a different Runtime"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backend_origin_rejects_unsafe_authority_changes() {
|
||||
for invalid in [
|
||||
"ftp://example.com",
|
||||
"https://user@example.com",
|
||||
"https://example.com/api",
|
||||
"https://example.com/?query=1",
|
||||
"https://example.com/#fragment",
|
||||
] {
|
||||
assert!(BackendOrigin::parse(invalid).is_err(), "accepted {invalid}");
|
||||
}
|
||||
assert_ne!(
|
||||
BackendOrigin::parse("http://localhost:8787").unwrap(),
|
||||
BackendOrigin::parse("http://127.0.0.1:8787").unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn token_lookup_distinguishes_missing_malformed_mismatch_and_expired() {
|
||||
let missing = temp_path("missing");
|
||||
assert!(matches!(
|
||||
BackendApiClient::from_token_file("http://localhost:8787", &missing),
|
||||
Err(BackendApiClientError::TokenFileMissing { .. })
|
||||
));
|
||||
|
||||
let malformed = temp_path("malformed");
|
||||
fs::write(&malformed, b"not json").unwrap();
|
||||
assert!(matches!(
|
||||
BackendApiClient::from_token_file("http://localhost:8787", &malformed),
|
||||
Err(BackendApiClientError::TokenFileMalformed { .. })
|
||||
));
|
||||
|
||||
let mismatch = temp_path("mismatch");
|
||||
write_fixture(
|
||||
&mismatch,
|
||||
serde_json::json!({"tokens": {"http://localhost:8787": {
|
||||
"token_type": "Bearer", "access_token": "secret"
|
||||
}}}),
|
||||
);
|
||||
assert!(matches!(
|
||||
BackendApiClient::from_token_file("http://127.0.0.1:8787", &mismatch),
|
||||
Err(BackendApiClientError::TokenEntryMissing { .. })
|
||||
));
|
||||
|
||||
let expired = temp_path("expired");
|
||||
write_fixture(
|
||||
&expired,
|
||||
serde_json::json!({"tokens": {"http://localhost:8787": {
|
||||
"token_type": "Bearer",
|
||||
"access_token": "secret",
|
||||
"expires_at": "2000-01-01T00:00:00Z"
|
||||
}}}),
|
||||
);
|
||||
assert!(matches!(
|
||||
BackendApiClient::from_token_file("http://localhost:8787", &expired),
|
||||
Err(BackendApiClientError::TokenExpired { .. })
|
||||
));
|
||||
|
||||
for path in [malformed, mismatch, expired] {
|
||||
let _ = fs::remove_file(path);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn token_write_and_lookup_share_origin_normalization() {
|
||||
let path = temp_path("normalized-write");
|
||||
save_backend_token_to_file(
|
||||
"HTTP://Example.COM:80////",
|
||||
"Bearer",
|
||||
"normalized-secret",
|
||||
None,
|
||||
&path,
|
||||
)
|
||||
.unwrap();
|
||||
let contents = fs::read_to_string(&path).unwrap();
|
||||
assert!(contents.contains("\"http://example.com\""));
|
||||
let client = BackendApiClient::from_token_file("http://example.com/", &path).unwrap();
|
||||
assert_eq!(
|
||||
client.authorization_header_value(),
|
||||
"Bearer normalized-secret"
|
||||
);
|
||||
fs::remove_file(path).unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn client_debug_and_errors_never_include_token_value() {
|
||||
let client = BackendApiClient::from_access_token_for_test(
|
||||
"http://localhost:8787",
|
||||
"never-print-this-token",
|
||||
)
|
||||
.unwrap();
|
||||
assert!(!format!("{client:?}").contains("never-print-this-token"));
|
||||
assert!(
|
||||
!BackendApiClientError::Unauthorized {
|
||||
origin: client.origin().clone()
|
||||
}
|
||||
.to_string()
|
||||
.contains("never-print-this-token")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn authenticated_requests_follow_only_same_origin_redirects() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
|
||||
let origin = format!("http://{}", listener.local_addr().unwrap());
|
||||
let handle = thread::spawn(move || {
|
||||
for response in [
|
||||
"HTTP/1.1 302 Found\r\nLocation: /final\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
|
||||
"HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
|
||||
] {
|
||||
let (mut stream, _) = listener.accept().unwrap();
|
||||
let mut request = vec![0; 4096];
|
||||
let read = stream.read(&mut request).unwrap();
|
||||
let request = String::from_utf8_lossy(&request[..read]).to_ascii_lowercase();
|
||||
assert!(request.contains("authorization: bearer redirect-secret\r\n"));
|
||||
stream.write_all(response.as_bytes()).unwrap();
|
||||
}
|
||||
});
|
||||
let client =
|
||||
BackendApiClient::from_access_token_for_test(&origin, "redirect-secret").unwrap();
|
||||
let response = client
|
||||
.blocking_request(Method::GET, "/start")
|
||||
.unwrap()
|
||||
.send()
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
handle.join().unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn authenticated_requests_reject_cross_origin_redirects_without_leaking_token() {
|
||||
let source = TcpListener::bind("127.0.0.1:0").unwrap();
|
||||
let target = TcpListener::bind("127.0.0.1:0").unwrap();
|
||||
target.set_nonblocking(true).unwrap();
|
||||
let source_origin = format!("http://{}", source.local_addr().unwrap());
|
||||
let target_origin = format!("http://{}", target.local_addr().unwrap());
|
||||
let location = format!("{target_origin}/capture");
|
||||
let handle = thread::spawn(move || {
|
||||
let (mut stream, _) = source.accept().unwrap();
|
||||
let mut request = vec![0; 4096];
|
||||
let read = stream.read(&mut request).unwrap();
|
||||
let request = String::from_utf8_lossy(&request[..read]).to_ascii_lowercase();
|
||||
assert!(request.contains("authorization: bearer redirect-secret\r\n"));
|
||||
let response = format!(
|
||||
"HTTP/1.1 302 Found\r\nLocation: {location}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
|
||||
);
|
||||
stream.write_all(response.as_bytes()).unwrap();
|
||||
});
|
||||
let client =
|
||||
BackendApiClient::from_access_token_for_test(&source_origin, "redirect-secret")
|
||||
.unwrap();
|
||||
let error = client
|
||||
.blocking_request(Method::GET, "/start")
|
||||
.unwrap()
|
||||
.send()
|
||||
.unwrap_err();
|
||||
let message = error.to_string();
|
||||
assert!(message.contains("redirect"));
|
||||
assert!(!message.contains("redirect-secret"));
|
||||
handle.join().unwrap();
|
||||
thread::sleep(Duration::from_millis(20));
|
||||
assert!(matches!(
|
||||
target.accept(),
|
||||
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn status_diagnostics_distinguish_unauthorized_and_forbidden() {
|
||||
let client =
|
||||
BackendApiClient::from_access_token_for_test("http://localhost:8787", "secret")
|
||||
.unwrap();
|
||||
assert!(matches!(
|
||||
client.check_status(StatusCode::UNAUTHORIZED),
|
||||
Err(BackendApiClientError::Unauthorized { .. })
|
||||
));
|
||||
assert!(matches!(
|
||||
client.check_status(StatusCode::FORBIDDEN),
|
||||
Err(BackendApiClientError::Forbidden { .. })
|
||||
));
|
||||
}
|
||||
}
|
||||
@@ -1,7 +1,11 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
use crate::BackendOrigin;
|
||||
use serde::Deserialize;
|
||||
use std::fmt;
|
||||
use std::time::Duration;
|
||||
|
||||
use workspace_api::{DeviceLoginPollRequest, DeviceLoginPollStatus, DeviceLoginStartRequest};
|
||||
pub use workspace_api::{DeviceLoginPollResponse, DeviceLoginStartResponse};
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct BackendAuthTarget {
|
||||
pub base_url: String,
|
||||
@@ -9,9 +13,11 @@ pub struct BackendAuthTarget {
|
||||
|
||||
impl BackendAuthTarget {
|
||||
pub fn new(base_url: impl Into<String>) -> Self {
|
||||
Self {
|
||||
base_url: base_url.into(),
|
||||
}
|
||||
let base_url = base_url.into();
|
||||
let base_url = BackendOrigin::parse(&base_url)
|
||||
.map(|origin| origin.to_string())
|
||||
.unwrap_or(base_url);
|
||||
Self { base_url }
|
||||
}
|
||||
|
||||
fn api_url(&self, path: &str) -> String {
|
||||
@@ -25,23 +31,6 @@ impl BackendAuthTarget {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Deserialize, PartialEq, Eq)]
|
||||
pub struct DeviceLoginStartResponse {
|
||||
pub device_code: String,
|
||||
pub user_code: String,
|
||||
pub verification_uri: String,
|
||||
pub verification_uri_complete: String,
|
||||
pub expires_in: u64,
|
||||
pub interval: u64,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Deserialize, PartialEq, Eq)]
|
||||
pub struct DeviceLoginPollResponse {
|
||||
pub status: String,
|
||||
pub access_token: Option<String>,
|
||||
pub token_type: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum BackendAuthClientError {
|
||||
Http(reqwest::Error),
|
||||
@@ -71,16 +60,6 @@ impl From<reqwest::Error> for BackendAuthClientError {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct DeviceLoginStartRequest<'a> {
|
||||
client_name: Option<&'a str>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct DeviceLoginPollRequest<'a> {
|
||||
device_code: &'a str,
|
||||
}
|
||||
|
||||
pub async fn start_device_login(
|
||||
target: &BackendAuthTarget,
|
||||
client_name: Option<&str>,
|
||||
@@ -88,7 +67,9 @@ pub async fn start_device_login(
|
||||
let client = reqwest::Client::new();
|
||||
let response = client
|
||||
.post(target.api_url("/api/auth/device-login/start"))
|
||||
.json(&DeviceLoginStartRequest { client_name })
|
||||
.json(&DeviceLoginStartRequest {
|
||||
client_name: client_name.map(ToOwned::to_owned),
|
||||
})
|
||||
.send()
|
||||
.await?;
|
||||
parse_json_response(response).await
|
||||
@@ -101,12 +82,38 @@ pub async fn poll_device_login(
|
||||
let client = reqwest::Client::new();
|
||||
let response = client
|
||||
.post(target.api_url("/api/auth/device-login/poll"))
|
||||
.json(&DeviceLoginPollRequest { device_code })
|
||||
.json(&DeviceLoginPollRequest {
|
||||
device_code: device_code.to_string(),
|
||||
})
|
||||
.send()
|
||||
.await?;
|
||||
parse_json_response(response).await
|
||||
}
|
||||
|
||||
fn device_login_poll_result(
|
||||
response: DeviceLoginPollResponse,
|
||||
) -> Result<Option<String>, BackendAuthClientError> {
|
||||
match response.status {
|
||||
DeviceLoginPollStatus::Approved => response
|
||||
.access_token
|
||||
.ok_or(BackendAuthClientError::MissingAccessToken)
|
||||
.map(Some),
|
||||
DeviceLoginPollStatus::Expired => Err(BackendAuthClientError::BackendStatus {
|
||||
status: 410,
|
||||
body: "device login expired".to_string(),
|
||||
}),
|
||||
DeviceLoginPollStatus::Denied => Err(BackendAuthClientError::BackendStatus {
|
||||
status: 403,
|
||||
body: "device login was denied".to_string(),
|
||||
}),
|
||||
DeviceLoginPollStatus::Consumed => Err(BackendAuthClientError::BackendStatus {
|
||||
status: 409,
|
||||
body: "device login was already consumed".to_string(),
|
||||
}),
|
||||
DeviceLoginPollStatus::Pending => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn wait_for_device_login(
|
||||
target: &BackendAuthTarget,
|
||||
device_code: &str,
|
||||
@@ -116,25 +123,8 @@ pub async fn wait_for_device_login(
|
||||
let started = std::time::Instant::now();
|
||||
loop {
|
||||
let response = poll_device_login(target, device_code).await?;
|
||||
match response.status.as_str() {
|
||||
"approved" => {
|
||||
return response
|
||||
.access_token
|
||||
.ok_or(BackendAuthClientError::MissingAccessToken);
|
||||
}
|
||||
"expired" => {
|
||||
return Err(BackendAuthClientError::BackendStatus {
|
||||
status: 410,
|
||||
body: "device login expired".to_string(),
|
||||
});
|
||||
}
|
||||
"consumed" => {
|
||||
return Err(BackendAuthClientError::BackendStatus {
|
||||
status: 409,
|
||||
body: "device login was already consumed".to_string(),
|
||||
});
|
||||
}
|
||||
_ => {}
|
||||
if let Some(access_token) = device_login_poll_result(response)? {
|
||||
return Ok(access_token);
|
||||
}
|
||||
if started.elapsed() >= expires_in {
|
||||
return Err(BackendAuthClientError::BackendStatus {
|
||||
@@ -159,3 +149,81 @@ async fn parse_json_response<T: for<'de> Deserialize<'de>>(
|
||||
}
|
||||
Ok(response.json::<T>().await?)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use workspace_api::DeviceAccessTokenType;
|
||||
|
||||
fn poll_response(status: DeviceLoginPollStatus) -> DeviceLoginPollResponse {
|
||||
DeviceLoginPollResponse {
|
||||
status,
|
||||
access_token: None,
|
||||
token_type: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn device_login_start_response_enforces_shared_expiry_bounds() {
|
||||
let valid = serde_json::json!({
|
||||
"device_code": "device-secret",
|
||||
"user_code": "ABCD-EFGH",
|
||||
"verification_uri": "https://yoi.example/login/device",
|
||||
"verification_uri_complete": "https://yoi.example/login/device?user_code=ABCD-EFGH",
|
||||
"expires_in": 600,
|
||||
"interval": 5
|
||||
});
|
||||
assert!(serde_json::from_value::<DeviceLoginStartResponse>(valid.clone()).is_ok());
|
||||
|
||||
let mut expired = valid;
|
||||
expired["expires_in"] = serde_json::json!(0);
|
||||
assert!(serde_json::from_value::<DeviceLoginStartResponse>(expired).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn device_login_poll_response_rejects_unknown_status() {
|
||||
assert!(
|
||||
serde_json::from_value::<DeviceLoginPollResponse>(
|
||||
serde_json::json!({"status": "future_status"}),
|
||||
)
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn device_login_poll_result_handles_pending_and_terminal_states() {
|
||||
assert!(
|
||||
device_login_poll_result(poll_response(DeviceLoginPollStatus::Pending))
|
||||
.unwrap()
|
||||
.is_none()
|
||||
);
|
||||
|
||||
let approved = DeviceLoginPollResponse {
|
||||
status: DeviceLoginPollStatus::Approved,
|
||||
access_token: Some("access-secret".to_string()),
|
||||
token_type: Some(DeviceAccessTokenType::Bearer),
|
||||
};
|
||||
assert_eq!(
|
||||
device_login_poll_result(approved).unwrap(),
|
||||
Some("access-secret".to_string())
|
||||
);
|
||||
assert!(matches!(
|
||||
device_login_poll_result(poll_response(DeviceLoginPollStatus::Approved)),
|
||||
Err(BackendAuthClientError::MissingAccessToken)
|
||||
));
|
||||
|
||||
for (status, expected_http_status) in [
|
||||
(DeviceLoginPollStatus::Expired, 410),
|
||||
(DeviceLoginPollStatus::Denied, 403),
|
||||
(DeviceLoginPollStatus::Consumed, 409),
|
||||
] {
|
||||
assert!(matches!(
|
||||
device_login_poll_result(poll_response(status)),
|
||||
Err(BackendAuthClientError::BackendStatus {
|
||||
status,
|
||||
..
|
||||
}) if status == expected_http_status
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,12 +1,11 @@
|
||||
use futures::{SinkExt, StreamExt};
|
||||
use protocol::stream::{decode_event, encode_method};
|
||||
use protocol::{ErrorCode, Event, Method};
|
||||
use std::collections::VecDeque;
|
||||
use crate::transport::websocket::{Socket as WebSocket, SocketError as WebSocketError};
|
||||
use crate::{BackendApiClient, BackendApiClientError, Client};
|
||||
use reqwest::Method as HttpMethod;
|
||||
use serde::Deserialize;
|
||||
use std::fmt;
|
||||
use tokio::sync::mpsc;
|
||||
use tokio_tungstenite::connect_async;
|
||||
use tokio_tungstenite::tungstenite::Message as TungsteniteMessage;
|
||||
pub use workdir::workspace::WorkingDirectorySummary as BackendWorkingDirectorySummary;
|
||||
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
|
||||
use tokio_tungstenite::tungstenite::http::HeaderValue;
|
||||
use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION;
|
||||
pub use workspace_api::{
|
||||
Diagnostic as BackendDiagnostic, DiagnosticSeverity as BackendDiagnosticSeverity,
|
||||
ListResponse as BackendRuntimeListResponse, RuntimeSummary as BackendRuntimeSummary,
|
||||
@@ -15,6 +14,11 @@ pub use workspace_api::{
|
||||
WorkerRestoreResponse as BackendWorkerRestoreResponse,
|
||||
WorkerRestoreResult as BackendWorkerRestoreResult, WorkerSummary as BackendWorkerSummary,
|
||||
WorkerWorkspaceSummary as BackendWorkerWorkspaceSummary,
|
||||
WorkingDirectoryCreateRequest as BackendWorkingDirectoryCreateRequest,
|
||||
WorkingDirectoryCreateResponse as BackendWorkingDirectoryCreateResponse,
|
||||
WorkingDirectoryDetailResponse as BackendWorkingDirectoryDetailResponse,
|
||||
WorkingDirectoryListResponse as BackendWorkingDirectoryListResponse,
|
||||
WorkingDirectorySummary as BackendWorkingDirectorySummary,
|
||||
};
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
@@ -48,6 +52,123 @@ impl BackendRuntimeTarget {
|
||||
pub fn display_label(&self) -> String {
|
||||
format!("{}:{}", self.runtime_id, self.worker_id)
|
||||
}
|
||||
|
||||
pub async fn upload_file(
|
||||
&self,
|
||||
file_name: &str,
|
||||
media_type: &str,
|
||||
content: Vec<u8>,
|
||||
) -> Result<protocol::UploadedFileRef, BackendRuntimeClientError> {
|
||||
self.upload_file_with_id(
|
||||
&uuid::Uuid::now_v7().to_string(),
|
||||
file_name,
|
||||
media_type,
|
||||
content,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn upload_file_with_id(
|
||||
&self,
|
||||
upload_id: &str,
|
||||
file_name: &str,
|
||||
media_type: &str,
|
||||
content: Vec<u8>,
|
||||
) -> Result<protocol::UploadedFileRef, BackendRuntimeClientError> {
|
||||
let api = BackendApiClient::from_stored_token(&self.base_url)?;
|
||||
let worker_path = format!(
|
||||
"/api/w/{}/runtimes/{}/workers/{}",
|
||||
path_segment_encode(&self.workspace_id),
|
||||
path_segment_encode(&self.runtime_id),
|
||||
path_segment_encode(&self.worker_id),
|
||||
);
|
||||
let grant_path = format!(
|
||||
"{worker_path}/attachment-upload-grants?file_name={}&media_type={}&upload_id={}",
|
||||
path_segment_encode(file_name),
|
||||
path_segment_encode(media_type),
|
||||
path_segment_encode(&upload_id),
|
||||
);
|
||||
let grant_response = api
|
||||
.request(HttpMethod::POST, &grant_path)?
|
||||
.send()
|
||||
.await
|
||||
.map_err(BackendRuntimeClientError::Http)?;
|
||||
api.check_status(grant_response.status())?;
|
||||
let grant = grant_response
|
||||
.json::<AttachmentUploadGrantResponse>()
|
||||
.await
|
||||
.map_err(BackendRuntimeClientError::Http)?;
|
||||
let upload_path = format!(
|
||||
"{worker_path}/attachment-uploads/{}",
|
||||
path_segment_encode(&grant.upload_id),
|
||||
);
|
||||
let response = api
|
||||
.request(HttpMethod::PUT, &upload_path)?
|
||||
.body(content)
|
||||
.send()
|
||||
.await
|
||||
.map_err(BackendRuntimeClientError::Http)?;
|
||||
api.check_status(response.status())?;
|
||||
response
|
||||
.json::<UploadedFileResponse>()
|
||||
.await
|
||||
.map(|response| response.file)
|
||||
.map_err(BackendRuntimeClientError::Http)
|
||||
}
|
||||
|
||||
pub async fn cancel_file_upload(
|
||||
&self,
|
||||
upload_id: &str,
|
||||
) -> Result<(), BackendRuntimeClientError> {
|
||||
let api = BackendApiClient::from_stored_token(&self.base_url)?;
|
||||
let path = format!(
|
||||
"/api/w/{}/runtimes/{}/workers/{}/attachment-uploads/{}",
|
||||
path_segment_encode(&self.workspace_id),
|
||||
path_segment_encode(&self.runtime_id),
|
||||
path_segment_encode(&self.worker_id),
|
||||
path_segment_encode(upload_id),
|
||||
);
|
||||
let response = api
|
||||
.request(HttpMethod::DELETE, &path)?
|
||||
.send()
|
||||
.await
|
||||
.map_err(BackendRuntimeClientError::Http)?;
|
||||
api.check_status(response.status())?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn delete_uploaded_file(
|
||||
&self,
|
||||
artifact_id: &str,
|
||||
) -> Result<(), BackendRuntimeClientError> {
|
||||
let api = BackendApiClient::from_stored_token(&self.base_url)?;
|
||||
let path = format!(
|
||||
"/api/w/{}/runtimes/{}/workers/{}/attachments/{}",
|
||||
path_segment_encode(&self.workspace_id),
|
||||
path_segment_encode(&self.runtime_id),
|
||||
path_segment_encode(&self.worker_id),
|
||||
path_segment_encode(artifact_id),
|
||||
);
|
||||
let response = api
|
||||
.request(HttpMethod::DELETE, &path)?
|
||||
.send()
|
||||
.await
|
||||
.map_err(BackendRuntimeClientError::Http)?;
|
||||
api.check_status(response.status())?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct AttachmentUploadGrantResponse {
|
||||
upload_id: String,
|
||||
#[allow(dead_code)]
|
||||
expires_at_ms: u64,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct UploadedFileResponse {
|
||||
file: protocol::UploadedFileRef,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
@@ -101,32 +222,33 @@ impl BackendRuntimeListTarget {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct BackendRuntimeClient {
|
||||
target: BackendRuntimeTarget,
|
||||
command_tx: mpsc::UnboundedSender<Method>,
|
||||
events: mpsc::UnboundedReceiver<Event>,
|
||||
diagnostics: VecDeque<Event>,
|
||||
_protocol_task: tokio::task::JoinHandle<()>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum BackendRuntimeClientError {
|
||||
InvalidTarget(String),
|
||||
Api(BackendApiClientError),
|
||||
Http(reqwest::Error),
|
||||
Protocol(String),
|
||||
}
|
||||
|
||||
impl fmt::Display for BackendRuntimeClientError {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
Self::InvalidTarget(message) => f.write_str(message),
|
||||
Self::Api(error) => write!(f, "{error}"),
|
||||
Self::Http(error) => write!(f, "{error}"),
|
||||
Self::Protocol(message) => f.write_str(message),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for BackendRuntimeClientError {}
|
||||
|
||||
impl From<BackendApiClientError> for BackendRuntimeClientError {
|
||||
fn from(error: BackendApiClientError) -> Self {
|
||||
Self::Api(error)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<reqwest::Error> for BackendRuntimeClientError {
|
||||
fn from(error: reqwest::Error) -> Self {
|
||||
Self::Http(error)
|
||||
@@ -137,7 +259,7 @@ pub async fn list_backend_workers(
|
||||
target: &BackendRuntimeListTarget,
|
||||
) -> Result<BackendRuntimeListResponse<BackendWorkerSummary>, BackendRuntimeClientError> {
|
||||
validate_list_target(target)?;
|
||||
let http = reqwest::Client::new();
|
||||
let api = BackendApiClient::from_stored_token(&target.base_url)?;
|
||||
if let Some(runtime_id) = target.runtime_id.as_deref() {
|
||||
let path = backend_runtime_workers_path(
|
||||
target
|
||||
@@ -146,12 +268,9 @@ pub async fn list_backend_workers(
|
||||
.expect("validated Backend Workspace scope"),
|
||||
runtime_id,
|
||||
);
|
||||
let url = join_base_and_path(&target.base_url, &path);
|
||||
return Ok(http
|
||||
.get(url)
|
||||
.send()
|
||||
.await?
|
||||
.error_for_status()?
|
||||
let response = api.request(HttpMethod::GET, &path)?.send().await?;
|
||||
api.check_status(response.status())?;
|
||||
return Ok(response
|
||||
.json::<BackendRuntimeListResponse<BackendWorkerSummary>>()
|
||||
.await?);
|
||||
}
|
||||
@@ -162,12 +281,9 @@ pub async fn list_backend_workers(
|
||||
.as_deref()
|
||||
.expect("validated Backend Workspace scope"),
|
||||
);
|
||||
let runtime_url = join_base_and_path(&target.base_url, &runtime_path);
|
||||
let runtimes = http
|
||||
.get(runtime_url)
|
||||
.send()
|
||||
.await?
|
||||
.error_for_status()?
|
||||
let response = api.request(HttpMethod::GET, &runtime_path)?.send().await?;
|
||||
api.check_status(response.status())?;
|
||||
let runtimes = response
|
||||
.json::<BackendRuntimeListResponse<BackendRuntimeSummary>>()
|
||||
.await?;
|
||||
|
||||
@@ -181,29 +297,43 @@ pub async fn list_backend_workers(
|
||||
.expect("validated Backend Workspace scope"),
|
||||
&runtime.runtime_id,
|
||||
);
|
||||
let url = join_base_and_path(&target.base_url, &path);
|
||||
match http
|
||||
.get(url)
|
||||
.send()
|
||||
.await
|
||||
.and_then(|response| response.error_for_status())
|
||||
{
|
||||
Ok(response) => {
|
||||
let response = response
|
||||
.json::<BackendRuntimeListResponse<BackendWorkerSummary>>()
|
||||
.await?;
|
||||
diagnostics.extend(response.diagnostics);
|
||||
items.extend(response.items);
|
||||
let response = match api.request(HttpMethod::GET, &path)?.send().await {
|
||||
Ok(response) => response,
|
||||
Err(error) => {
|
||||
diagnostics.push(BackendDiagnostic {
|
||||
code: "runtime_worker_list_failed".to_string(),
|
||||
severity: BackendDiagnosticSeverity::Error,
|
||||
message: format!(
|
||||
"failed to list workers for runtime {}: {error}",
|
||||
runtime.runtime_id
|
||||
),
|
||||
});
|
||||
continue;
|
||||
}
|
||||
Err(error) => diagnostics.push(BackendDiagnostic {
|
||||
};
|
||||
if matches!(
|
||||
response.status(),
|
||||
reqwest::StatusCode::UNAUTHORIZED | reqwest::StatusCode::FORBIDDEN
|
||||
) {
|
||||
api.check_status(response.status())?;
|
||||
}
|
||||
if !response.status().is_success() {
|
||||
diagnostics.push(BackendDiagnostic {
|
||||
code: "runtime_worker_list_failed".to_string(),
|
||||
severity: BackendDiagnosticSeverity::Error,
|
||||
message: format!(
|
||||
"failed to list workers for runtime {}: {error}",
|
||||
runtime.runtime_id
|
||||
"failed to list workers for runtime {}: Backend returned HTTP {}",
|
||||
runtime.runtime_id,
|
||||
response.status().as_u16()
|
||||
),
|
||||
}),
|
||||
});
|
||||
continue;
|
||||
}
|
||||
let response = response
|
||||
.json::<BackendRuntimeListResponse<BackendWorkerSummary>>()
|
||||
.await?;
|
||||
diagnostics.extend(response.diagnostics);
|
||||
items.extend(response.items);
|
||||
}
|
||||
|
||||
Ok(BackendRuntimeListResponse {
|
||||
@@ -224,7 +354,7 @@ pub async fn list_backend_stopped_workers(
|
||||
"stopped worker listing requires a runtime id".to_string(),
|
||||
));
|
||||
};
|
||||
let http = reqwest::Client::new();
|
||||
let api = BackendApiClient::from_stored_token(&target.base_url)?;
|
||||
let path = backend_runtime_workers_path(
|
||||
target
|
||||
.workspace_id
|
||||
@@ -232,12 +362,12 @@ pub async fn list_backend_stopped_workers(
|
||||
.expect("validated Backend Workspace scope"),
|
||||
runtime_id,
|
||||
);
|
||||
let url = join_base_and_path(&target.base_url, &format!("{path}?status=stopped"));
|
||||
Ok(http
|
||||
.get(url)
|
||||
let response = api
|
||||
.request(HttpMethod::GET, &format!("{path}?status=stopped"))?
|
||||
.send()
|
||||
.await?
|
||||
.error_for_status()?
|
||||
.await?;
|
||||
api.check_status(response.status())?;
|
||||
Ok(response
|
||||
.json::<BackendRuntimeListResponse<BackendWorkerSummary>>()
|
||||
.await?)
|
||||
}
|
||||
@@ -246,166 +376,61 @@ pub async fn restore_backend_worker(
|
||||
target: &BackendRuntimeTarget,
|
||||
) -> Result<BackendWorkerRestoreResponse, BackendRuntimeClientError> {
|
||||
validate_target(target)?;
|
||||
let http = reqwest::Client::new();
|
||||
let api = BackendApiClient::from_stored_token(&target.base_url)?;
|
||||
let path = backend_runtime_worker_restore_path(
|
||||
&target.workspace_id,
|
||||
&target.runtime_id,
|
||||
&target.worker_id,
|
||||
);
|
||||
let url = join_base_and_path(&target.base_url, &path);
|
||||
Ok(http
|
||||
.post(url)
|
||||
let response = api
|
||||
.request(HttpMethod::POST, &path)?
|
||||
.json(&serde_json::json!({}))
|
||||
.send()
|
||||
.await?
|
||||
.error_for_status()?
|
||||
.json::<BackendWorkerRestoreResponse>()
|
||||
.await?)
|
||||
.await?;
|
||||
let response = api.require_success(response).await?;
|
||||
Ok(response.json::<BackendWorkerRestoreResponse>().await?)
|
||||
}
|
||||
|
||||
impl BackendRuntimeClient {
|
||||
pub async fn connect(target: BackendRuntimeTarget) -> Result<Self, BackendRuntimeClientError> {
|
||||
validate_target(&target)?;
|
||||
let (event_tx, rx) = mpsc::unbounded_channel();
|
||||
let (command_tx, command_rx) = mpsc::unbounded_channel();
|
||||
|
||||
let protocol_target = target.clone();
|
||||
let protocol_event_tx = event_tx.clone();
|
||||
let protocol_task = tokio::spawn(async move {
|
||||
run_worker_protocol_transport(protocol_target, command_rx, protocol_event_tx).await;
|
||||
});
|
||||
|
||||
Ok(Self {
|
||||
target,
|
||||
command_tx,
|
||||
events: rx,
|
||||
diagnostics: VecDeque::new(),
|
||||
_protocol_task: protocol_task,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn try_next_event(&mut self) -> Option<Event> {
|
||||
if let Some(event) = self.diagnostics.pop_front() {
|
||||
return Some(event);
|
||||
}
|
||||
self.events.try_recv().ok()
|
||||
}
|
||||
|
||||
pub async fn next_event(&mut self) -> Option<Event> {
|
||||
if let Some(event) = self.diagnostics.pop_front() {
|
||||
return Some(event);
|
||||
}
|
||||
self.events.recv().await
|
||||
}
|
||||
|
||||
pub async fn send(&mut self, method: &Method) -> Result<(), BackendRuntimeClientError> {
|
||||
self.command_tx.send(method.clone()).map_err(|_| {
|
||||
BackendRuntimeClientError::InvalidTarget(format!(
|
||||
"Backend protocol command stream is closed for {}",
|
||||
self.target.display_label()
|
||||
))
|
||||
})?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for BackendRuntimeClient {
|
||||
fn drop(&mut self) {
|
||||
self._protocol_task.abort();
|
||||
}
|
||||
}
|
||||
|
||||
async fn run_worker_protocol_transport(
|
||||
pub async fn connect_backend_runtime(
|
||||
target: BackendRuntimeTarget,
|
||||
mut commands: mpsc::UnboundedReceiver<Method>,
|
||||
tx: mpsc::UnboundedSender<Event>,
|
||||
) {
|
||||
let url = protocol_ws_url(&target);
|
||||
match connect_async(&url).await {
|
||||
Ok((ws, _)) => {
|
||||
let (mut sink, mut stream) = ws.split();
|
||||
loop {
|
||||
tokio::select! {
|
||||
maybe_method = commands.recv() => {
|
||||
let Some(method) = maybe_method else {
|
||||
break;
|
||||
};
|
||||
match encode_method(&method) {
|
||||
Ok(text) => {
|
||||
if let Err(error) = sink.send(TungsteniteMessage::Text(text.into())).await {
|
||||
let _ = tx.send(diagnostic_event(format!(
|
||||
"Backend protocol command send failed for {}: {error}",
|
||||
target.display_label()
|
||||
)));
|
||||
break;
|
||||
}
|
||||
}
|
||||
Err(error) => {
|
||||
let _ = tx.send(diagnostic_event(format!(
|
||||
"Backend protocol command could not serialize method for {}: {error}",
|
||||
target.display_label()
|
||||
)));
|
||||
}
|
||||
}
|
||||
}
|
||||
frame = stream.next() => {
|
||||
match frame {
|
||||
Some(Ok(TungsteniteMessage::Text(text))) => {
|
||||
match decode_event(&text) {
|
||||
Ok(event) => {
|
||||
let _ = tx.send(event);
|
||||
}
|
||||
Err(error) => {
|
||||
let _ = tx.send(diagnostic_event(format!(
|
||||
"Backend protocol response was not valid Event JSON for {}: {error}",
|
||||
target.display_label()
|
||||
)));
|
||||
}
|
||||
}
|
||||
}
|
||||
Some(Ok(TungsteniteMessage::Close(_))) | None => {
|
||||
let _ = tx.send(diagnostic_event(format!(
|
||||
"Backend protocol command stream closed for {}",
|
||||
target.display_label()
|
||||
)));
|
||||
break;
|
||||
}
|
||||
Some(Ok(TungsteniteMessage::Ping(_)))
|
||||
| Some(Ok(TungsteniteMessage::Pong(_)))
|
||||
| Some(Ok(TungsteniteMessage::Binary(_)))
|
||||
| Some(Ok(TungsteniteMessage::Frame(_))) => {}
|
||||
Some(Err(error)) => {
|
||||
let _ = tx.send(diagnostic_event(format!(
|
||||
"Backend protocol WebSocket error for {}: {error}",
|
||||
target.display_label()
|
||||
)));
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
) -> Result<Client<WebSocket>, BackendRuntimeClientError> {
|
||||
validate_target(&target)?;
|
||||
let api = BackendApiClient::from_stored_token(&target.base_url)?;
|
||||
let request = protocol_ws_request(&target, &api).map_err(|error| {
|
||||
BackendRuntimeClientError::Protocol(format!(
|
||||
"Backend protocol request could not be constructed for {}: {error}",
|
||||
target.display_label()
|
||||
))
|
||||
})?;
|
||||
match WebSocket::connect(request).await {
|
||||
Ok(socket) => Ok(Client::new(socket)),
|
||||
Err(WebSocketError::WebSocket(error)) => Err(BackendRuntimeClientError::Protocol(
|
||||
protocol_connect_error_message(&target, &api, &error),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn protocol_connect_error_message(
|
||||
target: &BackendRuntimeTarget,
|
||||
api: &BackendApiClient,
|
||||
error: &tokio_tungstenite::tungstenite::Error,
|
||||
) -> String {
|
||||
if let tokio_tungstenite::tungstenite::Error::Http(response) = error {
|
||||
if let Ok(status) = reqwest::StatusCode::from_u16(response.status().as_u16()) {
|
||||
if matches!(
|
||||
status,
|
||||
reqwest::StatusCode::UNAUTHORIZED | reqwest::StatusCode::FORBIDDEN
|
||||
) {
|
||||
if let Err(error) = api.check_status(status) {
|
||||
return error.to_string();
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(error) => {
|
||||
let _ = tx.send(diagnostic_event(format!(
|
||||
"Backend protocol WebSocket connect failed for {}: {error}",
|
||||
target.display_label()
|
||||
)));
|
||||
while commands.recv().await.is_some() {
|
||||
let _ = tx.send(diagnostic_event(format!(
|
||||
"Backend protocol command was not sent because command stream is unavailable for {}",
|
||||
target.display_label()
|
||||
)));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn diagnostic_event(message: impl Into<String>) -> Event {
|
||||
Event::Error {
|
||||
code: ErrorCode::Internal,
|
||||
message: message.into(),
|
||||
}
|
||||
format!(
|
||||
"Backend protocol WebSocket connect failed for {}: {error}",
|
||||
target.display_label()
|
||||
)
|
||||
}
|
||||
|
||||
fn validate_target(target: &BackendRuntimeTarget) -> Result<(), BackendRuntimeClientError> {
|
||||
@@ -496,6 +521,19 @@ fn backend_runtime_worker_restore_path(
|
||||
)
|
||||
}
|
||||
|
||||
fn protocol_ws_request(
|
||||
target: &BackendRuntimeTarget,
|
||||
api: &BackendApiClient,
|
||||
) -> Result<tokio_tungstenite::tungstenite::http::Request<()>, String> {
|
||||
let mut request = protocol_ws_url(target)
|
||||
.into_client_request()
|
||||
.map_err(|error| error.to_string())?;
|
||||
let value = HeaderValue::from_str(&api.authorization_header_value())
|
||||
.map_err(|_| "saved Backend token is not a valid Authorization header".to_string())?;
|
||||
request.headers_mut().insert(AUTHORIZATION, value);
|
||||
Ok(request)
|
||||
}
|
||||
|
||||
fn protocol_ws_url(target: &BackendRuntimeTarget) -> String {
|
||||
let path = format!(
|
||||
"/api/w/{}/runtimes/{}/workers/{}/protocol/ws",
|
||||
@@ -557,6 +595,26 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn protocol_request_attaches_saved_bearer_authorization() {
|
||||
let target = BackendRuntimeTarget::new(
|
||||
"http://127.0.0.1:8787/",
|
||||
"workspace alpha",
|
||||
"runtime/one",
|
||||
"worker one",
|
||||
);
|
||||
let api = BackendApiClient::from_access_token_for_test(
|
||||
"http://127.0.0.1:8787",
|
||||
"websocket-secret",
|
||||
)
|
||||
.unwrap();
|
||||
let request = protocol_ws_request(&target, &api).unwrap();
|
||||
assert_eq!(
|
||||
request.headers().get(AUTHORIZATION).unwrap(),
|
||||
"Bearer websocket-secret"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backend_worker_summary_decodes_current_occupied_workdir_contract() {
|
||||
let payload = serde_json::json!({
|
||||
@@ -572,7 +630,7 @@ mod tests {
|
||||
"capabilities": {"can_stop": true, "can_spawn_followup": false},
|
||||
"working_directory": {
|
||||
"working_directory_id": "wd-1",
|
||||
"repository_id": "main",
|
||||
"repository_key": "main",
|
||||
"materializer_kind": "local_git_worktree",
|
||||
"status": "active",
|
||||
"occupied_by": {
|
||||
@@ -585,13 +643,11 @@ mod tests {
|
||||
});
|
||||
|
||||
let worker: BackendWorkerSummary = serde_json::from_value(payload.clone()).unwrap();
|
||||
let occupied_by = worker
|
||||
.working_directory
|
||||
.unwrap()
|
||||
.occupied_by
|
||||
.expect("occupied Workdir");
|
||||
assert_eq!(occupied_by.worker.runtime_id, "arcadia");
|
||||
assert_eq!(occupied_by.worker.worker_id, "worker-opaque-64");
|
||||
let workdir = worker.working_directory.unwrap();
|
||||
assert_eq!(workdir.repository_key, "main");
|
||||
let occupied_by = workdir.occupied_by.expect("occupied Workdir");
|
||||
assert_eq!(occupied_by.runtime_id, "arcadia");
|
||||
assert_eq!(occupied_by.worker_id, "worker-opaque-64");
|
||||
|
||||
let mut stale = payload;
|
||||
stale["working_directory"]["occupied_by"]["runtime_worker_id"] = serde_json::json!(64);
|
||||
|
||||
@@ -1,18 +1,17 @@
|
||||
use crate::{BackendApiClient, BackendApiClientError};
|
||||
use reqwest::Method;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::fmt;
|
||||
use workspace_api::{RepositoryObservedStatus, RepositorySource};
|
||||
use workspace_api::{
|
||||
WorkspaceCatalogListResponse, WorkspaceCreateResponse, WorkspaceRepositoryRecord,
|
||||
WorkspaceSummary,
|
||||
};
|
||||
|
||||
const DEFAULT_WORKSPACE_LIMIT: usize = 200;
|
||||
|
||||
#[derive(Debug, Clone, Deserialize, PartialEq, Eq)]
|
||||
pub struct BackendWorkspace {
|
||||
pub workspace_id: String,
|
||||
pub owner_account_id: Option<String>,
|
||||
pub display_name: String,
|
||||
pub state: String,
|
||||
pub created_at: String,
|
||||
pub updated_at: String,
|
||||
}
|
||||
pub type BackendWorkspace = WorkspaceSummary;
|
||||
pub type CreateBackendWorkspaceResponse = WorkspaceCreateResponse;
|
||||
pub type CreateBackendWorkspaceRepositoryRecord = WorkspaceRepositoryRecord;
|
||||
|
||||
#[derive(Debug, Clone, Serialize, PartialEq, Eq)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
@@ -30,30 +29,6 @@ pub struct CreateBackendWorkspaceRepository {
|
||||
pub default_ref: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Deserialize, PartialEq, Eq)]
|
||||
pub struct CreateBackendWorkspaceResponse {
|
||||
pub workspace: BackendWorkspace,
|
||||
pub repository: CreateBackendWorkspaceRepositoryRecord,
|
||||
pub config_revision: u64,
|
||||
pub request_fingerprint: String,
|
||||
pub replayed: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Deserialize, PartialEq, Eq)]
|
||||
pub struct CreateBackendWorkspaceRepositoryRecord {
|
||||
pub workspace_id: String,
|
||||
pub repository_id: String,
|
||||
pub name: String,
|
||||
pub kind: String,
|
||||
pub provider: Option<String>,
|
||||
pub source: RepositorySource,
|
||||
pub default_ref: Option<String>,
|
||||
pub source_revision: u64,
|
||||
pub source_fingerprint: String,
|
||||
pub observed_status: RepositoryObservedStatus,
|
||||
pub observed_at: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct BackendWorkspaceCatalogTarget {
|
||||
pub base_url: String,
|
||||
@@ -70,7 +45,7 @@ impl BackendWorkspaceCatalogTarget {
|
||||
#[derive(Debug)]
|
||||
pub enum BackendWorkspaceClientError {
|
||||
InvalidTarget(String),
|
||||
RequestFailed { status: u16, message: String },
|
||||
Api(BackendApiClientError),
|
||||
Http(reqwest::Error),
|
||||
}
|
||||
|
||||
@@ -78,9 +53,7 @@ impl fmt::Display for BackendWorkspaceClientError {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
Self::InvalidTarget(message) => f.write_str(message),
|
||||
Self::RequestFailed { status, message } => {
|
||||
write!(f, "Backend request failed with HTTP {status}: {message}")
|
||||
}
|
||||
Self::Api(error) => write!(f, "{error}"),
|
||||
Self::Http(error) => write!(f, "{error}"),
|
||||
}
|
||||
}
|
||||
@@ -88,6 +61,12 @@ impl fmt::Display for BackendWorkspaceClientError {
|
||||
|
||||
impl std::error::Error for BackendWorkspaceClientError {}
|
||||
|
||||
impl From<BackendApiClientError> for BackendWorkspaceClientError {
|
||||
fn from(error: BackendApiClientError) -> Self {
|
||||
Self::Api(error)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<reqwest::Error> for BackendWorkspaceClientError {
|
||||
fn from(error: reqwest::Error) -> Self {
|
||||
Self::Http(error)
|
||||
@@ -97,56 +76,72 @@ impl From<reqwest::Error> for BackendWorkspaceClientError {
|
||||
pub async fn list_backend_workspaces(
|
||||
target: &BackendWorkspaceCatalogTarget,
|
||||
) -> Result<Vec<BackendWorkspace>, BackendWorkspaceClientError> {
|
||||
validate_target(target)?;
|
||||
let url = format!(
|
||||
"{}/api/workspaces?limit={DEFAULT_WORKSPACE_LIMIT}",
|
||||
target.base_url.trim_end_matches('/')
|
||||
);
|
||||
let response = reqwest::Client::new().get(url).send().await?;
|
||||
let response = require_success(response).await?;
|
||||
Ok(response.json::<Vec<BackendWorkspace>>().await?)
|
||||
let client = BackendApiClient::from_stored_token(&target.base_url)?;
|
||||
list_backend_workspaces_with_client(&client).await
|
||||
}
|
||||
|
||||
async fn list_backend_workspaces_with_client(
|
||||
client: &BackendApiClient,
|
||||
) -> Result<Vec<BackendWorkspace>, BackendWorkspaceClientError> {
|
||||
let response = client
|
||||
.request(
|
||||
Method::GET,
|
||||
&format!("/api/workspaces?limit={DEFAULT_WORKSPACE_LIMIT}"),
|
||||
)?
|
||||
.send()
|
||||
.await?;
|
||||
client.check_status(response.status())?;
|
||||
Ok(response.json::<WorkspaceCatalogListResponse>().await?.0)
|
||||
}
|
||||
|
||||
pub async fn create_backend_workspace(
|
||||
target: &BackendWorkspaceCatalogTarget,
|
||||
request: &CreateBackendWorkspaceRequest,
|
||||
) -> Result<CreateBackendWorkspaceResponse, BackendWorkspaceClientError> {
|
||||
validate_target(target)?;
|
||||
let url = format!("{}/api/workspaces", target.base_url.trim_end_matches('/'));
|
||||
let response = reqwest::Client::new()
|
||||
.post(url)
|
||||
let client = BackendApiClient::from_stored_token(&target.base_url)?;
|
||||
let response = client
|
||||
.request(Method::POST, "/api/workspaces")?
|
||||
.json(request)
|
||||
.send()
|
||||
.await?;
|
||||
let response = require_success(response).await?;
|
||||
client.check_status(response.status())?;
|
||||
Ok(response.json::<CreateBackendWorkspaceResponse>().await?)
|
||||
}
|
||||
|
||||
async fn require_success(
|
||||
response: reqwest::Response,
|
||||
) -> Result<reqwest::Response, BackendWorkspaceClientError> {
|
||||
if response.status().is_success() {
|
||||
return Ok(response);
|
||||
}
|
||||
let status = response.status().as_u16();
|
||||
let message = response.text().await.unwrap_or_default();
|
||||
Err(BackendWorkspaceClientError::RequestFailed { status, message })
|
||||
}
|
||||
|
||||
fn validate_target(
|
||||
target: &BackendWorkspaceCatalogTarget,
|
||||
) -> Result<(), BackendWorkspaceClientError> {
|
||||
if !(target.base_url.starts_with("http://") || target.base_url.starts_with("https://")) {
|
||||
return Err(BackendWorkspaceClientError::InvalidTarget(
|
||||
"Backend API base URL must start with http:// or https://".to_string(),
|
||||
));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::io::{Read, Write};
|
||||
use std::net::TcpListener;
|
||||
use std::thread;
|
||||
|
||||
#[tokio::test]
|
||||
async fn workspace_catalog_request_uses_shared_bearer_client() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
|
||||
let base_url = format!("http://{}", listener.local_addr().unwrap());
|
||||
let handle = thread::spawn(move || {
|
||||
let (mut stream, _) = listener.accept().unwrap();
|
||||
let mut request = vec![0; 4096];
|
||||
let read = stream.read(&mut request).unwrap();
|
||||
let request = String::from_utf8_lossy(&request[..read]).to_ascii_lowercase();
|
||||
assert!(request.starts_with("get /api/workspaces?limit=200 "));
|
||||
assert!(request.contains("authorization: bearer catalog-secret\r\n"));
|
||||
stream
|
||||
.write_all(
|
||||
b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 2\r\nConnection: close\r\n\r\n[]",
|
||||
)
|
||||
.unwrap();
|
||||
});
|
||||
let client =
|
||||
BackendApiClient::from_access_token_for_test(&base_url, "catalog-secret").unwrap();
|
||||
assert!(
|
||||
list_backend_workspaces_with_client(&client)
|
||||
.await
|
||||
.unwrap()
|
||||
.is_empty()
|
||||
);
|
||||
handle.join().unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn create_request_keeps_operation_key_for_exact_retry() {
|
||||
|
||||
@@ -0,0 +1,137 @@
|
||||
use std::error::Error;
|
||||
use std::fmt;
|
||||
|
||||
use protocol::stream::{decode_event, encode_method};
|
||||
use protocol::{Event, Method};
|
||||
|
||||
use crate::transport::Socket;
|
||||
|
||||
/// Typed Worker protocol client over an injected message transport.
|
||||
pub struct Client<T> {
|
||||
socket: T,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum ClientError<E> {
|
||||
Transport(E),
|
||||
Protocol(serde_json::Error),
|
||||
}
|
||||
|
||||
impl<T> Client<T> {
|
||||
pub fn new(socket: T) -> Self {
|
||||
Self { socket }
|
||||
}
|
||||
|
||||
pub fn into_inner(self) -> T {
|
||||
self.socket
|
||||
}
|
||||
}
|
||||
|
||||
impl<T: Socket> Client<T> {
|
||||
pub async fn send(&mut self, method: &Method) -> Result<(), ClientError<T::Error>> {
|
||||
let message = encode_method(method).map_err(ClientError::Protocol)?;
|
||||
self.socket
|
||||
.send(message)
|
||||
.await
|
||||
.map_err(ClientError::Transport)
|
||||
}
|
||||
|
||||
pub async fn next_event(&mut self) -> Result<Option<Event>, ClientError<T::Error>> {
|
||||
self.socket
|
||||
.next()
|
||||
.await
|
||||
.map_err(ClientError::Transport)?
|
||||
.map(|message| decode_event(&message).map_err(ClientError::Protocol))
|
||||
.transpose()
|
||||
}
|
||||
|
||||
pub fn try_next_event(&mut self) -> Result<Option<Event>, ClientError<T::Error>> {
|
||||
self.socket
|
||||
.try_next()
|
||||
.map_err(ClientError::Transport)?
|
||||
.map(|message| decode_event(&message).map_err(ClientError::Protocol))
|
||||
.transpose()
|
||||
}
|
||||
}
|
||||
|
||||
impl<E: fmt::Display> fmt::Display for ClientError<E> {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
Self::Transport(error) => write!(formatter, "Worker transport error: {error}"),
|
||||
Self::Protocol(error) => write!(formatter, "Worker protocol error: {error}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<E: Error + 'static> Error for ClientError<E> {
|
||||
fn source(&self) -> Option<&(dyn Error + 'static)> {
|
||||
match self {
|
||||
Self::Transport(error) => Some(error),
|
||||
Self::Protocol(error) => Some(error),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::VecDeque;
|
||||
use std::convert::Infallible;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use protocol::stream::{decode_method, encode_event};
|
||||
use protocol::{Event, Method, WorkerStatus};
|
||||
|
||||
use super::Client;
|
||||
use crate::transport::Socket;
|
||||
|
||||
#[derive(Default)]
|
||||
struct TestSocket {
|
||||
sent: Vec<String>,
|
||||
incoming: VecDeque<String>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Socket for TestSocket {
|
||||
type Error = Infallible;
|
||||
|
||||
async fn send(&mut self, message: String) -> Result<(), Self::Error> {
|
||||
self.sent.push(message);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn next(&mut self) -> Result<Option<String>, Self::Error> {
|
||||
Ok(self.incoming.pop_front())
|
||||
}
|
||||
|
||||
fn try_next(&mut self) -> Result<Option<String>, Self::Error> {
|
||||
Ok(self.incoming.pop_front())
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn encodes_methods_and_decodes_events_above_transport() {
|
||||
let mut socket = TestSocket::default();
|
||||
socket.incoming.push_back(
|
||||
encode_event(&Event::Status {
|
||||
status: WorkerStatus::Idle,
|
||||
})
|
||||
.expect("encode event"),
|
||||
);
|
||||
let mut client = Client::new(socket);
|
||||
|
||||
client
|
||||
.send(&Method::run_text("hello"))
|
||||
.await
|
||||
.expect("send method");
|
||||
assert!(matches!(
|
||||
decode_method(&client.socket.sent[0]),
|
||||
Ok(Method::Run { .. })
|
||||
));
|
||||
assert!(matches!(
|
||||
client.next_event().await,
|
||||
Ok(Some(Event::Status {
|
||||
status: WorkerStatus::Idle
|
||||
}))
|
||||
));
|
||||
}
|
||||
}
|
||||
+23
-32
@@ -1,57 +1,48 @@
|
||||
//! Worker プロトコルを喋るクライアント。
|
||||
//! Backend Workspace/Runtime と既存 Worker protocol へ接続するクライアント。
|
||||
//!
|
||||
//! - [`WorkerClient`]: 既存 worker の Unix ソケットへ接続して `Method` を送り、
|
||||
//! `Event` を受け取る低レベル接続。
|
||||
//! - [`spawn`]: worker バイナリをサブプロセスとして起動し、`YOI-READY`
|
||||
//! ハンドシェイクが終わるまで待つフロー。subprocess を立ち上げる必要が
|
||||
//! ない呼び出し側 (=既存 worker に attach する場合) は使わなくてよい。
|
||||
//!
|
||||
//! TUI / GUI / E2E ハーネスはこの crate に依存して protocol を喋る。
|
||||
//! Standalone execution is owned by the `standalone` crate and does not spawn
|
||||
//! a Worker subprocess through this crate.
|
||||
|
||||
pub mod backend_auth;
|
||||
pub mod backend_api;
|
||||
mod backend_auth;
|
||||
pub mod backend_runtime;
|
||||
pub mod backend_workspace;
|
||||
pub mod runtime_command;
|
||||
pub mod spawn;
|
||||
mod client;
|
||||
pub mod target;
|
||||
pub mod ticket_role;
|
||||
mod worker_client;
|
||||
pub mod transport;
|
||||
mod workspace_product;
|
||||
|
||||
pub use backend_api::{
|
||||
BackendApiClient, BackendApiClientError, BackendOrigin, backend_token_file_path,
|
||||
save_backend_token,
|
||||
};
|
||||
pub use backend_auth::{
|
||||
BackendAuthClientError, BackendAuthTarget, DeviceLoginPollResponse, DeviceLoginStartResponse,
|
||||
poll_device_login, start_device_login, wait_for_device_login,
|
||||
};
|
||||
pub use backend_runtime::{
|
||||
BackendDiagnostic, BackendDiagnosticSeverity, BackendRuntimeClient, BackendRuntimeClientError,
|
||||
BackendDiagnostic, BackendDiagnosticSeverity, BackendRuntimeClientError,
|
||||
BackendRuntimeListResponse, BackendRuntimeListTarget, BackendRuntimeSummary,
|
||||
BackendRuntimeTarget, BackendWorkerCapabilitySummary, BackendWorkerImplementationSummary,
|
||||
BackendWorkerRestoreResponse, BackendWorkerRestoreResult, BackendWorkerSummary,
|
||||
BackendWorkerWorkspaceSummary, BackendWorkingDirectorySummary, list_backend_stopped_workers,
|
||||
list_backend_workers, restore_backend_worker,
|
||||
BackendWorkerWorkspaceSummary, BackendWorkingDirectorySummary, connect_backend_runtime,
|
||||
list_backend_stopped_workers, list_backend_workers, restore_backend_worker,
|
||||
};
|
||||
pub use backend_workspace::{
|
||||
BackendWorkspace, BackendWorkspaceCatalogTarget, BackendWorkspaceClientError,
|
||||
CreateBackendWorkspaceRepository, CreateBackendWorkspaceRequest,
|
||||
CreateBackendWorkspaceResponse, create_backend_workspace, list_backend_workspaces,
|
||||
};
|
||||
pub use runtime_command::WorkerRuntimeCommand;
|
||||
pub use client::{Client, ClientError};
|
||||
pub use target::{
|
||||
BackendTarget, Dashboard, LocalTarget, ResolvedTarget, Target, TargetError, TargetKind,
|
||||
WorkerByName, WorkerConnection, WorkerConnectionSelector, WorkerList, WorkerListRequest,
|
||||
WorkerResume, WorkerSpawn,
|
||||
BackendTarget, Dashboard, ResolvedTarget, StandaloneTarget, StandaloneWorkerListIntent,
|
||||
StandaloneWorkerResumeIntent, Target, TargetError, TargetKind, WorkerConnection,
|
||||
WorkerConnectionSelector, WorkerList, WorkerListRequest, WorkerSpawn,
|
||||
};
|
||||
|
||||
pub use spawn::{
|
||||
SpawnConfig, SpawnError, SpawnReady, WorkerProcessLaunchConfig, WorkerProcessLaunchOptions,
|
||||
spawn_worker, spawn_worker_with_options,
|
||||
pub use workspace_api::{
|
||||
CompanionCancelRequest, CompanionLifecycleState, CompanionMessageDisposition,
|
||||
CompanionMessageRequest, CompanionMessageResponse, CompanionStatusResponse,
|
||||
CompanionTranscriptItem, CompanionTranscriptProjection, CompanionTranscriptRole,
|
||||
CompanionTransportSummary, ObjectiveDetail, ObjectiveSummary,
|
||||
};
|
||||
pub use ticket_role::{
|
||||
TicketRef, TicketRoleLaunchContext, TicketRoleLaunchError, TicketRoleLaunchOptions,
|
||||
TicketRoleLaunchPlan, TicketRoleLaunchResult, TicketRolePreRunWarning,
|
||||
launch_ticket_role_worker, launch_ticket_role_worker_with_options, plan_ticket_role_launch,
|
||||
plan_ticket_role_launch_with_config,
|
||||
};
|
||||
pub use worker_client::WorkerClient;
|
||||
pub use workspace_api::{ObjectiveDetail, ObjectiveSummary};
|
||||
pub use workspace_product::BackendWorkspaceProductClient;
|
||||
|
||||
@@ -1,435 +0,0 @@
|
||||
//! Worker runtime command をサブプロセスとして立ち上げ、`YOI-READY` を待つ
|
||||
//! ハンドシェイク。
|
||||
//!
|
||||
//! - 親プロセス (TUI / GUI / E2E) は profile/default/typed restore flags を
|
||||
//! 指定してこの関数に渡す。worker はそれを受けて socket を bind し、stderr に
|
||||
//! `YOI-READY\t<name>\t<socket>` を吐く。
|
||||
//! - 待機中の stderr 行は `progress` コールバック越しに呼び出し側へ流す。
|
||||
//! UI の進捗表示や E2E のログ収集はここで賄う。
|
||||
//! - `kill_on_drop = false` + `process_group(0)` により、親プロセス
|
||||
//! ライフサイクルから切り離した detached worker を作る。ready 後の lifecycle
|
||||
//! 管理は runtime ディレクトリ / socket を介して行う。
|
||||
|
||||
use std::io;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::process::Stdio;
|
||||
use std::time::Duration;
|
||||
|
||||
use crate::WorkerRuntimeCommand;
|
||||
use tokio::process::Command;
|
||||
use uuid::Uuid;
|
||||
|
||||
const READY_PREFIX: &str = "YOI-READY\t";
|
||||
const READY_TIMEOUT: Duration = Duration::from_secs(20);
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct WorkerProcessLaunchConfig {
|
||||
pub runtime_command: WorkerRuntimeCommand,
|
||||
/// `worker.name` として使う識別子。runtime ディレクトリ
|
||||
/// (`manifest::paths::worker_runtime_dir`) の解決と、ready 行に乗る
|
||||
/// 名前との突き合わせに使う。
|
||||
pub worker_name: String,
|
||||
/// Optional reusable Profile selector. Worker identity is always supplied
|
||||
/// separately with `--worker`; profile selection must not imply a name.
|
||||
pub profile: Option<String>,
|
||||
/// Explicit runtime workspace root. The child receives it via
|
||||
/// `--workspace` so startup does not infer workspace identity from the
|
||||
/// parent process cwd.
|
||||
pub workspace_root: PathBuf,
|
||||
/// Optional child process cwd. This is not runtime workspace identity and
|
||||
/// is not passed as a CLI argument; the child observes it as its ordinary
|
||||
/// process current directory.
|
||||
pub cwd: Option<PathBuf>,
|
||||
/// `Some(id)` のとき `--session <id>` を付与し、当該セッションから
|
||||
/// resume させる。
|
||||
pub resume_from: Option<Uuid>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Default)]
|
||||
pub struct WorkerProcessLaunchOptions {
|
||||
/// Extra child CLI arguments supplied by an upper resolver layer. The
|
||||
/// low-level launch config intentionally does not model Ticket IDs,
|
||||
/// Ticket roles, orchestration roles, executable authority, or raw
|
||||
/// browser-provided profile/cwd/workspace inputs.
|
||||
pub extra_args: Vec<String>,
|
||||
}
|
||||
|
||||
impl WorkerProcessLaunchOptions {
|
||||
pub fn with_hidden_arg(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
|
||||
self.extra_args.extend([name.into(), value.into()]);
|
||||
self
|
||||
}
|
||||
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.extra_args.is_empty()
|
||||
}
|
||||
}
|
||||
|
||||
pub type SpawnConfig = WorkerProcessLaunchConfig;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct SpawnReady {
|
||||
pub worker_name: String,
|
||||
pub socket_path: PathBuf,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum SpawnError {
|
||||
Io(io::Error),
|
||||
/// runtime ディレクトリが解決できなかった (環境変数未設定等)。
|
||||
RuntimeDirUnavailable,
|
||||
WorkerLaunchFailed {
|
||||
command: WorkerRuntimeCommand,
|
||||
source: io::Error,
|
||||
},
|
||||
WorkerExitedEarly {
|
||||
stderr_tail: String,
|
||||
},
|
||||
Timeout,
|
||||
}
|
||||
|
||||
impl std::fmt::Display for SpawnError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::Io(e) => write!(f, "io error: {e}"),
|
||||
Self::RuntimeDirUnavailable => write!(
|
||||
f,
|
||||
"could not resolve runtime directory (set YOI_HOME, YOI_RUNTIME_DIR, XDG_RUNTIME_DIR, or HOME)"
|
||||
),
|
||||
Self::WorkerLaunchFailed { command, source } => write!(
|
||||
f,
|
||||
"failed to launch worker runtime command `{command}`: {source}"
|
||||
),
|
||||
Self::WorkerExitedEarly { stderr_tail } => {
|
||||
if stderr_tail.is_empty() {
|
||||
write!(f, "worker exited before becoming ready")
|
||||
} else {
|
||||
write!(f, "worker exited before becoming ready: {stderr_tail}")
|
||||
}
|
||||
}
|
||||
Self::Timeout => write!(
|
||||
f,
|
||||
"worker did not become ready within {}s",
|
||||
READY_TIMEOUT.as_secs()
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for SpawnError {
|
||||
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
|
||||
match self {
|
||||
Self::Io(error) | Self::WorkerLaunchFailed { source: error, .. } => Some(error),
|
||||
Self::RuntimeDirUnavailable | Self::WorkerExitedEarly { .. } | Self::Timeout => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<io::Error> for SpawnError {
|
||||
fn from(e: io::Error) -> Self {
|
||||
Self::Io(e)
|
||||
}
|
||||
}
|
||||
|
||||
fn runtime_args(
|
||||
config: &WorkerProcessLaunchConfig,
|
||||
options: &WorkerProcessLaunchOptions,
|
||||
) -> Vec<String> {
|
||||
let mut args = vec![
|
||||
"--workspace".to_string(),
|
||||
config.workspace_root.display().to_string(),
|
||||
];
|
||||
if let Some(id) = config.resume_from {
|
||||
args.extend([
|
||||
"--session".to_string(),
|
||||
id.to_string(),
|
||||
"--worker".to_string(),
|
||||
config.worker_name.clone(),
|
||||
]);
|
||||
} else {
|
||||
args.extend(["--worker".to_string(), config.worker_name.clone()]);
|
||||
if let Some(profile) = &config.profile {
|
||||
args.extend(["--profile".to_string(), profile.clone()]);
|
||||
}
|
||||
}
|
||||
args.extend(options.extra_args.clone());
|
||||
args
|
||||
}
|
||||
|
||||
/// worker を spawn し、`YOI-READY` ハンドシェイクが終わるまで待つ。
|
||||
///
|
||||
/// `progress` は ready 行を見つけるまでに観測した stderr の各行で呼ばれる
|
||||
/// (ready 行自体は除外される)。UI の表示更新や E2E ログ取得に使う。
|
||||
pub async fn spawn_worker<F>(
|
||||
config: WorkerProcessLaunchConfig,
|
||||
progress: F,
|
||||
) -> Result<SpawnReady, SpawnError>
|
||||
where
|
||||
F: FnMut(&str),
|
||||
{
|
||||
spawn_worker_with_options(config, WorkerProcessLaunchOptions::default(), progress).await
|
||||
}
|
||||
|
||||
pub async fn spawn_worker_with_options<F>(
|
||||
config: WorkerProcessLaunchConfig,
|
||||
options: WorkerProcessLaunchOptions,
|
||||
mut progress: F,
|
||||
) -> Result<SpawnReady, SpawnError>
|
||||
where
|
||||
F: FnMut(&str),
|
||||
{
|
||||
let worker_runtime_dir = manifest::paths::worker_runtime_dir(&config.worker_name)
|
||||
.ok_or(SpawnError::RuntimeDirUnavailable)?;
|
||||
std::fs::create_dir_all(&worker_runtime_dir).map_err(SpawnError::Io)?;
|
||||
let stderr_path = worker_runtime_dir.join("stderr.log");
|
||||
let stderr_file = std::fs::File::create(&stderr_path).map_err(SpawnError::Io)?;
|
||||
|
||||
let mut command = Command::new(config.runtime_command.program());
|
||||
command
|
||||
.args(config.runtime_command.prefix_args())
|
||||
.current_dir(config.cwd.as_ref().unwrap_or(&config.workspace_root))
|
||||
.stdin(Stdio::null())
|
||||
.stdout(Stdio::null())
|
||||
.stderr(Stdio::from(stderr_file))
|
||||
.process_group(0);
|
||||
for arg in runtime_args(&config, &options) {
|
||||
command.arg(arg);
|
||||
}
|
||||
let mut child = command
|
||||
.spawn()
|
||||
.map_err(|source| SpawnError::WorkerLaunchFailed {
|
||||
command: config.runtime_command.clone(),
|
||||
source,
|
||||
})?;
|
||||
|
||||
// Default `kill_on_drop = false` plus `process_group(0)` makes this
|
||||
// a detached Worker once startup succeeds: dropping the handle does not
|
||||
// terminate it, and terminal-generated signals for the parent's
|
||||
// process group do not hit the Worker. Runtime state/socket files are
|
||||
// the source of truth after that point.
|
||||
let ready = match wait_for_ready_file(&mut progress, &stderr_path, &mut child).await {
|
||||
Ok(ready) => ready,
|
||||
Err(e) => {
|
||||
let _ = child.start_kill();
|
||||
let _ = child.wait().await;
|
||||
return Err(e);
|
||||
}
|
||||
};
|
||||
tokio::spawn(async move {
|
||||
let _ = child.wait().await;
|
||||
});
|
||||
Ok(ready)
|
||||
}
|
||||
|
||||
async fn wait_for_ready_file<F>(
|
||||
progress: &mut F,
|
||||
stderr_path: &Path,
|
||||
child: &mut tokio::process::Child,
|
||||
) -> Result<SpawnReady, SpawnError>
|
||||
where
|
||||
F: FnMut(&str),
|
||||
{
|
||||
let mut tail = StderrTail::new();
|
||||
let deadline = tokio::time::Instant::now() + READY_TIMEOUT;
|
||||
let mut offset = 0usize;
|
||||
|
||||
loop {
|
||||
let content = match tokio::fs::read_to_string(stderr_path).await {
|
||||
Ok(content) => content,
|
||||
Err(e) if e.kind() == io::ErrorKind::NotFound => String::new(),
|
||||
Err(e) => return Err(SpawnError::Io(e)),
|
||||
};
|
||||
if content.len() > offset {
|
||||
for line in content[offset..].lines() {
|
||||
if let Some(rest) = line.strip_prefix(READY_PREFIX) {
|
||||
let mut parts = rest.splitn(2, '\t');
|
||||
let worker_name = parts.next().unwrap_or("").to_string();
|
||||
let socket_str = parts.next().unwrap_or("").to_string();
|
||||
if worker_name.is_empty() || socket_str.is_empty() {
|
||||
return Err(SpawnError::WorkerExitedEarly {
|
||||
stderr_tail: format!("malformed ready line: {line}"),
|
||||
});
|
||||
}
|
||||
let socket_path = PathBuf::from(socket_str);
|
||||
wait_for_socket(
|
||||
&socket_path,
|
||||
deadline,
|
||||
child,
|
||||
stderr_path,
|
||||
&mut tail,
|
||||
&mut offset,
|
||||
)
|
||||
.await?;
|
||||
return Ok(SpawnReady {
|
||||
worker_name,
|
||||
socket_path,
|
||||
});
|
||||
}
|
||||
tail.push(line);
|
||||
progress(line);
|
||||
}
|
||||
offset = content.len();
|
||||
}
|
||||
|
||||
if tokio::time::Instant::now() >= deadline {
|
||||
return Err(SpawnError::Timeout);
|
||||
}
|
||||
tokio::select! {
|
||||
status = child.wait() => {
|
||||
let _ = status;
|
||||
// Worker は exit 直前に最終 stderr 行を flush することがある。
|
||||
// child.wait() が解決した後に再読みして、原因行を取りこ
|
||||
// ぼさず WorkerExitedEarly に載せる。
|
||||
drain_stderr_into_tail(stderr_path, &mut tail, &mut offset).await;
|
||||
return Err(SpawnError::WorkerExitedEarly {
|
||||
stderr_tail: tail.into_string(),
|
||||
});
|
||||
}
|
||||
_ = tokio::time::sleep(Duration::from_millis(100)) => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn wait_for_socket(
|
||||
socket_path: &Path,
|
||||
deadline: tokio::time::Instant,
|
||||
child: &mut tokio::process::Child,
|
||||
stderr_path: &Path,
|
||||
tail: &mut StderrTail,
|
||||
offset: &mut usize,
|
||||
) -> Result<(), SpawnError> {
|
||||
loop {
|
||||
match tokio::net::UnixStream::connect(socket_path).await {
|
||||
Ok(_) => return Ok(()),
|
||||
Err(e)
|
||||
if e.kind() == io::ErrorKind::NotFound
|
||||
|| e.kind() == io::ErrorKind::ConnectionRefused => {}
|
||||
Err(e) => return Err(SpawnError::Io(e)),
|
||||
}
|
||||
if tokio::time::Instant::now() >= deadline {
|
||||
return Err(SpawnError::Timeout);
|
||||
}
|
||||
tokio::select! {
|
||||
status = child.wait() => {
|
||||
let _ = status;
|
||||
drain_stderr_into_tail(stderr_path, tail, offset).await;
|
||||
return Err(SpawnError::WorkerExitedEarly {
|
||||
stderr_tail: tail.as_string(),
|
||||
});
|
||||
}
|
||||
_ = tokio::time::sleep(Duration::from_millis(50)) => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn drain_stderr_into_tail(stderr_path: &Path, tail: &mut StderrTail, offset: &mut usize) {
|
||||
let Ok(content) = tokio::fs::read_to_string(stderr_path).await else {
|
||||
return;
|
||||
};
|
||||
if content.len() <= *offset {
|
||||
return;
|
||||
}
|
||||
for line in content[*offset..].lines() {
|
||||
if !line.starts_with(READY_PREFIX) {
|
||||
tail.push(line);
|
||||
}
|
||||
}
|
||||
*offset = content.len();
|
||||
}
|
||||
|
||||
struct StderrTail {
|
||||
lines: std::collections::VecDeque<String>,
|
||||
}
|
||||
|
||||
impl StderrTail {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
lines: std::collections::VecDeque::with_capacity(8),
|
||||
}
|
||||
}
|
||||
fn push(&mut self, line: &str) {
|
||||
if self.lines.len() == 8 {
|
||||
self.lines.pop_front();
|
||||
}
|
||||
self.lines.push_back(line.to_string());
|
||||
}
|
||||
fn as_string(&self) -> String {
|
||||
self.lines.iter().cloned().collect::<Vec<_>>().join(" | ")
|
||||
}
|
||||
fn into_string(self) -> String {
|
||||
self.lines.into_iter().collect::<Vec<_>>().join(" | ")
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::ffi::OsString;
|
||||
|
||||
fn base_config() -> WorkerProcessLaunchConfig {
|
||||
WorkerProcessLaunchConfig {
|
||||
runtime_command: WorkerRuntimeCommand::new("/bin/yoi", vec![OsString::from("worker")]),
|
||||
worker_name: "explicit-worker".to_string(),
|
||||
profile: Some("project:companion".to_string()),
|
||||
workspace_root: PathBuf::from("/work/other-project"),
|
||||
cwd: None,
|
||||
resume_from: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_args_keep_workspace_worker_and_profile_separate() {
|
||||
assert_eq!(
|
||||
runtime_args(&base_config(), &WorkerProcessLaunchOptions::default()),
|
||||
vec![
|
||||
"--workspace",
|
||||
"/work/other-project",
|
||||
"--worker",
|
||||
"explicit-worker",
|
||||
"--profile",
|
||||
"project:companion",
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_args_use_session_mode_without_profile_identity_alias() {
|
||||
let mut config = base_config();
|
||||
config.resume_from = Some(Uuid::nil());
|
||||
assert_eq!(
|
||||
runtime_args(&config, &WorkerProcessLaunchOptions::default()),
|
||||
vec![
|
||||
"--workspace",
|
||||
"/work/other-project",
|
||||
"--session",
|
||||
"00000000-0000-0000-0000-000000000000",
|
||||
"--worker",
|
||||
"explicit-worker",
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_args_include_upper_resolver_extra_args_without_child_cwd() {
|
||||
let mut config = base_config();
|
||||
config.cwd = Some(PathBuf::from("/work/main/.worktree/orchestration/yoi"));
|
||||
|
||||
assert_eq!(
|
||||
runtime_args(
|
||||
&config,
|
||||
&WorkerProcessLaunchOptions::default()
|
||||
.with_hidden_arg("--ticket-role", "orchestrator"),
|
||||
),
|
||||
vec![
|
||||
"--workspace",
|
||||
"/work/other-project",
|
||||
"--worker",
|
||||
"explicit-worker",
|
||||
"--profile",
|
||||
"project:companion",
|
||||
"--ticket-role",
|
||||
"orchestrator",
|
||||
]
|
||||
);
|
||||
}
|
||||
}
|
||||
+161
-194
@@ -1,16 +1,20 @@
|
||||
use std::fmt;
|
||||
use std::{fmt, path::PathBuf};
|
||||
|
||||
use crate::{BackendRuntimeListTarget, BackendRuntimeTarget, WorkerRuntimeCommand};
|
||||
use crate::{
|
||||
BackendApiClient, BackendApiClientError, BackendOrigin, BackendRuntimeListTarget,
|
||||
BackendRuntimeTarget,
|
||||
};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum TargetKind {
|
||||
Local,
|
||||
/// One-process Standalone authority with no Runtime or Workspace backend.
|
||||
Standalone,
|
||||
Backend,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum ResolvedTarget {
|
||||
Local,
|
||||
Standalone,
|
||||
Backend {
|
||||
base_url: String,
|
||||
workspace_id: String,
|
||||
@@ -20,7 +24,7 @@ pub enum ResolvedTarget {
|
||||
impl ResolvedTarget {
|
||||
pub fn kind(&self) -> TargetKind {
|
||||
match self {
|
||||
Self::Local => TargetKind::Local,
|
||||
Self::Standalone => TargetKind::Standalone,
|
||||
Self::Backend { .. } => TargetKind::Backend,
|
||||
}
|
||||
}
|
||||
@@ -29,31 +33,12 @@ impl ResolvedTarget {
|
||||
impl fmt::Display for TargetKind {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
Self::Local => f.write_str("local"),
|
||||
Self::Standalone => f.write_str("Standalone"),
|
||||
Self::Backend => f.write_str("Backend"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct LocalTarget;
|
||||
|
||||
impl LocalTarget {
|
||||
pub fn new() -> Self {
|
||||
Self
|
||||
}
|
||||
|
||||
fn runtime_command(&self) -> Result<WorkerRuntimeCommand, TargetError> {
|
||||
WorkerRuntimeCommand::resolve().map_err(TargetError::local_runtime_command)
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for LocalTarget {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct BackendTarget {
|
||||
pub base_url: String,
|
||||
@@ -62,11 +47,19 @@ pub struct BackendTarget {
|
||||
|
||||
impl BackendTarget {
|
||||
pub fn new(base_url: impl Into<String>, workspace_id: Option<impl Into<String>>) -> Self {
|
||||
let base_url = base_url.into();
|
||||
let base_url = BackendOrigin::parse(&base_url)
|
||||
.map(|origin| origin.to_string())
|
||||
.unwrap_or(base_url);
|
||||
Self {
|
||||
base_url: base_url.into(),
|
||||
base_url,
|
||||
workspace_id: workspace_id.map(Into::into),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn authenticated_client(&self) -> Result<BackendApiClient, BackendApiClientError> {
|
||||
BackendApiClient::from_stored_token(&self.base_url)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
@@ -108,34 +101,31 @@ impl WorkerConnectionSelector {
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct WorkerSpawn {
|
||||
pub runtime_command: WorkerRuntimeCommand,
|
||||
pub state_dir: PathBuf,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct WorkerByName {
|
||||
pub runtime_command: WorkerRuntimeCommand,
|
||||
pub struct StandaloneWorkerListIntent {
|
||||
pub state_dir: PathBuf,
|
||||
pub cwd: PathBuf,
|
||||
pub include_all: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct WorkerResume {
|
||||
pub runtime_command: WorkerRuntimeCommand,
|
||||
pub struct StandaloneWorkerResumeIntent {
|
||||
pub state_dir: PathBuf,
|
||||
pub worker_id: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum Dashboard {
|
||||
Local {
|
||||
runtime_command: WorkerRuntimeCommand,
|
||||
},
|
||||
Backend {
|
||||
base_url: String,
|
||||
workspace_id: String,
|
||||
},
|
||||
pub struct Dashboard {
|
||||
pub base_url: String,
|
||||
pub workspace_id: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct WorkerList {
|
||||
pub local_runtime_command: Option<WorkerRuntimeCommand>,
|
||||
pub backend_target: Option<BackendRuntimeListTarget>,
|
||||
pub backend_target: BackendRuntimeListTarget,
|
||||
pub include_stopped: bool,
|
||||
}
|
||||
|
||||
@@ -161,12 +151,6 @@ impl TargetError {
|
||||
message: format!("invalid {target} target: {}", message.into()),
|
||||
}
|
||||
}
|
||||
|
||||
fn local_runtime_command(error: std::io::Error) -> Self {
|
||||
Self {
|
||||
message: format!("failed to resolve local Worker runtime command: {error}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for TargetError {
|
||||
@@ -183,71 +167,40 @@ pub trait Target: fmt::Debug + Send + Sync {
|
||||
/// Resolve the target once for Workspace product-state operations.
|
||||
///
|
||||
/// Backend targets must carry an explicit Workspace identity. Callers use
|
||||
/// this value instead of rediscovering Backend/local authority from cwd or
|
||||
/// process configuration after command dispatch.
|
||||
/// this value instead of rediscovering authority from cwd or process
|
||||
/// configuration after command dispatch.
|
||||
fn resolve(&self) -> Result<ResolvedTarget, TargetError>;
|
||||
|
||||
fn spawn_worker(&self) -> Result<WorkerSpawn, TargetError>;
|
||||
|
||||
fn worker_by_name(&self) -> Result<WorkerByName, TargetError>;
|
||||
|
||||
fn resume_worker(&self) -> Result<WorkerResume, TargetError>;
|
||||
|
||||
fn dashboard(&self) -> Result<Dashboard, TargetError>;
|
||||
|
||||
fn list_workers(&self, request: WorkerListRequest) -> Result<WorkerList, TargetError>;
|
||||
|
||||
fn connect_worker(
|
||||
&self,
|
||||
selector: WorkerConnectionSelector,
|
||||
) -> Result<WorkerConnection, TargetError>;
|
||||
}
|
||||
|
||||
impl Target for LocalTarget {
|
||||
fn kind(&self) -> TargetKind {
|
||||
TargetKind::Local
|
||||
}
|
||||
|
||||
fn resolve(&self) -> Result<ResolvedTarget, TargetError> {
|
||||
Ok(ResolvedTarget::Local)
|
||||
}
|
||||
|
||||
fn spawn_worker(&self) -> Result<WorkerSpawn, TargetError> {
|
||||
Ok(WorkerSpawn {
|
||||
runtime_command: self.runtime_command()?,
|
||||
})
|
||||
Err(TargetError::unsupported("Worker spawn", self.kind()))
|
||||
}
|
||||
|
||||
fn worker_by_name(&self) -> Result<WorkerByName, TargetError> {
|
||||
Ok(WorkerByName {
|
||||
runtime_command: self.runtime_command()?,
|
||||
})
|
||||
fn standalone_worker_list(
|
||||
&self,
|
||||
_include_all: bool,
|
||||
) -> Result<StandaloneWorkerListIntent, TargetError> {
|
||||
Err(TargetError::unsupported(
|
||||
"standalone Worker listing",
|
||||
self.kind(),
|
||||
))
|
||||
}
|
||||
|
||||
fn resume_worker(&self) -> Result<WorkerResume, TargetError> {
|
||||
Ok(WorkerResume {
|
||||
runtime_command: self.runtime_command()?,
|
||||
})
|
||||
fn standalone_worker_resume(
|
||||
&self,
|
||||
_worker_id: String,
|
||||
) -> Result<StandaloneWorkerResumeIntent, TargetError> {
|
||||
Err(TargetError::unsupported(
|
||||
"standalone Worker restore",
|
||||
self.kind(),
|
||||
))
|
||||
}
|
||||
|
||||
fn dashboard(&self) -> Result<Dashboard, TargetError> {
|
||||
Ok(Dashboard::Local {
|
||||
runtime_command: self.runtime_command()?,
|
||||
})
|
||||
Err(TargetError::unsupported("Worker dashboard", self.kind()))
|
||||
}
|
||||
|
||||
fn list_workers(&self, request: WorkerListRequest) -> Result<WorkerList, TargetError> {
|
||||
if request.runtime_id.is_some() {
|
||||
return Err(TargetError::unsupported(
|
||||
"Explicit runtime id for local worker listing",
|
||||
self.kind(),
|
||||
));
|
||||
}
|
||||
Ok(WorkerList {
|
||||
local_runtime_command: Some(self.runtime_command()?),
|
||||
backend_target: None,
|
||||
include_stopped: request.include_stopped,
|
||||
})
|
||||
fn list_workers(&self, _request: WorkerListRequest) -> Result<WorkerList, TargetError> {
|
||||
Err(TargetError::unsupported("Worker listing", self.kind()))
|
||||
}
|
||||
|
||||
fn connect_worker(
|
||||
@@ -261,6 +214,59 @@ impl Target for LocalTarget {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct StandaloneTarget {
|
||||
state_dir: PathBuf,
|
||||
}
|
||||
|
||||
impl StandaloneTarget {
|
||||
#[must_use]
|
||||
pub fn new(state_dir: impl Into<PathBuf>) -> Self {
|
||||
Self {
|
||||
state_dir: state_dir.into(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Target for StandaloneTarget {
|
||||
fn kind(&self) -> TargetKind {
|
||||
TargetKind::Standalone
|
||||
}
|
||||
|
||||
fn resolve(&self) -> Result<ResolvedTarget, TargetError> {
|
||||
Ok(ResolvedTarget::Standalone)
|
||||
}
|
||||
|
||||
fn spawn_worker(&self) -> Result<WorkerSpawn, TargetError> {
|
||||
Ok(WorkerSpawn {
|
||||
state_dir: self.state_dir.clone(),
|
||||
})
|
||||
}
|
||||
|
||||
fn standalone_worker_list(
|
||||
&self,
|
||||
include_all: bool,
|
||||
) -> Result<StandaloneWorkerListIntent, TargetError> {
|
||||
let cwd = std::env::current_dir()
|
||||
.map_err(|error| TargetError::invalid(self.kind(), error.to_string()))?;
|
||||
Ok(StandaloneWorkerListIntent {
|
||||
state_dir: self.state_dir.clone(),
|
||||
cwd,
|
||||
include_all,
|
||||
})
|
||||
}
|
||||
|
||||
fn standalone_worker_resume(
|
||||
&self,
|
||||
worker_id: String,
|
||||
) -> Result<StandaloneWorkerResumeIntent, TargetError> {
|
||||
Ok(StandaloneWorkerResumeIntent {
|
||||
state_dir: self.state_dir.clone(),
|
||||
worker_id,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl Target for BackendTarget {
|
||||
fn kind(&self) -> TargetKind {
|
||||
TargetKind::Backend
|
||||
@@ -279,42 +285,27 @@ impl Target for BackendTarget {
|
||||
})
|
||||
}
|
||||
|
||||
fn spawn_worker(&self) -> Result<WorkerSpawn, TargetError> {
|
||||
Err(TargetError::unsupported("Worker spawn", self.kind()))
|
||||
}
|
||||
|
||||
fn worker_by_name(&self) -> Result<WorkerByName, TargetError> {
|
||||
Err(TargetError::unsupported(
|
||||
"Worker name attachment",
|
||||
self.kind(),
|
||||
))
|
||||
}
|
||||
|
||||
fn resume_worker(&self) -> Result<WorkerResume, TargetError> {
|
||||
Err(TargetError::unsupported("Worker resume", self.kind()))
|
||||
}
|
||||
|
||||
fn dashboard(&self) -> Result<Dashboard, TargetError> {
|
||||
match self.resolve()? {
|
||||
ResolvedTarget::Backend {
|
||||
base_url,
|
||||
workspace_id,
|
||||
} => Ok(Dashboard::Backend {
|
||||
base_url,
|
||||
workspace_id,
|
||||
}),
|
||||
ResolvedTarget::Local => unreachable!("BackendTarget cannot resolve as Local"),
|
||||
}
|
||||
let ResolvedTarget::Backend {
|
||||
base_url,
|
||||
workspace_id,
|
||||
} = self.resolve()?
|
||||
else {
|
||||
unreachable!("BackendTarget resolves only Backend authority")
|
||||
};
|
||||
Ok(Dashboard {
|
||||
base_url,
|
||||
workspace_id,
|
||||
})
|
||||
}
|
||||
|
||||
fn list_workers(&self, request: WorkerListRequest) -> Result<WorkerList, TargetError> {
|
||||
Ok(WorkerList {
|
||||
local_runtime_command: None,
|
||||
backend_target: Some(BackendRuntimeListTarget::new(
|
||||
backend_target: BackendRuntimeListTarget::new(
|
||||
self.base_url.clone(),
|
||||
self.workspace_id.clone(),
|
||||
request.runtime_id,
|
||||
)),
|
||||
),
|
||||
include_stopped: request.include_stopped,
|
||||
})
|
||||
}
|
||||
@@ -371,8 +362,34 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn local_target_resolves_local_product_state_authority() {
|
||||
assert_eq!(LocalTarget::new().resolve().unwrap(), ResolvedTarget::Local);
|
||||
fn standalone_target_carries_in_process_state_without_runtime_command() {
|
||||
let target = StandaloneTarget::new("/tmp/yoi-standalone-state");
|
||||
|
||||
assert_eq!(target.kind(), TargetKind::Standalone);
|
||||
assert_eq!(target.resolve().unwrap(), ResolvedTarget::Standalone);
|
||||
assert_eq!(
|
||||
target.spawn_worker().unwrap(),
|
||||
WorkerSpawn {
|
||||
state_dir: PathBuf::from("/tmp/yoi-standalone-state"),
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn standalone_target_never_exposes_workspace_worker_operations() {
|
||||
let target = StandaloneTarget::new("/tmp/yoi-standalone-state");
|
||||
|
||||
assert_eq!(
|
||||
target
|
||||
.list_workers(WorkerListRequest::new(None))
|
||||
.unwrap_err()
|
||||
.to_string(),
|
||||
"Worker listing is not supported by Standalone target"
|
||||
);
|
||||
assert_eq!(
|
||||
target.dashboard().unwrap_err().to_string(),
|
||||
"Worker dashboard is not supported by Standalone target"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -381,26 +398,13 @@ mod tests {
|
||||
|
||||
assert_eq!(
|
||||
target.dashboard().unwrap(),
|
||||
Dashboard::Backend {
|
||||
Dashboard {
|
||||
base_url: "http://127.0.0.1:8787".to_string(),
|
||||
workspace_id: "workspace-a".to_string(),
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backend_target_rejects_dashboard_without_workspace_selection() {
|
||||
let target = BackendTarget::new("http://127.0.0.1:8787", None::<String>);
|
||||
|
||||
assert!(
|
||||
target
|
||||
.dashboard()
|
||||
.unwrap_err()
|
||||
.to_string()
|
||||
.contains("workspace selection is required")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backend_target_builds_worker_list() {
|
||||
let target = BackendTarget::new("http://127.0.0.1:8787", Some("workspace-a"));
|
||||
@@ -408,26 +412,13 @@ mod tests {
|
||||
.list_workers(WorkerListRequest::new(Some("runtime-a".to_string())))
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(workers.backend_target.base_url, "http://127.0.0.1:8787");
|
||||
assert_eq!(
|
||||
workers.backend_target.as_ref().unwrap().base_url,
|
||||
"http://127.0.0.1:8787"
|
||||
);
|
||||
assert_eq!(
|
||||
workers
|
||||
.backend_target
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.workspace_id
|
||||
.as_deref(),
|
||||
workers.backend_target.workspace_id.as_deref(),
|
||||
Some("workspace-a")
|
||||
);
|
||||
assert_eq!(
|
||||
workers
|
||||
.backend_target
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.runtime_id
|
||||
.as_deref(),
|
||||
workers.backend_target.runtime_id.as_deref(),
|
||||
Some("runtime-a")
|
||||
);
|
||||
}
|
||||
@@ -446,41 +437,17 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backend_target_rejects_worker_connection_before_workspace_selection() {
|
||||
let target = BackendTarget::new("http://127.0.0.1:8787", None::<String>);
|
||||
let error =
|
||||
match target.connect_worker(WorkerConnectionSelector::new("runtime-a", "worker-b")) {
|
||||
Ok(_) => panic!("unscoped connection must fail"),
|
||||
Err(error) => error,
|
||||
};
|
||||
fn standalone_target_builds_explicit_worker_intents() {
|
||||
let target = StandaloneTarget::new("/tmp/yoi-client-workers");
|
||||
let list = target.standalone_worker_list(true).unwrap();
|
||||
assert_eq!(list.state_dir, PathBuf::from("/tmp/yoi-client-workers"));
|
||||
assert!(list.include_all);
|
||||
assert!(list.cwd.is_absolute());
|
||||
|
||||
assert!(
|
||||
error
|
||||
.to_string()
|
||||
.contains("workspace selection is required")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backend_target_rejects_local_worker_operations() {
|
||||
let target = BackendTarget::new("http://127.0.0.1:8787", None::<String>);
|
||||
let err = target.spawn_worker().unwrap_err();
|
||||
|
||||
assert_eq!(
|
||||
err.to_string(),
|
||||
"Worker spawn is not supported by Backend target"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn local_target_builds_local_worker_list() {
|
||||
let target = LocalTarget::new();
|
||||
let workers = target
|
||||
.list_workers(WorkerListRequest::with_stopped(None))
|
||||
let resume = target
|
||||
.standalone_worker_resume("019d1234-0000-7000-8000-000000000000".to_string())
|
||||
.unwrap();
|
||||
|
||||
assert!(workers.local_runtime_command.is_some());
|
||||
assert!(workers.backend_target.is_none());
|
||||
assert!(workers.include_stopped);
|
||||
assert_eq!(resume.state_dir, list.state_dir);
|
||||
assert_eq!(resume.worker_id, "019d1234-0000-7000-8000-000000000000");
|
||||
}
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,115 @@
|
||||
use async_trait::async_trait;
|
||||
use thiserror::Error;
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use super::Socket as SocketContract;
|
||||
|
||||
const CHANNEL_CAPACITY: usize = 256;
|
||||
|
||||
pub struct Socket {
|
||||
outgoing: mpsc::Sender<String>,
|
||||
incoming: mpsc::Receiver<String>,
|
||||
}
|
||||
|
||||
/// Host-side endpoint paired with an in-process client transport.
|
||||
pub struct Peer {
|
||||
incoming: mpsc::Receiver<String>,
|
||||
outgoing: mpsc::Sender<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
|
||||
pub enum SocketError {
|
||||
#[error("in-process Worker protocol transport closed")]
|
||||
Closed,
|
||||
}
|
||||
|
||||
impl Socket {
|
||||
pub fn pair() -> (Self, Peer) {
|
||||
let (client_tx, peer_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
||||
let (peer_tx, client_rx) = mpsc::channel(CHANNEL_CAPACITY);
|
||||
(
|
||||
Self {
|
||||
outgoing: client_tx,
|
||||
incoming: client_rx,
|
||||
},
|
||||
Peer {
|
||||
incoming: peer_rx,
|
||||
outgoing: peer_tx,
|
||||
},
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl SocketContract for Socket {
|
||||
type Error = SocketError;
|
||||
|
||||
async fn send(&mut self, message: String) -> Result<(), Self::Error> {
|
||||
self.outgoing
|
||||
.send(message)
|
||||
.await
|
||||
.map_err(|_| SocketError::Closed)
|
||||
}
|
||||
|
||||
async fn next(&mut self) -> Result<Option<String>, Self::Error> {
|
||||
Ok(self.incoming.recv().await)
|
||||
}
|
||||
|
||||
fn try_next(&mut self) -> Result<Option<String>, Self::Error> {
|
||||
match self.incoming.try_recv() {
|
||||
Ok(message) => Ok(Some(message)),
|
||||
Err(mpsc::error::TryRecvError::Empty | mpsc::error::TryRecvError::Disconnected) => {
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Peer {
|
||||
pub async fn next(&mut self) -> Option<String> {
|
||||
self.incoming.recv().await
|
||||
}
|
||||
|
||||
pub async fn send(&self, message: String) -> Result<(), String> {
|
||||
self.outgoing.send(message).await.map_err(|error| error.0)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use protocol::stream::{decode_method, encode_event};
|
||||
use protocol::{Event, Method, WorkerStatus};
|
||||
|
||||
use super::Socket;
|
||||
use crate::Client;
|
||||
|
||||
#[tokio::test]
|
||||
async fn pair_carries_typed_protocol_through_generic_client() {
|
||||
let (socket, mut peer) = Socket::pair();
|
||||
let mut client = Client::new(socket);
|
||||
|
||||
client
|
||||
.send(&Method::run_text("hello"))
|
||||
.await
|
||||
.expect("send method");
|
||||
assert!(matches!(
|
||||
peer.next().await.as_deref().map(decode_method),
|
||||
Some(Ok(Method::Run { .. }))
|
||||
));
|
||||
|
||||
peer.send(
|
||||
encode_event(&Event::Status {
|
||||
status: WorkerStatus::Idle,
|
||||
})
|
||||
.expect("encode event"),
|
||||
)
|
||||
.await
|
||||
.expect("send event");
|
||||
assert!(matches!(
|
||||
client.next_event().await,
|
||||
Ok(Some(Event::Status {
|
||||
status: WorkerStatus::Idle
|
||||
}))
|
||||
));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
use std::error::Error;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
pub mod in_process;
|
||||
pub mod unix_socket;
|
||||
pub mod websocket;
|
||||
|
||||
/// Message-oriented transport for one Worker protocol connection.
|
||||
///
|
||||
/// Implementations own physical framing. `client::Client` owns the typed
|
||||
/// Method/Event protocol encoding layered on top of these UTF-8 messages.
|
||||
#[async_trait]
|
||||
pub trait Socket {
|
||||
type Error: Error + Send + Sync + 'static;
|
||||
|
||||
async fn send(&mut self, message: String) -> Result<(), Self::Error>;
|
||||
|
||||
async fn next(&mut self) -> Result<Option<String>, Self::Error>;
|
||||
|
||||
fn try_next(&mut self) -> Result<Option<String>, Self::Error>;
|
||||
}
|
||||
@@ -0,0 +1,172 @@
|
||||
use std::io;
|
||||
use std::path::Path;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
|
||||
use tokio::net::UnixStream;
|
||||
use tokio::sync::mpsc;
|
||||
use tokio::task::JoinHandle;
|
||||
|
||||
use super::Socket as SocketContract;
|
||||
|
||||
pub struct Socket {
|
||||
writer: tokio::io::WriteHalf<UnixStream>,
|
||||
messages: mpsc::Receiver<io::Result<String>>,
|
||||
reader_task: JoinHandle<()>,
|
||||
}
|
||||
|
||||
impl Socket {
|
||||
pub async fn connect(path: &Path) -> io::Result<Self> {
|
||||
let stream = UnixStream::connect(path).await?;
|
||||
let (reader, writer) = tokio::io::split(stream);
|
||||
let (message_tx, messages) = mpsc::channel(256);
|
||||
let reader_task = tokio::spawn(async move {
|
||||
let mut lines = BufReader::new(reader).lines();
|
||||
loop {
|
||||
match lines.next_line().await {
|
||||
Ok(Some(message)) if message.trim().is_empty() => {}
|
||||
Ok(Some(message)) => {
|
||||
if message_tx.send(Ok(message)).await.is_err() {
|
||||
return;
|
||||
}
|
||||
}
|
||||
Ok(None) => return,
|
||||
Err(error) => {
|
||||
let _ = message_tx.send(Err(error)).await;
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
Ok(Self {
|
||||
writer,
|
||||
messages,
|
||||
reader_task,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl SocketContract for Socket {
|
||||
type Error = io::Error;
|
||||
|
||||
async fn send(&mut self, message: String) -> Result<(), Self::Error> {
|
||||
self.writer.write_all(message.as_bytes()).await?;
|
||||
self.writer.write_all(b"\n").await?;
|
||||
self.writer.flush().await
|
||||
}
|
||||
|
||||
async fn next(&mut self) -> Result<Option<String>, Self::Error> {
|
||||
match self.messages.recv().await {
|
||||
Some(message) => message.map(Some),
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
fn try_next(&mut self) -> Result<Option<String>, Self::Error> {
|
||||
match self.messages.try_recv() {
|
||||
Ok(message) => message.map(Some),
|
||||
Err(mpsc::error::TryRecvError::Empty | mpsc::error::TryRecvError::Disconnected) => {
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for Socket {
|
||||
fn drop(&mut self) {
|
||||
self.reader_task.abort();
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::io::ErrorKind;
|
||||
use std::time::Duration;
|
||||
|
||||
use protocol::stream::{decode_method, encode_event};
|
||||
use protocol::{Event, Method, WorkerStatus};
|
||||
use tempfile::tempdir;
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::UnixListener;
|
||||
|
||||
use super::*;
|
||||
use crate::Client;
|
||||
|
||||
async fn assert_peer_closed(stream: &mut UnixStream, reason: &str) {
|
||||
let mut buf = [0_u8; 1];
|
||||
match tokio::time::timeout(Duration::from_secs(1), stream.read(&mut buf))
|
||||
.await
|
||||
.expect(reason)
|
||||
{
|
||||
Ok(0) => {}
|
||||
Err(error) if error.kind() == ErrorKind::ConnectionReset => {}
|
||||
Ok(n) => panic!("server should observe peer close, read {n} byte(s)"),
|
||||
Err(error) => panic!("server read failed unexpectedly: {error}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn client_receives_events_over_unix_socket() {
|
||||
let socket_dir = tempdir().unwrap();
|
||||
let socket_path = socket_dir.path().join("events.sock");
|
||||
let listener = UnixListener::bind(&socket_path).unwrap();
|
||||
let server = tokio::spawn(async move {
|
||||
let (mut stream, _) = listener.accept().await.unwrap();
|
||||
let event = encode_event(&Event::Status {
|
||||
status: WorkerStatus::Idle,
|
||||
})
|
||||
.unwrap();
|
||||
stream.write_all(event.as_bytes()).await.unwrap();
|
||||
stream.write_all(b"\n").await.unwrap();
|
||||
});
|
||||
|
||||
let mut client = Client::new(Socket::connect(&socket_path).await.unwrap());
|
||||
let event = tokio::time::timeout(Duration::from_secs(1), client.next_event())
|
||||
.await
|
||||
.expect("client should receive event while alive")
|
||||
.expect("transport should succeed");
|
||||
assert!(matches!(
|
||||
event,
|
||||
Some(Event::Status {
|
||||
status: WorkerStatus::Idle
|
||||
})
|
||||
));
|
||||
server.await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn client_sends_methods_over_unix_socket() {
|
||||
let socket_dir = tempdir().unwrap();
|
||||
let socket_path = socket_dir.path().join("send.sock");
|
||||
let listener = UnixListener::bind(&socket_path).unwrap();
|
||||
let server = tokio::spawn(async move {
|
||||
let (reader, _) = listener.accept().await.unwrap();
|
||||
BufReader::new(reader).lines().next_line().await.unwrap()
|
||||
});
|
||||
|
||||
let mut client = Client::new(Socket::connect(&socket_path).await.unwrap());
|
||||
client
|
||||
.send(&Method::run_text("hello"))
|
||||
.await
|
||||
.expect("send method");
|
||||
|
||||
let received = server.await.unwrap().expect("method message");
|
||||
assert!(matches!(decode_method(&received), Ok(Method::Run { .. })));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn dropping_socket_closes_server_connection() {
|
||||
let socket_dir = tempdir().unwrap();
|
||||
let socket_path = socket_dir.path().join("drop.sock");
|
||||
let listener = UnixListener::bind(&socket_path).unwrap();
|
||||
let server = tokio::spawn(async move {
|
||||
let (mut stream, _) = listener.accept().await.unwrap();
|
||||
assert_peer_closed(&mut stream, "dropped socket should close promptly").await;
|
||||
});
|
||||
|
||||
let socket = Socket::connect(&socket_path).await.unwrap();
|
||||
drop(socket);
|
||||
server.await.unwrap();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,140 @@
|
||||
use async_trait::async_trait;
|
||||
use futures::{SinkExt, StreamExt};
|
||||
use thiserror::Error;
|
||||
use tokio::net::TcpStream;
|
||||
use tokio::sync::mpsc;
|
||||
use tokio::task::JoinHandle;
|
||||
use tokio_tungstenite::tungstenite::http::Request;
|
||||
use tokio_tungstenite::tungstenite::{self, Message};
|
||||
use tokio_tungstenite::{MaybeTlsStream, WebSocketStream, connect_async};
|
||||
|
||||
use super::Socket as SocketContract;
|
||||
|
||||
type Writer = futures::stream::SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, Message>;
|
||||
|
||||
pub struct Socket {
|
||||
writer: Writer,
|
||||
messages: mpsc::Receiver<Result<String, SocketError>>,
|
||||
reader_task: JoinHandle<()>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum SocketError {
|
||||
#[error("WebSocket transport failed: {0}")]
|
||||
WebSocket(#[from] tungstenite::Error),
|
||||
}
|
||||
|
||||
impl Socket {
|
||||
pub async fn connect(request: Request<()>) -> Result<Self, SocketError> {
|
||||
let (stream, _) = connect_async(request).await?;
|
||||
let (writer, mut reader) = stream.split();
|
||||
let (message_tx, messages) = mpsc::channel(256);
|
||||
let reader_task = tokio::spawn(async move {
|
||||
loop {
|
||||
match reader.next().await {
|
||||
Some(Ok(Message::Text(message))) => {
|
||||
if message_tx.send(Ok(message.to_string())).await.is_err() {
|
||||
return;
|
||||
}
|
||||
}
|
||||
Some(Ok(Message::Close(_))) | None => return,
|
||||
Some(Ok(
|
||||
Message::Binary(_)
|
||||
| Message::Ping(_)
|
||||
| Message::Pong(_)
|
||||
| Message::Frame(_),
|
||||
)) => {}
|
||||
Some(Err(error)) => {
|
||||
let _ = message_tx.send(Err(SocketError::WebSocket(error))).await;
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
Ok(Self {
|
||||
writer,
|
||||
messages,
|
||||
reader_task,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl SocketContract for Socket {
|
||||
type Error = SocketError;
|
||||
|
||||
async fn send(&mut self, message: String) -> Result<(), Self::Error> {
|
||||
self.writer.send(Message::Text(message.into())).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn next(&mut self) -> Result<Option<String>, Self::Error> {
|
||||
match self.messages.recv().await {
|
||||
Some(message) => message.map(Some),
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
fn try_next(&mut self) -> Result<Option<String>, Self::Error> {
|
||||
match self.messages.try_recv() {
|
||||
Ok(message) => message.map(Some),
|
||||
Err(mpsc::error::TryRecvError::Empty | mpsc::error::TryRecvError::Disconnected) => {
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for Socket {
|
||||
fn drop(&mut self) {
|
||||
self.reader_task.abort();
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use futures::{SinkExt, StreamExt};
|
||||
use protocol::stream::{decode_method, encode_event};
|
||||
use protocol::{Event, Method, WorkerStatus};
|
||||
use tokio::net::TcpListener;
|
||||
use tokio_tungstenite::accept_async;
|
||||
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
|
||||
|
||||
use super::*;
|
||||
use crate::Client;
|
||||
|
||||
#[tokio::test]
|
||||
async fn carries_typed_protocol_through_generic_client() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let server = tokio::spawn(async move {
|
||||
let (stream, _) = listener.accept().await.unwrap();
|
||||
let mut socket = accept_async(stream).await.unwrap();
|
||||
let message = socket.next().await.unwrap().unwrap();
|
||||
assert!(matches!(
|
||||
message,
|
||||
Message::Text(ref text)
|
||||
if matches!(decode_method(text), Ok(Method::Run { .. }))
|
||||
));
|
||||
let event = encode_event(&Event::Status {
|
||||
status: WorkerStatus::Idle,
|
||||
})
|
||||
.unwrap();
|
||||
socket.send(Message::Text(event.into())).await.unwrap();
|
||||
});
|
||||
|
||||
let request = format!("ws://{address}").into_client_request().unwrap();
|
||||
let mut client = Client::new(Socket::connect(request).await.unwrap());
|
||||
client
|
||||
.send(&Method::run_text("hello"))
|
||||
.await
|
||||
.expect("send method");
|
||||
assert!(matches!(
|
||||
client.next_event().await,
|
||||
Ok(Some(Event::Status {
|
||||
status: WorkerStatus::Idle
|
||||
}))
|
||||
));
|
||||
server.await.unwrap();
|
||||
}
|
||||
}
|
||||
@@ -1,186 +0,0 @@
|
||||
use std::io;
|
||||
use std::path::Path;
|
||||
|
||||
use protocol::stream::{JsonLineReader, JsonLineWriter};
|
||||
use protocol::{Event, Method};
|
||||
use tokio::net::UnixStream;
|
||||
use tokio::sync::mpsc;
|
||||
use tokio::task::JoinHandle;
|
||||
|
||||
pub struct WorkerClient {
|
||||
writer: JsonLineWriter<tokio::io::WriteHalf<UnixStream>>,
|
||||
event_rx: mpsc::Receiver<Event>,
|
||||
reader_task: JoinHandle<()>,
|
||||
}
|
||||
|
||||
impl WorkerClient {
|
||||
pub async fn connect(path: &Path) -> Result<Self, io::Error> {
|
||||
let stream = UnixStream::connect(path).await?;
|
||||
let (reader, writer) = tokio::io::split(stream);
|
||||
let writer = JsonLineWriter::new(writer);
|
||||
|
||||
let (event_tx, event_rx) = mpsc::channel::<Event>(256);
|
||||
|
||||
let reader_task = tokio::spawn(async move {
|
||||
let mut reader = JsonLineReader::new(reader);
|
||||
while let Ok(Some(event)) = reader.next::<Event>().await {
|
||||
if event_tx.send(event).await.is_err() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
Ok(Self {
|
||||
writer,
|
||||
event_rx,
|
||||
reader_task,
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn send(&mut self, method: &Method) -> Result<(), io::Error> {
|
||||
self.writer.write(method).await
|
||||
}
|
||||
|
||||
pub fn try_next_event(&mut self) -> Option<Event> {
|
||||
self.event_rx.try_recv().ok()
|
||||
}
|
||||
|
||||
pub async fn next_event(&mut self) -> Option<Event> {
|
||||
self.event_rx.recv().await
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for WorkerClient {
|
||||
fn drop(&mut self) {
|
||||
self.reader_task.abort();
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::io::ErrorKind;
|
||||
use std::time::Duration;
|
||||
|
||||
use protocol::{Segment, WorkerStatus};
|
||||
use tempfile::tempdir;
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::UnixListener;
|
||||
|
||||
use super::*;
|
||||
|
||||
async fn assert_peer_closed(stream: &mut UnixStream, reason: &str) {
|
||||
let mut buf = [0_u8; 1];
|
||||
match tokio::time::timeout(Duration::from_secs(1), stream.read(&mut buf))
|
||||
.await
|
||||
.expect(reason)
|
||||
{
|
||||
Ok(0) => {}
|
||||
Err(error) if error.kind() == ErrorKind::ConnectionReset => {}
|
||||
Ok(n) => panic!("server should observe peer close, read {n} byte(s)"),
|
||||
Err(error) => panic!("server read failed unexpectedly: {error}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn receives_events_while_client_is_alive() {
|
||||
let socket_dir = tempdir().unwrap();
|
||||
let socket_path = socket_dir.path().join("events.sock");
|
||||
let listener = UnixListener::bind(&socket_path).unwrap();
|
||||
let server = tokio::spawn(async move {
|
||||
let (stream, _) = listener.accept().await.unwrap();
|
||||
let mut writer = JsonLineWriter::new(stream);
|
||||
writer
|
||||
.write(&Event::Status {
|
||||
status: WorkerStatus::Idle,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
});
|
||||
|
||||
let mut client = WorkerClient::connect(&socket_path).await.unwrap();
|
||||
|
||||
let event = tokio::time::timeout(Duration::from_secs(1), client.next_event())
|
||||
.await
|
||||
.expect("client should receive event while alive");
|
||||
assert!(matches!(
|
||||
event,
|
||||
Some(Event::Status {
|
||||
status: WorkerStatus::Idle
|
||||
})
|
||||
));
|
||||
server.await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn send_writes_methods_while_client_is_alive() {
|
||||
let socket_dir = tempdir().unwrap();
|
||||
let socket_path = socket_dir.path().join("send.sock");
|
||||
let listener = UnixListener::bind(&socket_path).unwrap();
|
||||
let server = tokio::spawn(async move {
|
||||
let (stream, _) = listener.accept().await.unwrap();
|
||||
let mut reader = JsonLineReader::new(stream);
|
||||
reader.next::<Method>().await.unwrap()
|
||||
});
|
||||
|
||||
let mut client = WorkerClient::connect(&socket_path).await.unwrap();
|
||||
let method = Method::Run {
|
||||
input: vec![Segment::text("hello")],
|
||||
};
|
||||
client.send(&method).await.unwrap();
|
||||
|
||||
let received = tokio::time::timeout(Duration::from_secs(1), server)
|
||||
.await
|
||||
.expect("server should receive method while client is alive")
|
||||
.unwrap();
|
||||
match received {
|
||||
Some(Method::Run { input }) => assert_eq!(input, vec![Segment::text("hello")]),
|
||||
other => panic!("expected Run method, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn dropping_repeated_clients_closes_server_connections() {
|
||||
let socket_dir = tempdir().unwrap();
|
||||
let socket_path = socket_dir.path().join("drop.sock");
|
||||
let listener = UnixListener::bind(&socket_path).unwrap();
|
||||
let server = tokio::spawn(async move {
|
||||
for _ in 0..16 {
|
||||
let (mut stream, _) = listener.accept().await.unwrap();
|
||||
assert_peer_closed(
|
||||
&mut stream,
|
||||
"dropped client should close its socket promptly",
|
||||
)
|
||||
.await;
|
||||
}
|
||||
});
|
||||
|
||||
for _ in 0..16 {
|
||||
let client = WorkerClient::connect(&socket_path).await.unwrap();
|
||||
drop(client);
|
||||
}
|
||||
|
||||
server.await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn dropping_client_aborts_blocked_reader_task() {
|
||||
let socket_dir = tempdir().unwrap();
|
||||
let socket_path = socket_dir.path().join("blocked-reader.sock");
|
||||
let listener = UnixListener::bind(&socket_path).unwrap();
|
||||
let server = tokio::spawn(async move {
|
||||
let (mut stream, _) = listener.accept().await.unwrap();
|
||||
stream.write_all(b"{\"event\"").await.unwrap();
|
||||
assert_peer_closed(
|
||||
&mut stream,
|
||||
"aborting the blocked client reader should close the socket",
|
||||
)
|
||||
.await;
|
||||
});
|
||||
|
||||
let client = WorkerClient::connect(&socket_path).await.unwrap();
|
||||
tokio::task::yield_now().await;
|
||||
drop(client);
|
||||
|
||||
server.await.unwrap();
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,6 @@
|
||||
use reqwest::Method;
|
||||
use serde::Serialize;
|
||||
use serde::de::DeserializeOwned;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use ticket::{
|
||||
MarkdownText, NewOrchestrationPlanRecord, NewTicket, NewTicketEvent, NewTicketRelation,
|
||||
OrchestrationPlanKind, OrchestrationPlanRecord, Ticket, TicketBackend, TicketDependencyCheck,
|
||||
@@ -9,47 +9,25 @@ use ticket::{
|
||||
TicketRelationKind, TicketRelationView, TicketStateChange, TicketStateSelector, TicketSummary,
|
||||
};
|
||||
use workspace_api::{
|
||||
ListResponse, ObjectiveCreateRequest, ObjectiveDetail, ObjectiveEditRequest,
|
||||
ObjectiveLinkTicketRequest, ObjectiveStateRequest, ObjectiveSummary,
|
||||
TICKET_ORCHESTRATION_PLANS_QUERY_PATH, TICKET_RELATIONS_QUERY_PATH,
|
||||
BrowserCreateWorkerResponse, BrowserWorkspaceOrchestratorResponse,
|
||||
CreateWorkspaceWorkerRequest, ListResponse, MemoryDocumentResponse, MemoryStagingListResponse,
|
||||
ObjectiveCreateRequest, ObjectiveDetail, ObjectiveEditRequest, ObjectiveLinkTicketRequest,
|
||||
ObjectiveStateRequest, ObjectiveSummary, TICKET_ORCHESTRATION_PLANS_QUERY_PATH,
|
||||
TICKET_RELATIONS_QUERY_PATH, WorkerLaunchOptionsResponse,
|
||||
};
|
||||
|
||||
use crate::BackendWorkspaceClientError;
|
||||
use crate::{BackendApiClient, BackendWorkspaceClientError};
|
||||
|
||||
const DEFAULT_PRODUCT_LIST_LIMIT: usize = 1_000;
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct BackendWorkerLaunchOptions {
|
||||
runtimes: Vec<BackendWorkerLaunchRuntime>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct BackendWorkerLaunchRuntime {
|
||||
runtime_id: String,
|
||||
worker_creation_available: bool,
|
||||
working_directory_required: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct BackendCreateWorkerResponse {
|
||||
runtime_id: String,
|
||||
worker_id: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct BackendWorkspaceOrchestratorResponse {
|
||||
disposition: String,
|
||||
worker: Option<BackendCreateWorkerResponse>,
|
||||
}
|
||||
|
||||
/// Workspace-scoped Backend client for Ticket and Objective product state.
|
||||
///
|
||||
/// Construction requires both the selected Backend URL and Workspace identity.
|
||||
/// Callers should derive these once from `Target::resolve()` and must not retry
|
||||
/// failed requests against repository-local state.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct BackendWorkspaceProductClient {
|
||||
base_url: String,
|
||||
api: BackendApiClient,
|
||||
workspace_id: String,
|
||||
}
|
||||
|
||||
@@ -58,22 +36,32 @@ impl BackendWorkspaceProductClient {
|
||||
base_url: impl Into<String>,
|
||||
workspace_id: impl Into<String>,
|
||||
) -> Result<Self, BackendWorkspaceClientError> {
|
||||
let base_url = base_url.into().trim_end_matches('/').to_string();
|
||||
if base_url.is_empty() {
|
||||
return Err(BackendWorkspaceClientError::InvalidTarget(
|
||||
"Backend base URL must not be empty".into(),
|
||||
));
|
||||
}
|
||||
let base_url = base_url.into();
|
||||
let api = BackendApiClient::from_stored_token(&base_url)?;
|
||||
let workspace_id = workspace_id.into();
|
||||
if workspace_id.trim().is_empty() {
|
||||
return Err(BackendWorkspaceClientError::InvalidTarget(
|
||||
"Backend Workspace identity must not be empty".into(),
|
||||
));
|
||||
}
|
||||
Ok(Self {
|
||||
base_url,
|
||||
workspace_id,
|
||||
})
|
||||
Ok(Self { api, workspace_id })
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
fn new_with_access_token(
|
||||
base_url: impl Into<String>,
|
||||
workspace_id: impl Into<String>,
|
||||
access_token: &str,
|
||||
) -> Result<Self, BackendWorkspaceClientError> {
|
||||
let base_url = base_url.into();
|
||||
let api = BackendApiClient::from_access_token_for_test(&base_url, access_token)?;
|
||||
let workspace_id = workspace_id.into();
|
||||
if workspace_id.trim().is_empty() {
|
||||
return Err(BackendWorkspaceClientError::InvalidTarget(
|
||||
"Backend Workspace identity must not be empty".into(),
|
||||
));
|
||||
}
|
||||
Ok(Self { api, workspace_id })
|
||||
}
|
||||
|
||||
pub fn workspace_id(&self) -> &str {
|
||||
@@ -253,11 +241,22 @@ impl BackendWorkspaceProductClient {
|
||||
)
|
||||
}
|
||||
|
||||
pub fn memory_document(&self) -> Result<MemoryDocumentResponse, BackendWorkspaceClientError> {
|
||||
self.get_json("/memory")
|
||||
}
|
||||
|
||||
pub fn list_memory_staging(
|
||||
&self,
|
||||
limit: usize,
|
||||
) -> Result<MemoryStagingListResponse, BackendWorkspaceClientError> {
|
||||
self.get_json(&format!("/memory/staging?limit={limit}"))
|
||||
}
|
||||
|
||||
pub fn launch_ticket_intake(
|
||||
&self,
|
||||
ticket_id: &str,
|
||||
) -> Result<String, BackendWorkspaceClientError> {
|
||||
let options: BackendWorkerLaunchOptions = self.get_json("/workers/launch-options")?;
|
||||
let options: WorkerLaunchOptionsResponse = self.get_json("/workers/launch-options")?;
|
||||
let runtime = options
|
||||
.runtimes
|
||||
.iter()
|
||||
@@ -268,19 +267,19 @@ impl BackendWorkspaceProductClient {
|
||||
.to_string(),
|
||||
)
|
||||
})?;
|
||||
let response: BackendCreateWorkerResponse = self.send_json(
|
||||
Method::POST,
|
||||
"/workers",
|
||||
Some(&serde_json::json!({
|
||||
"runtime_id": runtime.runtime_id,
|
||||
"display_name": format!("intake-{ticket_id}"),
|
||||
"profile": "builtin:intake",
|
||||
"initial_submit": [{
|
||||
"kind": "text",
|
||||
"content": format!("Please handle intake for Ticket {ticket_id}.")
|
||||
}]
|
||||
})),
|
||||
)?;
|
||||
let request = CreateWorkspaceWorkerRequest {
|
||||
runtime_id: runtime.runtime_id.clone(),
|
||||
display_name: format!("intake-{ticket_id}"),
|
||||
profile: Some("builtin:intake".to_string()),
|
||||
ticket_assignment: None,
|
||||
initial_submit: vec![protocol::Segment::Text {
|
||||
content: format!("Please handle intake for Ticket {ticket_id}."),
|
||||
}],
|
||||
working_directory: None,
|
||||
control_operation_id: None,
|
||||
};
|
||||
let response: BrowserCreateWorkerResponse =
|
||||
self.send_json(Method::POST, "/workers", Some(&request))?;
|
||||
Ok(format!(
|
||||
"Started Intake Worker {}/{} for Ticket {ticket_id}",
|
||||
response.runtime_id, response.worker_id
|
||||
@@ -288,7 +287,7 @@ impl BackendWorkspaceProductClient {
|
||||
}
|
||||
|
||||
pub fn start_workspace_orchestrator(&self) -> Result<String, BackendWorkspaceClientError> {
|
||||
let response: BackendWorkspaceOrchestratorResponse =
|
||||
let response: BrowserWorkspaceOrchestratorResponse =
|
||||
self.send_json::<(), _>(Method::POST, "/orchestrator", None)?;
|
||||
let worker = response.worker.ok_or_else(|| {
|
||||
BackendWorkspaceClientError::InvalidTarget(
|
||||
@@ -316,7 +315,7 @@ impl BackendWorkspaceProductClient {
|
||||
body: Option<&B>,
|
||||
) -> Result<R, BackendWorkspaceClientError> {
|
||||
let response = self.request(method, path, body)?.send()?;
|
||||
let response = ensure_success(response)?;
|
||||
self.api.check_status(response.status())?;
|
||||
response.json().map_err(BackendWorkspaceClientError::Http)
|
||||
}
|
||||
|
||||
@@ -326,7 +325,8 @@ impl BackendWorkspaceProductClient {
|
||||
path: &str,
|
||||
body: Option<&B>,
|
||||
) -> Result<(), BackendWorkspaceClientError> {
|
||||
ensure_success(self.request(method, path, body)?.send()?)?;
|
||||
let response = self.request(method, path, body)?.send()?;
|
||||
self.api.check_status(response.status())?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -336,14 +336,12 @@ impl BackendWorkspaceProductClient {
|
||||
path: &str,
|
||||
body: Option<&B>,
|
||||
) -> Result<reqwest::blocking::RequestBuilder, BackendWorkspaceClientError> {
|
||||
let client = reqwest::blocking::Client::builder().build()?;
|
||||
let url = format!(
|
||||
"{}/api/w/{}/{}",
|
||||
self.base_url,
|
||||
let path = format!(
|
||||
"/api/w/{}/{}",
|
||||
encode_path_segment(&self.workspace_id),
|
||||
path.trim_start_matches('/')
|
||||
);
|
||||
let request = client.request(method, url);
|
||||
let request = self.api.blocking_request(method, &path)?;
|
||||
Ok(match body {
|
||||
Some(body) => request.json(body),
|
||||
None => request,
|
||||
@@ -588,19 +586,6 @@ fn ticket_client_error(error: BackendWorkspaceClientError) -> TicketError {
|
||||
TicketError::Sqlite(format!("Backend request failed: {error}"))
|
||||
}
|
||||
|
||||
fn ensure_success(
|
||||
response: reqwest::blocking::Response,
|
||||
) -> Result<reqwest::blocking::Response, BackendWorkspaceClientError> {
|
||||
if response.status().is_success() {
|
||||
return Ok(response);
|
||||
}
|
||||
let status = response.status().as_u16();
|
||||
let message = response
|
||||
.text()
|
||||
.unwrap_or_else(|_| "Backend request failed".to_string());
|
||||
Err(BackendWorkspaceClientError::RequestFailed { status, message })
|
||||
}
|
||||
|
||||
fn ticket_reference(id: &TicketIdOrSlug) -> String {
|
||||
match id {
|
||||
TicketIdOrSlug::Id(id) => id.to_string(),
|
||||
@@ -695,27 +680,111 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn objective_list_uses_workspace_scoped_backend_route() {
|
||||
let body = r#"{"workspace_id":"workspace-a","limit":1000,"items":[],"source":"sqlite","diagnostics":[]}"#;
|
||||
fn memory_document_uses_shared_workspace_scoped_response() {
|
||||
let body = r##"{"body_md":"# Memory\\n","created_at":"2026-09-01T00:00:00Z","updated_at":"2026-09-02T00:00:00Z","bytes":10,"record_source":"workspace-sqlite"}"##;
|
||||
let (base_url, request, handle) = one_response_server("200 OK", body);
|
||||
let client = BackendWorkspaceProductClient::new(base_url, "workspace-a").unwrap();
|
||||
let client = BackendWorkspaceProductClient::new_with_access_token(
|
||||
base_url,
|
||||
"workspace-a",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let response = client.list_objectives(1_000).unwrap();
|
||||
let response = client.memory_document().unwrap();
|
||||
|
||||
assert!(response.items.is_empty());
|
||||
assert_eq!(response.record_source, "workspace-sqlite");
|
||||
assert!(
|
||||
request
|
||||
.recv()
|
||||
.unwrap()
|
||||
.starts_with("GET /api/w/workspace-a/objectives?limit=1000 ")
|
||||
.starts_with("GET /api/w/workspace-a/memory ")
|
||||
);
|
||||
handle.join().unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn memory_staging_uses_shared_dto_with_typed_origin() {
|
||||
let body = r#"{"limit":10,"returned_count":1,"total_valid_count":1,"invalid_count":0,"truncated":false,"order":"imported_at_desc_candidate_id_asc","record_authority":"sqlite_workspace_authority.memory_staging","items":[{"id":"candidate-1","byte_len":128,"record":{"schema_version":1,"id":"candidate-1","extract_run_id":"run-1","source":{"segment_id":"segment-1","range":[1,2]},"kind":"decision","claim":"Keep typed provenance.","why_useful":"Prevents trust loss.","staleness":null,"evidence":[],"source_refs":[{"session_id":"session-1","segment_id":"segment-1","entry_range":[1,2],"evidence_id":"evidence-1","origin":{"kind":"worker_input","workspace_id":"workspace-a","runtime_id":"runtime-1","worker_id":"worker-1"},"evidence_kind":"worker_session_entry","label":null,"summary":null}]}}],"diagnostics":[]}"#;
|
||||
let (base_url, request, handle) = one_response_server("200 OK", body);
|
||||
let client = BackendWorkspaceProductClient::new_with_access_token(
|
||||
base_url,
|
||||
"workspace-a",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let response = client.list_memory_staging(10).unwrap();
|
||||
|
||||
assert_eq!(
|
||||
response.items[0].record.source_refs[0]
|
||||
.origin
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.kind,
|
||||
workspace_api::MemoryEvidenceOriginKind::WorkerInput
|
||||
);
|
||||
assert!(
|
||||
request
|
||||
.recv()
|
||||
.unwrap()
|
||||
.starts_with("GET /api/w/workspace-a/memory/staging?limit=10 ")
|
||||
);
|
||||
handle.join().unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn memory_staging_rejects_unknown_origin_kind() {
|
||||
let body = r#"{"limit":10,"returned_count":1,"total_valid_count":1,"invalid_count":0,"truncated":false,"order":"order","record_authority":"authority","items":[{"id":"candidate-1","byte_len":1,"record":{"schema_version":1,"id":"candidate-1","extract_run_id":"run-1","source":{"segment_id":"segment-1","range":[1,2]},"kind":"decision","claim":"claim","why_useful":"useful","staleness":null,"evidence":[],"source_refs":[{"session_id":null,"segment_id":null,"entry_range":null,"evidence_id":null,"origin":{"kind":"future_origin"},"evidence_kind":null,"label":null,"summary":null}]}}],"diagnostics":[]}"#;
|
||||
let (base_url, request, handle) = one_response_server("200 OK", body);
|
||||
let client = BackendWorkspaceProductClient::new_with_access_token(
|
||||
base_url,
|
||||
"workspace-a",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let error = client.list_memory_staging(10).unwrap_err();
|
||||
|
||||
assert!(matches!(error, BackendWorkspaceClientError::Http(_)));
|
||||
assert!(
|
||||
request
|
||||
.recv()
|
||||
.unwrap()
|
||||
.starts_with("GET /api/w/workspace-a/memory/staging?limit=10 ")
|
||||
);
|
||||
handle.join().unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn objective_list_uses_workspace_scoped_backend_route() {
|
||||
let body = r#"{"workspace_id":"workspace-a","limit":1000,"items":[],"source":"sqlite","diagnostics":[]}"#;
|
||||
let (base_url, request, handle) = one_response_server("200 OK", body);
|
||||
let client = BackendWorkspaceProductClient::new_with_access_token(
|
||||
base_url,
|
||||
"workspace-a",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let response = client.list_objectives(1_000).unwrap();
|
||||
|
||||
assert!(response.items.is_empty());
|
||||
let request = request.recv().unwrap();
|
||||
assert!(request.starts_with("GET /api/w/workspace-a/objectives?limit=1000 "));
|
||||
assert!(request.contains("authorization: Bearer test-backend-token\r\n"));
|
||||
handle.join().unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backend_mutation_failure_is_returned_without_local_fallback() {
|
||||
let (base_url, request, handle) = one_response_server("403 Forbidden", "denied");
|
||||
let client = BackendWorkspaceProductClient::new(base_url, "workspace-a").unwrap();
|
||||
let (base_url, request, handle) =
|
||||
one_response_server("403 Forbidden", "test-backend-token");
|
||||
let client = BackendWorkspaceProductClient::new_with_access_token(
|
||||
base_url,
|
||||
"workspace-a",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let error = client
|
||||
.create_objective(&ObjectiveCreateRequest {
|
||||
@@ -727,6 +796,7 @@ mod tests {
|
||||
.unwrap_err();
|
||||
|
||||
assert!(error.to_string().contains("403"));
|
||||
assert!(!error.to_string().contains("test-backend-token"));
|
||||
assert!(
|
||||
request
|
||||
.recv()
|
||||
@@ -739,7 +809,12 @@ mod tests {
|
||||
#[test]
|
||||
fn ticket_relation_query_uses_workspace_scoped_backend_route() {
|
||||
let (base_url, request, handle) = one_response_server("200 OK", "[]");
|
||||
let client = BackendWorkspaceProductClient::new(base_url, "workspace-a").unwrap();
|
||||
let client = BackendWorkspaceProductClient::new_with_access_token(
|
||||
base_url,
|
||||
"workspace-a",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let relations = client
|
||||
.query_ticket_relations(
|
||||
@@ -758,7 +833,12 @@ mod tests {
|
||||
#[test]
|
||||
fn orchestration_plan_query_uses_workspace_scoped_backend_route() {
|
||||
let (base_url, request, handle) = one_response_server("200 OK", "[]");
|
||||
let client = BackendWorkspaceProductClient::new(base_url, "workspace-a").unwrap();
|
||||
let client = BackendWorkspaceProductClient::new_with_access_token(
|
||||
base_url,
|
||||
"workspace-a",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let records = TicketBackend::query_orchestration_plan_records(&client, None, None).unwrap();
|
||||
|
||||
@@ -777,14 +857,19 @@ mod tests {
|
||||
let (base_url, requests, handle) = response_sequence_server(vec![
|
||||
(
|
||||
"200 OK",
|
||||
r#"{"runtimes":[{"runtime_id":"embedded","worker_creation_available":true,"working_directory_required":false}]}"#,
|
||||
r#"{"workspace_id":"workspace-a","runtimes":[{"runtime_id":"embedded","display_name":"Embedded","built_in":true,"worker_creation_available":true,"working_directory_required":false,"status":"connected","diagnostics":[]}],"default_profile":null,"profiles":[],"repositories":[],"working_directories":[],"diagnostics":[]}"#,
|
||||
),
|
||||
(
|
||||
"200 OK",
|
||||
r#"{"runtime_id":"embedded","worker_id":"worker-1"}"#,
|
||||
r#"{"workspace_id":"workspace-a","runtime_id":"embedded","worker_id":"worker-1","console_href":"/w/workspace-a/workers/worker-1","worker":{"runtime_id":"embedded","worker_id":"worker-1","host_id":"embedded","display_name":"Intake","label":"worker-1","profile":"builtin:intake","singleton_key":null,"tags":[],"workspace":{"visibility":"workspace","identity":"workspace-a","workspace_id":"workspace-a"},"state":"idle","last_seen_at":null,"pinned":false,"retention_state":"active","implementation":{"kind":"runtime","display_hint":"Runtime Worker"},"capabilities":{"can_stop":true,"can_spawn_followup":false},"diagnostics":[]},"diagnostics":[]}"#,
|
||||
),
|
||||
]);
|
||||
let client = BackendWorkspaceProductClient::new(base_url, "workspace-a").unwrap();
|
||||
let client = BackendWorkspaceProductClient::new_with_access_token(
|
||||
base_url,
|
||||
"workspace-a",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let status = client.launch_ticket_intake("T-1").unwrap();
|
||||
|
||||
@@ -804,9 +889,14 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn workspace_orchestrator_launch_uses_scoped_backend_route() {
|
||||
let body = r#"{"disposition":"created","worker":{"runtime_id":"embedded","worker_id":"worker-2"}}"#;
|
||||
let body = r#"{"workspace_id":"workspace-a","online":true,"disposition":"created","worker":{"runtime_id":"embedded","worker_id":"worker-2","host_id":"embedded","display_name":"Orchestrator","label":"worker-2","profile":"builtin:orchestrator","singleton_key":"workspace-orchestrator","tags":[],"workspace":{"visibility":"workspace","identity":"workspace-a","workspace_id":"workspace-a"},"state":"idle","last_seen_at":null,"pinned":true,"retention_state":"active","implementation":{"kind":"runtime","display_hint":"Runtime Worker"},"capabilities":{"can_stop":true,"can_spawn_followup":false},"diagnostics":[]},"diagnostics":[]}"#;
|
||||
let (base_url, request, handle) = one_response_server("200 OK", body);
|
||||
let client = BackendWorkspaceProductClient::new(base_url, "workspace-a").unwrap();
|
||||
let client = BackendWorkspaceProductClient::new_with_access_token(
|
||||
base_url,
|
||||
"workspace-a",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let status = client.start_workspace_orchestrator().unwrap();
|
||||
|
||||
@@ -822,7 +912,12 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn product_client_requires_workspace_identity() {
|
||||
let error = BackendWorkspaceProductClient::new("http://127.0.0.1:8787", "").unwrap_err();
|
||||
let error = BackendWorkspaceProductClient::new_with_access_token(
|
||||
"http://127.0.0.1:8787",
|
||||
"",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap_err();
|
||||
assert!(error.to_string().contains("Workspace identity"));
|
||||
}
|
||||
|
||||
|
||||
@@ -101,20 +101,24 @@ pub fn complete_current(
|
||||
let utf8_byte_offset = utf16_to_utf8_offset(&source, utf16_offset)?;
|
||||
let result = session_environment(snapshot.clone())
|
||||
.complete_config(&entrypoint, &source, utf8_byte_offset, explicit)
|
||||
.map_err(|error| JsValue::from_str(&format!("{error:?}")))?
|
||||
.map(|result| WasmCompletionResult {
|
||||
from: result.from,
|
||||
items: result
|
||||
.items
|
||||
.into_iter()
|
||||
.map(|item| WasmCompletionItem {
|
||||
label: item.label,
|
||||
kind: format!("{:?}", item.kind).to_lowercase(),
|
||||
detail: item.detail,
|
||||
priority: item.priority,
|
||||
})
|
||||
.collect(),
|
||||
});
|
||||
.map_err(|error| JsValue::from_str(&format!("{error:?}")))?;
|
||||
let result = result
|
||||
.map(|result| {
|
||||
Ok::<WasmCompletionResult, JsValue>(WasmCompletionResult {
|
||||
from: utf8_to_utf16_offset(&source, result.from)?,
|
||||
items: result
|
||||
.items
|
||||
.into_iter()
|
||||
.map(|item| WasmCompletionItem {
|
||||
label: item.label,
|
||||
kind: format!("{:?}", item.kind).to_lowercase(),
|
||||
detail: item.detail,
|
||||
priority: item.priority,
|
||||
})
|
||||
.collect(),
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
encode(result)
|
||||
})
|
||||
}
|
||||
@@ -177,6 +181,16 @@ fn utf16_to_utf8_offset(source: &str, utf16_offset: usize) -> Result<usize, JsVa
|
||||
}
|
||||
}
|
||||
|
||||
fn utf8_to_utf16_offset(source: &str, utf8_offset: usize) -> Result<usize, JsValue> {
|
||||
if utf8_offset > source.len() {
|
||||
return Err(JsValue::from_str("UTF-8 offset is outside the source"));
|
||||
}
|
||||
if !source.is_char_boundary(utf8_offset) {
|
||||
return Err(JsValue::from_str("UTF-8 offset splits a character"));
|
||||
}
|
||||
Ok(source[..utf8_offset].encode_utf16().count())
|
||||
}
|
||||
|
||||
fn decode<T: serde::de::DeserializeOwned>(value: JsValue) -> Result<T, JsValue> {
|
||||
from_value(value).map_err(|error| JsValue::from_str(&error.to_string()))
|
||||
}
|
||||
|
||||
@@ -1203,6 +1203,9 @@ impl SnapshotEnvironment {
|
||||
{
|
||||
let mut member_source = format!("{WORKSPACE_CONFIG_SCHEMA_GLOBAL}.");
|
||||
member_source.push_str(&context.schema_path.join("."));
|
||||
if !context.schema_path.is_empty() && context.from == utf8_byte_offset {
|
||||
member_source.push('.');
|
||||
}
|
||||
let mut completion = LanguageService::new(self).complete(
|
||||
entrypoint.as_str(),
|
||||
&member_source,
|
||||
@@ -1961,6 +1964,31 @@ mod tests {
|
||||
.iter()
|
||||
.any(|item| item.label == "default_profile")
|
||||
);
|
||||
|
||||
let blank_nested_source = "{ profile = { } } as WorkspaceConfigSchema";
|
||||
let blank_nested_cursor = blank_nested_source.find("{ }").unwrap() + 2;
|
||||
let blank_nested = environment
|
||||
.complete_config(
|
||||
&path("main.dcdl"),
|
||||
blank_nested_source,
|
||||
blank_nested_cursor,
|
||||
true,
|
||||
)
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(blank_nested.from, blank_nested_cursor);
|
||||
assert!(
|
||||
blank_nested
|
||||
.items
|
||||
.iter()
|
||||
.any(|item| item.label == "default_profile")
|
||||
);
|
||||
assert!(
|
||||
!blank_nested
|
||||
.items
|
||||
.iter()
|
||||
.any(|item| item.label == "profile")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -157,10 +157,28 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
fn grep_request(path: &str, pattern: &str) -> GrepRequest {
|
||||
GrepRequest {
|
||||
pattern: pattern.to_string(),
|
||||
path: FsPath::new(path).unwrap(),
|
||||
glob: None,
|
||||
file_type: None,
|
||||
case_insensitive: false,
|
||||
before_context: 0,
|
||||
after_context: 0,
|
||||
multiline: false,
|
||||
output_mode: GrepOutputMode::Content,
|
||||
limit: 10,
|
||||
offset: 0,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn logical_paths_reject_absolute_parent_and_backslash_forms() {
|
||||
assert!(FsPath::new("src/lib.rs").is_ok());
|
||||
assert!(FsPath::new("/tmp/file").is_err());
|
||||
assert!(FsPath::new_scoped("/tmp/file").is_ok());
|
||||
assert!(FsPath::new_scoped("/tmp/../secret").is_err());
|
||||
assert!(FsPath::new("../file").is_err());
|
||||
assert!(FsPath::new("src\\lib.rs").is_err());
|
||||
}
|
||||
@@ -279,4 +297,313 @@ mod tests {
|
||||
assert_eq!(grep.matched_files, 2);
|
||||
assert!(!grep.output.contains("c.txt"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grep_accepts_a_direct_file_without_searching_siblings() {
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
let selected = temp.path().join("selected.txt");
|
||||
std::fs::write(&selected, "before\nneedle selected\nafter\n").unwrap();
|
||||
std::fs::write(temp.path().join("sibling.txt"), "needle sibling\n").unwrap();
|
||||
let root = temp.path().canonicalize().unwrap();
|
||||
let readable = RootAccess(root.clone());
|
||||
|
||||
let mut request = grep_request("selected.txt", "needle");
|
||||
request.before_context = 1;
|
||||
request.after_context = 1;
|
||||
let direct = run_grep(&root, selected, request, &readable).unwrap();
|
||||
|
||||
assert_eq!(direct.match_count, 1);
|
||||
assert_eq!(direct.matched_files, 1);
|
||||
assert_eq!(
|
||||
direct.output,
|
||||
concat!(
|
||||
"selected.txt\n",
|
||||
" 1 │ before\n",
|
||||
" > 2 │ needle selected\n",
|
||||
" 3 │ after\n",
|
||||
)
|
||||
);
|
||||
assert!(!direct.output.contains("sibling"));
|
||||
|
||||
let directory = run_grep(
|
||||
&root,
|
||||
root.clone(),
|
||||
GrepRequest {
|
||||
pattern: "needle".to_string(),
|
||||
path: FsPath::root(),
|
||||
glob: None,
|
||||
file_type: None,
|
||||
case_insensitive: false,
|
||||
before_context: 0,
|
||||
after_context: 0,
|
||||
multiline: false,
|
||||
output_mode: GrepOutputMode::Content,
|
||||
limit: 10,
|
||||
offset: 0,
|
||||
},
|
||||
&readable,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(directory.match_count, 2);
|
||||
assert_eq!(directory.matched_files, 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grep_direct_file_applies_glob_and_type_filters_for_every_output_mode() {
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
let nested = temp.path().join("nested");
|
||||
std::fs::create_dir(&nested).unwrap();
|
||||
let selected = nested.join("selected.rs");
|
||||
std::fs::write(&selected, "needle one\nneedle two\n").unwrap();
|
||||
let root = temp.path().canonicalize().unwrap();
|
||||
let readable = RootAccess(root.clone());
|
||||
|
||||
for mode in [
|
||||
GrepOutputMode::Content,
|
||||
GrepOutputMode::FilesWithMatches,
|
||||
GrepOutputMode::Count,
|
||||
] {
|
||||
for (glob, file_type) in [(Some("other/*.rs"), None), (None, Some("python"))] {
|
||||
let mut request = grep_request("nested/selected.rs", "needle");
|
||||
request.output_mode = mode;
|
||||
request.glob = glob.map(str::to_string);
|
||||
request.file_type = file_type.map(str::to_string);
|
||||
|
||||
let excluded = run_grep(&root, selected.clone(), request, &readable).unwrap();
|
||||
assert_eq!(excluded.output, "", "mode {mode:?}");
|
||||
assert_eq!(excluded.match_count, 0, "mode {mode:?}");
|
||||
assert_eq!(excluded.matched_files, 0, "mode {mode:?}");
|
||||
assert!(!excluded.truncated, "mode {mode:?}");
|
||||
}
|
||||
|
||||
let mut request = grep_request("nested/selected.rs", "needle");
|
||||
request.output_mode = mode;
|
||||
request.glob = Some("nested/*.rs".to_string());
|
||||
request.file_type = Some("rust".to_string());
|
||||
let matched = run_grep(&root, selected.clone(), request, &readable).unwrap();
|
||||
|
||||
match mode {
|
||||
GrepOutputMode::Content => {
|
||||
assert_eq!(matched.match_count, 2);
|
||||
assert_eq!(matched.matched_files, 1);
|
||||
assert!(matched.output.starts_with("nested/selected.rs\n"));
|
||||
assert!(matched.output.contains("> 1 │ needle one"));
|
||||
assert!(matched.output.contains("> 2 │ needle two"));
|
||||
}
|
||||
GrepOutputMode::FilesWithMatches => {
|
||||
assert_eq!(matched.match_count, 1);
|
||||
assert_eq!(matched.matched_files, 1);
|
||||
assert_eq!(matched.output, "nested/selected.rs\n");
|
||||
}
|
||||
GrepOutputMode::Count => {
|
||||
assert_eq!(matched.match_count, 2);
|
||||
assert_eq!(matched.matched_files, 1);
|
||||
assert_eq!(matched.output, "nested/selected.rs:2\n");
|
||||
}
|
||||
}
|
||||
assert!(!matched.truncated, "mode {mode:?}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grep_direct_file_preserves_explicit_hidden_and_gitignored_behavior() {
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
let hidden = temp.path().join(".hidden.rs");
|
||||
let ignored = temp.path().join("ignored.rs");
|
||||
std::fs::write(&hidden, "needle hidden\n").unwrap();
|
||||
std::fs::write(&ignored, "needle ignored\n").unwrap();
|
||||
std::fs::write(temp.path().join(".gitignore"), "ignored.rs\n").unwrap();
|
||||
let root = temp.path().canonicalize().unwrap();
|
||||
let readable = RootAccess(root.clone());
|
||||
|
||||
for (path, expected) in [
|
||||
(".hidden.rs", "needle hidden"),
|
||||
("ignored.rs", "needle ignored"),
|
||||
] {
|
||||
let result = run_grep(
|
||||
&root,
|
||||
root.join(path),
|
||||
grep_request(path, "needle"),
|
||||
&readable,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(result.match_count, 1, "path {path}");
|
||||
assert!(result.output.contains(expected), "path {path}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grep_direct_file_preserves_case_multiline_and_bounds() {
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
let selected = temp.path().join("selected.txt");
|
||||
std::fs::write(&selected, "NEEDLE first\nstart\nfinish\nneedle last\n").unwrap();
|
||||
let root = temp.path().canonicalize().unwrap();
|
||||
let readable = RootAccess(root.clone());
|
||||
|
||||
let mut case_request = grep_request("selected.txt", "needle");
|
||||
case_request.case_insensitive = true;
|
||||
case_request.offset = 1;
|
||||
case_request.limit = 1;
|
||||
let bounded = run_grep(&root, selected.clone(), case_request, &readable).unwrap();
|
||||
assert_eq!(bounded.match_count, 1);
|
||||
assert!(!bounded.output.contains("NEEDLE first"));
|
||||
assert!(bounded.output.contains("needle last"));
|
||||
assert!(bounded.truncated);
|
||||
|
||||
let mut multiline_request = grep_request("selected.txt", "start\\nfinish");
|
||||
multiline_request.multiline = true;
|
||||
let multiline = run_grep(&root, selected, multiline_request, &readable).unwrap();
|
||||
assert_eq!(multiline.match_count, 1);
|
||||
assert!(multiline.output.contains("start\nfinish"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grep_returns_not_found_for_a_missing_direct_path() {
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
let root = temp.path().canonicalize().unwrap();
|
||||
let missing = root.join("missing.txt");
|
||||
let readable = RootAccess(root.clone());
|
||||
|
||||
let error = run_grep(
|
||||
&root,
|
||||
missing.clone(),
|
||||
grep_request("missing.txt", "needle"),
|
||||
&readable,
|
||||
)
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(error, FsError::NotFound(path) if path == missing));
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
#[test]
|
||||
fn grep_keeps_direct_symlink_directory_and_broken_path_guards() {
|
||||
use std::os::unix::fs::symlink;
|
||||
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
let root = temp.path().canonicalize().unwrap();
|
||||
let readable = RootAccess(root.clone());
|
||||
std::fs::create_dir(root.join("target-dir")).unwrap();
|
||||
std::fs::write(root.join("target-file.rs"), "needle file\n").unwrap();
|
||||
symlink(root.join("target-file.rs"), root.join("file-link.rs")).unwrap();
|
||||
symlink(root.join("target-dir"), root.join("directory-link")).unwrap();
|
||||
symlink(root.join("missing-target"), root.join("broken-link")).unwrap();
|
||||
|
||||
let request = |path: &str| grep_request(path, "needle");
|
||||
|
||||
let file_result = run_grep(
|
||||
&root,
|
||||
root.join("file-link.rs"),
|
||||
request("file-link.rs"),
|
||||
&readable,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(file_result.match_count, 1);
|
||||
assert!(file_result.output.starts_with("file-link.rs\n"));
|
||||
|
||||
let directory_error = run_grep(
|
||||
&root,
|
||||
root.join("directory-link"),
|
||||
request("directory-link"),
|
||||
&readable,
|
||||
)
|
||||
.unwrap_err();
|
||||
assert!(matches!(
|
||||
directory_error,
|
||||
FsError::SymlinkDirectoryNotTraversed { tool: "Grep", path, .. }
|
||||
if path == root.join("directory-link")
|
||||
));
|
||||
|
||||
let broken_error = run_grep(
|
||||
&root,
|
||||
root.join("broken-link"),
|
||||
request("broken-link"),
|
||||
&readable,
|
||||
)
|
||||
.unwrap_err();
|
||||
assert!(matches!(
|
||||
broken_error,
|
||||
FsError::BrokenSymlink { path, .. } if path == root.join("broken-link")
|
||||
));
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
#[test]
|
||||
fn grep_rejects_a_direct_special_file_as_invalid_argument() {
|
||||
use std::os::unix::net::UnixListener;
|
||||
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
let socket = temp.path().join("grep.sock");
|
||||
let _listener = UnixListener::bind(&socket).unwrap();
|
||||
let root = temp.path().canonicalize().unwrap();
|
||||
let readable = RootAccess(root.clone());
|
||||
|
||||
let error = run_grep(
|
||||
&root,
|
||||
socket,
|
||||
grep_request("grep.sock", "needle"),
|
||||
&readable,
|
||||
)
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(
|
||||
error,
|
||||
FsError::InvalidArgument(message)
|
||||
if message.contains("must be a regular file or directory")
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grep_content_groups_lines_by_file_and_marks_matches() {
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
std::fs::write(
|
||||
temp.path().join("first.txt"),
|
||||
"before\nneedle one\nafter\nomitted one\nomitted two\nbefore distant\nneedle distant\nafter distant\n",
|
||||
)
|
||||
.unwrap();
|
||||
std::fs::write(temp.path().join("second.txt"), "needle two\n").unwrap();
|
||||
let root = temp.path().canonicalize().unwrap();
|
||||
let readable = RootAccess(root.clone());
|
||||
|
||||
let grep = run_grep(
|
||||
&root,
|
||||
root.clone(),
|
||||
GrepRequest {
|
||||
pattern: "needle".to_string(),
|
||||
path: FsPath::root(),
|
||||
glob: Some("*.txt".to_string()),
|
||||
output_mode: GrepOutputMode::Content,
|
||||
case_insensitive: false,
|
||||
before_context: 1,
|
||||
after_context: 1,
|
||||
multiline: false,
|
||||
file_type: None,
|
||||
limit: 20,
|
||||
offset: 0,
|
||||
},
|
||||
&readable,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(grep.match_count, 3);
|
||||
assert_eq!(grep.matched_files, 2);
|
||||
assert_eq!(
|
||||
grep.output,
|
||||
concat!(
|
||||
"first.txt\n",
|
||||
" 1 │ before\n",
|
||||
" > 2 │ needle one\n",
|
||||
" 3 │ after\n",
|
||||
" …\n",
|
||||
" 6 │ before distant\n",
|
||||
" > 7 │ needle distant\n",
|
||||
" 8 │ after distant\n",
|
||||
"\n",
|
||||
"second.txt\n",
|
||||
" > 1 │ needle two\n",
|
||||
)
|
||||
);
|
||||
assert_eq!(grep.output.matches("first.txt").count(), 1);
|
||||
assert_eq!(grep.output.matches("second.txt").count(), 1);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,7 +5,8 @@ use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::FsError;
|
||||
|
||||
/// Logical path relative to the bound Workdir root.
|
||||
/// Scope-checked filesystem path. Relative paths resolve below the bound
|
||||
/// Workdir root; absolute paths require an explicit matching scope rule.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize)]
|
||||
#[serde(transparent)]
|
||||
pub struct FsPath(String);
|
||||
@@ -16,11 +17,30 @@ impl<'de> Deserialize<'de> for FsPath {
|
||||
D: serde::Deserializer<'de>,
|
||||
{
|
||||
let value = String::deserialize(deserializer)?;
|
||||
Self::new(&value).map_err(serde::de::Error::custom)
|
||||
Self::new_scoped(&value).map_err(serde::de::Error::custom)
|
||||
}
|
||||
}
|
||||
|
||||
impl FsPath {
|
||||
/// Construct a path for a scope-checked operation that may target an
|
||||
/// explicitly granted absolute path outside the provider root.
|
||||
pub fn new_scoped(value: impl Into<String>) -> Result<Self, FsError> {
|
||||
let value = value.into();
|
||||
if !Path::new(&value).is_absolute() {
|
||||
return Self::new(value);
|
||||
}
|
||||
if value.contains('\\') {
|
||||
return Err(FsError::InvalidPath(value));
|
||||
}
|
||||
if Path::new(&value)
|
||||
.components()
|
||||
.any(|component| component == Component::ParentDir)
|
||||
{
|
||||
return Err(FsError::InvalidPath(value));
|
||||
}
|
||||
Ok(Self(value))
|
||||
}
|
||||
|
||||
pub fn root() -> Self {
|
||||
Self(String::new())
|
||||
}
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::fmt::Write as _;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use crate::FsAccessPolicy;
|
||||
@@ -5,8 +7,8 @@ use grep_regex::RegexMatcherBuilder;
|
||||
use grep_searcher::sinks::UTF8 as UTF8Sink;
|
||||
use grep_searcher::{BinaryDetection, Searcher, SearcherBuilder, Sink, SinkContext, SinkMatch};
|
||||
use ignore::WalkBuilder;
|
||||
use ignore::overrides::OverrideBuilder;
|
||||
use ignore::types::TypesBuilder;
|
||||
use ignore::overrides::{Override, OverrideBuilder};
|
||||
use ignore::types::{Types, TypesBuilder};
|
||||
|
||||
use crate::{FsError, GrepOutputMode, GrepRequest, GrepResult, direct_symlink};
|
||||
|
||||
@@ -57,20 +59,11 @@ impl GrepReport {
|
||||
}
|
||||
}
|
||||
GrepOutputMode::Content => {
|
||||
for line in &self.lines {
|
||||
let separator = if line.is_match { ':' } else { '-' };
|
||||
let path = logical_display(root, &line.path);
|
||||
if self.show_line_numbers
|
||||
&& let Some(number) = line.line_number
|
||||
{
|
||||
output.push_str(&format!(
|
||||
"{path}{separator}{number}{separator}{}\n",
|
||||
line.text
|
||||
));
|
||||
} else {
|
||||
output.push_str(&format!("{path}{separator}{}\n", line.text));
|
||||
}
|
||||
}
|
||||
output.push_str(&render_content_lines(
|
||||
root,
|
||||
&self.lines,
|
||||
self.show_line_numbers,
|
||||
));
|
||||
}
|
||||
}
|
||||
GrepResult {
|
||||
@@ -82,6 +75,48 @@ impl GrepReport {
|
||||
}
|
||||
}
|
||||
|
||||
fn render_content_lines(root: &Path, lines: &[ContentLine], show_line_numbers: bool) -> String {
|
||||
let mut grouped = BTreeMap::<&Path, Vec<&ContentLine>>::new();
|
||||
for line in lines {
|
||||
grouped.entry(&line.path).or_default().push(line);
|
||||
}
|
||||
|
||||
let mut output = String::new();
|
||||
for (file_index, (path, file_lines)) in grouped.into_iter().enumerate() {
|
||||
if file_index > 0 {
|
||||
output.push('\n');
|
||||
}
|
||||
let _ = writeln!(output, "{}", logical_display(root, path));
|
||||
|
||||
let number_width = file_lines
|
||||
.iter()
|
||||
.filter_map(|line| line.line_number)
|
||||
.map(|number| number.to_string().len())
|
||||
.max()
|
||||
.unwrap_or(1);
|
||||
let mut previous_line_end = None;
|
||||
for line in file_lines {
|
||||
if let (Some(previous_end), Some(number)) = (previous_line_end, line.line_number)
|
||||
&& number > previous_end
|
||||
{
|
||||
let _ = writeln!(output, " …");
|
||||
}
|
||||
|
||||
let marker = if line.is_match { '>' } else { ' ' };
|
||||
if show_line_numbers && let Some(number) = line.line_number {
|
||||
let _ = writeln!(output, " {marker} {number:>number_width$} │ {}", line.text);
|
||||
} else {
|
||||
let _ = writeln!(output, " {marker} │ {}", line.text);
|
||||
}
|
||||
previous_line_end = line
|
||||
.line_number
|
||||
.map(|number| number + line.text.split('\n').count() as u64);
|
||||
}
|
||||
}
|
||||
|
||||
output
|
||||
}
|
||||
|
||||
fn logical_display(root: &Path, path: &Path) -> String {
|
||||
path.strip_prefix(root)
|
||||
.unwrap_or(path)
|
||||
@@ -91,6 +126,38 @@ fn logical_display(root: &Path, path: &Path) -> String {
|
||||
|
||||
const DEFAULT_HEAD_LIMIT: usize = 250;
|
||||
|
||||
fn build_overrides(base: &Path, glob: Option<&str>) -> Result<Option<Override>, FsError> {
|
||||
let Some(glob) = glob else {
|
||||
return Ok(None);
|
||||
};
|
||||
let mut builder = OverrideBuilder::new(base);
|
||||
builder
|
||||
.add(glob)
|
||||
.map_err(|error| FsError::InvalidGlob(error.to_string()))?;
|
||||
builder
|
||||
.build()
|
||||
.map(Some)
|
||||
.map_err(|error| FsError::InvalidGlob(error.to_string()))
|
||||
}
|
||||
|
||||
fn build_types(file_type: Option<&str>) -> Result<Option<Types>, FsError> {
|
||||
let Some(file_type) = file_type else {
|
||||
return Ok(None);
|
||||
};
|
||||
let mut builder = TypesBuilder::new();
|
||||
builder.add_defaults();
|
||||
builder.select(file_type);
|
||||
builder
|
||||
.build()
|
||||
.map(Some)
|
||||
.map_err(|error| FsError::InvalidArgument(format!("invalid type {file_type}: {error}")))
|
||||
}
|
||||
|
||||
fn direct_file_selected(path: &Path, overrides: Option<&Override>, types: Option<&Types>) -> bool {
|
||||
!overrides.is_some_and(|filter| filter.matched(path, false).is_ignore())
|
||||
&& !types.is_some_and(|filter| filter.matched(path, false).is_ignore())
|
||||
}
|
||||
|
||||
struct GrepParams {
|
||||
pattern: String,
|
||||
path: Option<PathBuf>,
|
||||
@@ -186,13 +253,15 @@ pub fn run_grep(
|
||||
std::io::ErrorKind::NotFound => FsError::NotFound(base.clone()),
|
||||
_ => FsError::io(&base, e),
|
||||
})?;
|
||||
if !base_meta.is_dir() {
|
||||
if !base_meta.is_file() && !base_meta.is_dir() {
|
||||
return Err(FsError::InvalidArgument(format!(
|
||||
"grep search path is not a directory: {}",
|
||||
"grep search path must be a regular file or directory: {}",
|
||||
base.display()
|
||||
)));
|
||||
}
|
||||
if let Some(info) = symlink.as_ref() {
|
||||
if base_meta.is_dir()
|
||||
&& let Some(info) = symlink.as_ref()
|
||||
{
|
||||
return Err(FsError::SymlinkDirectoryNotTraversed {
|
||||
tool: "Grep",
|
||||
path: base.clone(),
|
||||
@@ -200,32 +269,9 @@ pub fn run_grep(
|
||||
});
|
||||
}
|
||||
|
||||
let mut wb = WalkBuilder::new(&base);
|
||||
wb.hidden(true)
|
||||
.git_ignore(true)
|
||||
.git_global(true)
|
||||
.git_exclude(true)
|
||||
.ignore(true)
|
||||
.parents(true)
|
||||
.follow_links(false);
|
||||
|
||||
if let Some(t) = p.file_type.as_deref() {
|
||||
let mut tb = TypesBuilder::new();
|
||||
tb.add_defaults();
|
||||
tb.select(t);
|
||||
let types = tb
|
||||
.build()
|
||||
.map_err(|e| FsError::InvalidArgument(format!("invalid type {t}: {e}")))?;
|
||||
wb.types(types);
|
||||
}
|
||||
if let Some(g) = p.glob.as_deref() {
|
||||
let mut ob = OverrideBuilder::new(&base);
|
||||
ob.add(g).map_err(|e| FsError::InvalidGlob(e.to_string()))?;
|
||||
let ov = ob
|
||||
.build()
|
||||
.map_err(|e| FsError::InvalidGlob(e.to_string()))?;
|
||||
wb.overrides(ov);
|
||||
}
|
||||
let filter_base = if base_meta.is_file() { root } else { &base };
|
||||
let types = build_types(p.file_type.as_deref())?;
|
||||
let overrides = build_overrides(filter_base, p.glob.as_deref())?;
|
||||
|
||||
let mode = p.output_mode.unwrap_or_default();
|
||||
let head_limit = p.head_limit.unwrap_or(DEFAULT_HEAD_LIMIT);
|
||||
@@ -240,74 +286,133 @@ pub fn run_grep(
|
||||
lines: Vec::new(),
|
||||
truncated: false,
|
||||
};
|
||||
let mut matching_files_seen = 0;
|
||||
let mut matches_seen = 0;
|
||||
|
||||
// Per-mode walker state.
|
||||
let mut matching_files_seen: usize = 0;
|
||||
let mut matches_seen: usize = 0;
|
||||
if base_meta.is_file() {
|
||||
if direct_file_selected(&base, overrides.as_ref(), types.as_ref()) {
|
||||
scan_path(
|
||||
&mut searcher,
|
||||
&matcher,
|
||||
&base,
|
||||
mode,
|
||||
&mut report,
|
||||
&mut matching_files_seen,
|
||||
&mut matches_seen,
|
||||
offset,
|
||||
head_limit,
|
||||
)?;
|
||||
}
|
||||
return Ok(report.into_result(root));
|
||||
}
|
||||
|
||||
'walker: for entry in wb.build().flatten() {
|
||||
if !entry.file_type().map(|t| t.is_file()).unwrap_or(false) {
|
||||
let mut walker = WalkBuilder::new(&base);
|
||||
walker
|
||||
.hidden(true)
|
||||
.git_ignore(true)
|
||||
.git_global(true)
|
||||
.git_exclude(true)
|
||||
.ignore(true)
|
||||
.parents(true)
|
||||
.follow_links(false);
|
||||
if let Some(types) = types {
|
||||
walker.types(types);
|
||||
}
|
||||
if let Some(overrides) = overrides {
|
||||
walker.overrides(overrides);
|
||||
}
|
||||
|
||||
for entry in walker.build().flatten() {
|
||||
if !entry
|
||||
.file_type()
|
||||
.map(|kind| kind.is_file())
|
||||
.unwrap_or(false)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
let path = entry.path();
|
||||
if !access.is_readable(path) {
|
||||
continue;
|
||||
}
|
||||
|
||||
match mode {
|
||||
GrepOutputMode::FilesWithMatches => {
|
||||
let hit = scan_any_match(&mut searcher, &matcher, path)?;
|
||||
if !hit {
|
||||
continue;
|
||||
}
|
||||
if matching_files_seen >= offset {
|
||||
report.files.push(path.to_path_buf());
|
||||
if report.files.len() >= head_limit {
|
||||
report.truncated = true;
|
||||
break 'walker;
|
||||
}
|
||||
}
|
||||
matching_files_seen += 1;
|
||||
}
|
||||
GrepOutputMode::Count => {
|
||||
let count = scan_count(&mut searcher, &matcher, path)?;
|
||||
if count == 0 {
|
||||
continue;
|
||||
}
|
||||
if matching_files_seen >= offset {
|
||||
report.counts.push((path.to_path_buf(), count));
|
||||
if report.counts.len() >= head_limit {
|
||||
report.truncated = true;
|
||||
break 'walker;
|
||||
}
|
||||
}
|
||||
matching_files_seen += 1;
|
||||
}
|
||||
GrepOutputMode::Content => {
|
||||
let before_count = matches_seen;
|
||||
let mut sink = ContentSink {
|
||||
path: path.to_path_buf(),
|
||||
lines: &mut report.lines,
|
||||
matches_seen: &mut matches_seen,
|
||||
offset,
|
||||
head_limit,
|
||||
};
|
||||
searcher
|
||||
.search_path(&matcher, path, &mut sink)
|
||||
.map_err(|e| FsError::io(path, e))?;
|
||||
// If we hit head_limit during this file, stop walking.
|
||||
if matches_seen >= offset.saturating_add(head_limit) && matches_seen > before_count
|
||||
{
|
||||
report.truncated = true;
|
||||
break 'walker;
|
||||
}
|
||||
}
|
||||
if scan_path(
|
||||
&mut searcher,
|
||||
&matcher,
|
||||
path,
|
||||
mode,
|
||||
&mut report,
|
||||
&mut matching_files_seen,
|
||||
&mut matches_seen,
|
||||
offset,
|
||||
head_limit,
|
||||
)? {
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(report.into_result(root))
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn scan_path(
|
||||
searcher: &mut Searcher,
|
||||
matcher: &grep_regex::RegexMatcher,
|
||||
path: &Path,
|
||||
mode: GrepOutputMode,
|
||||
report: &mut GrepReport,
|
||||
matching_files_seen: &mut usize,
|
||||
matches_seen: &mut usize,
|
||||
offset: usize,
|
||||
head_limit: usize,
|
||||
) -> Result<bool, FsError> {
|
||||
match mode {
|
||||
GrepOutputMode::FilesWithMatches => {
|
||||
if !scan_any_match(searcher, matcher, path)? {
|
||||
return Ok(false);
|
||||
}
|
||||
if *matching_files_seen >= offset {
|
||||
report.files.push(path.to_path_buf());
|
||||
if report.files.len() >= head_limit {
|
||||
report.truncated = true;
|
||||
return Ok(true);
|
||||
}
|
||||
}
|
||||
*matching_files_seen += 1;
|
||||
}
|
||||
GrepOutputMode::Count => {
|
||||
let count = scan_count(searcher, matcher, path)?;
|
||||
if count == 0 {
|
||||
return Ok(false);
|
||||
}
|
||||
if *matching_files_seen >= offset {
|
||||
report.counts.push((path.to_path_buf(), count));
|
||||
if report.counts.len() >= head_limit {
|
||||
report.truncated = true;
|
||||
return Ok(true);
|
||||
}
|
||||
}
|
||||
*matching_files_seen += 1;
|
||||
}
|
||||
GrepOutputMode::Content => {
|
||||
let before_count = *matches_seen;
|
||||
let mut sink = ContentSink {
|
||||
path: path.to_path_buf(),
|
||||
lines: &mut report.lines,
|
||||
matches_seen,
|
||||
offset,
|
||||
head_limit,
|
||||
};
|
||||
searcher
|
||||
.search_path(matcher, path, &mut sink)
|
||||
.map_err(|error| FsError::io(path, error))?;
|
||||
if *matches_seen >= offset.saturating_add(head_limit) && *matches_seen > before_count {
|
||||
report.truncated = true;
|
||||
return Ok(true);
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(false)
|
||||
}
|
||||
|
||||
fn scan_any_match(
|
||||
searcher: &mut Searcher,
|
||||
matcher: &grep_regex::RegexMatcher,
|
||||
|
||||
@@ -7,6 +7,7 @@ license.workspace = true
|
||||
[dependencies]
|
||||
arc-swap = "1"
|
||||
agen = { workspace = true }
|
||||
decodal.workspace = true
|
||||
protocol = { workspace = true }
|
||||
serde = { workspace = true, features = ["derive"] }
|
||||
serde_json = { workspace = true }
|
||||
|
||||
@@ -0,0 +1,318 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use decodal::{Data, Engine, ImportLoader, LoadedImport};
|
||||
use serde_json::{Map, Number, Value};
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
use crate::profile::ProfileError;
|
||||
|
||||
pub const BUILTIN_PROFILE_CATALOG_ID: &str = "builtin-profiles-v2";
|
||||
pub const BUILTIN_DEFAULT_PROFILE: &str = "builtin:default";
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct BuiltinProfileImport {
|
||||
pub specifier: &'static str,
|
||||
pub resolved_path: &'static str,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct BuiltinProfileResource {
|
||||
pub selector: Option<&'static str>,
|
||||
pub path: &'static str,
|
||||
pub source: &'static str,
|
||||
pub description: &'static str,
|
||||
pub imports: &'static [BuiltinProfileImport],
|
||||
}
|
||||
|
||||
const BASE_PATH: &str = "profiles/base.dcdl";
|
||||
const BASE_IMPORT: &[BuiltinProfileImport] = &[BuiltinProfileImport {
|
||||
specifier: "./base.dcdl",
|
||||
resolved_path: BASE_PATH,
|
||||
}];
|
||||
const NO_IMPORTS: &[BuiltinProfileImport] = &[];
|
||||
|
||||
pub const BUILTIN_PROFILE_RESOURCES: &[BuiltinProfileResource] = &[
|
||||
BuiltinProfileResource {
|
||||
selector: None,
|
||||
path: BASE_PATH,
|
||||
source: include_str!("../../../resources/profiles/base.dcdl"),
|
||||
description: "Shared built-in Profile defaults.",
|
||||
imports: NO_IMPORTS,
|
||||
},
|
||||
BuiltinProfileResource {
|
||||
selector: Some(BUILTIN_DEFAULT_PROFILE),
|
||||
path: "profiles/default.dcdl",
|
||||
source: include_str!("../../../resources/profiles/default.dcdl"),
|
||||
description: "Standalone Yoi coding profile.",
|
||||
imports: BASE_IMPORT,
|
||||
},
|
||||
BuiltinProfileResource {
|
||||
selector: Some("builtin:coder"),
|
||||
path: "profiles/coder.dcdl",
|
||||
source: include_str!("../../../resources/profiles/coder.dcdl"),
|
||||
description: "Ticket implementation with direct Reviewer SubWorkers.",
|
||||
imports: BASE_IMPORT,
|
||||
},
|
||||
BuiltinProfileResource {
|
||||
selector: Some("builtin:companion"),
|
||||
path: "profiles/companion.dcdl",
|
||||
source: include_str!("../../../resources/profiles/companion.dcdl"),
|
||||
description: "General assistance with Workspace tools.",
|
||||
imports: BASE_IMPORT,
|
||||
},
|
||||
BuiltinProfileResource {
|
||||
selector: Some("builtin:intake"),
|
||||
path: "profiles/intake.dcdl",
|
||||
source: include_str!("../../../resources/profiles/intake.dcdl"),
|
||||
description: "Read-only intake and planning.",
|
||||
imports: BASE_IMPORT,
|
||||
},
|
||||
BuiltinProfileResource {
|
||||
selector: Some("builtin:reviewer"),
|
||||
path: "profiles/reviewer.dcdl",
|
||||
source: include_str!("../../../resources/profiles/reviewer.dcdl"),
|
||||
description: "Independent review of a published Merge Request source.",
|
||||
imports: BASE_IMPORT,
|
||||
},
|
||||
BuiltinProfileResource {
|
||||
selector: Some("builtin:orchestrator"),
|
||||
path: "profiles/orchestrator.dcdl",
|
||||
source: include_str!("../../../resources/profiles/orchestrator.dcdl"),
|
||||
description: "Workspace orchestration and Worker control.",
|
||||
imports: BASE_IMPORT,
|
||||
},
|
||||
BuiltinProfileResource {
|
||||
selector: Some("builtin:memory-consolidation"),
|
||||
path: "profiles/memory-consolidation.dcdl",
|
||||
source: include_str!("../../../resources/profiles/memory-consolidation.dcdl"),
|
||||
description: "Internal Memory consolidation service.",
|
||||
imports: BASE_IMPORT,
|
||||
},
|
||||
];
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct BuiltinProfileCatalogSnapshot {
|
||||
pub id: &'static str,
|
||||
pub sources: BTreeMap<String, String>,
|
||||
pub entrypoints: BTreeMap<String, String>,
|
||||
pub imports: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
impl BuiltinProfileCatalogSnapshot {
|
||||
pub fn digest(&self) -> String {
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(self.id.as_bytes());
|
||||
for (path, source) in &self.sources {
|
||||
hasher.update((path.len() as u64).to_le_bytes());
|
||||
hasher.update(path.as_bytes());
|
||||
hasher.update((source.len() as u64).to_le_bytes());
|
||||
hasher.update(source.as_bytes());
|
||||
}
|
||||
for (selector, path) in &self.entrypoints {
|
||||
hasher.update((selector.len() as u64).to_le_bytes());
|
||||
hasher.update(selector.as_bytes());
|
||||
hasher.update((path.len() as u64).to_le_bytes());
|
||||
hasher.update(path.as_bytes());
|
||||
}
|
||||
for (request, resolved_path) in &self.imports {
|
||||
hasher.update((request.len() as u64).to_le_bytes());
|
||||
hasher.update(request.as_bytes());
|
||||
hasher.update((resolved_path.len() as u64).to_le_bytes());
|
||||
hasher.update(resolved_path.as_bytes());
|
||||
}
|
||||
format!("sha256:{:x}", hasher.finalize())
|
||||
}
|
||||
}
|
||||
|
||||
pub fn builtin_profile_catalog_snapshot() -> BuiltinProfileCatalogSnapshot {
|
||||
let mut sources = BTreeMap::new();
|
||||
let mut entrypoints = BTreeMap::new();
|
||||
let mut imports = BTreeMap::new();
|
||||
|
||||
for resource in BUILTIN_PROFILE_RESOURCES {
|
||||
sources.insert(resource.path.to_owned(), resource.source.to_owned());
|
||||
for import in resource.imports {
|
||||
imports.insert(
|
||||
format!("{}\0{}", resource.path, import.specifier),
|
||||
import.resolved_path.to_owned(),
|
||||
);
|
||||
}
|
||||
if let Some(selector) = resource.selector {
|
||||
entrypoints.insert(selector.to_owned(), resource.path.to_owned());
|
||||
}
|
||||
}
|
||||
|
||||
BuiltinProfileCatalogSnapshot {
|
||||
id: BUILTIN_PROFILE_CATALOG_ID,
|
||||
sources,
|
||||
entrypoints,
|
||||
imports,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn builtin_profile_entrypoints() -> impl Iterator<Item = &'static BuiltinProfileResource> {
|
||||
BUILTIN_PROFILE_RESOURCES
|
||||
.iter()
|
||||
.filter(|resource| resource.selector.is_some())
|
||||
}
|
||||
|
||||
pub(crate) fn resolve_builtin_profile_artifact(
|
||||
selector: &str,
|
||||
) -> Result<Option<Value>, ProfileError> {
|
||||
let catalog = builtin_profile_catalog_snapshot();
|
||||
let Some(entrypoint) = catalog.entrypoints.get(selector) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let source = catalog
|
||||
.sources
|
||||
.get(entrypoint)
|
||||
.expect("built-in Profile entrypoint must name a source")
|
||||
.clone();
|
||||
let mut engine = Engine::new(BuiltinProfileImportLoader {
|
||||
sources: catalog.sources,
|
||||
});
|
||||
let module = engine
|
||||
.add_root_source(entrypoint, entrypoint, &source)
|
||||
.map_err(|error| ProfileError::BuiltinProfileEvaluation {
|
||||
selector: selector.to_owned(),
|
||||
message: format!("{error:?}"),
|
||||
})?;
|
||||
let value =
|
||||
engine
|
||||
.eval_module(module)
|
||||
.map_err(|error| ProfileError::BuiltinProfileEvaluation {
|
||||
selector: selector.to_owned(),
|
||||
message: format!("{error:?}"),
|
||||
})?;
|
||||
let data =
|
||||
engine
|
||||
.materialize(&value)
|
||||
.map_err(|error| ProfileError::BuiltinProfileEvaluation {
|
||||
selector: selector.to_owned(),
|
||||
message: format!("{error:?}"),
|
||||
})?;
|
||||
Ok(Some(data_to_json(&data)))
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct BuiltinProfileImportLoader {
|
||||
sources: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
impl ImportLoader for BuiltinProfileImportLoader {
|
||||
fn load(
|
||||
&mut self,
|
||||
current_key: Option<&str>,
|
||||
specifier: &str,
|
||||
) -> decodal::Result<LoadedImport> {
|
||||
let current_key = current_key.ok_or_else(|| {
|
||||
decodal::Diagnostic::new(
|
||||
decodal::DiagnosticKind::Import,
|
||||
decodal::Span::default(),
|
||||
format!("built-in Profile import `{specifier}` has no source context"),
|
||||
)
|
||||
})?;
|
||||
let resolved = resolve_import_path(current_key, specifier).ok_or_else(|| {
|
||||
decodal::Diagnostic::new(
|
||||
decodal::DiagnosticKind::Import,
|
||||
decodal::Span::default(),
|
||||
format!("built-in Profile import `{specifier}` from `{current_key}` is invalid"),
|
||||
)
|
||||
})?;
|
||||
let source = self.sources.get(&resolved).ok_or_else(|| {
|
||||
decodal::Diagnostic::new(
|
||||
decodal::DiagnosticKind::Import,
|
||||
decodal::Span::default(),
|
||||
format!("built-in Profile import `{specifier}` from `{current_key}` was not found"),
|
||||
)
|
||||
})?;
|
||||
Ok(LoadedImport::source(
|
||||
resolved.clone(),
|
||||
resolved,
|
||||
source.clone(),
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
fn resolve_import_path(current_key: &str, specifier: &str) -> Option<String> {
|
||||
let current_parent = current_key
|
||||
.rsplit_once('/')
|
||||
.map_or("", |(parent, _)| parent);
|
||||
let joined = if let Some(relative) = specifier.strip_prefix("./") {
|
||||
format!("{current_parent}/{relative}")
|
||||
} else {
|
||||
return None;
|
||||
};
|
||||
if joined
|
||||
.split('/')
|
||||
.any(|segment| segment.is_empty() || segment == "." || segment == "..")
|
||||
{
|
||||
return None;
|
||||
}
|
||||
Some(joined)
|
||||
}
|
||||
|
||||
fn data_to_json(data: &Data) -> Value {
|
||||
match data {
|
||||
Data::Bool(value) => Value::Bool(*value),
|
||||
Data::Int(value) => Value::Number(Number::from(*value)),
|
||||
Data::Float(value) => Number::from_f64(*value)
|
||||
.map(Value::Number)
|
||||
.unwrap_or(Value::Null),
|
||||
Data::String(value) => Value::String(value.clone()),
|
||||
Data::Array(values) => Value::Array(values.iter().map(data_to_json).collect()),
|
||||
Data::Object(fields) => Value::Object(
|
||||
fields
|
||||
.iter()
|
||||
.map(|field| (field.name.clone(), data_to_json(&field.value)))
|
||||
.collect::<Map<_, _>>(),
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn catalog_has_one_explicit_entrypoint_for_each_builtin_profile() {
|
||||
let catalog = builtin_profile_catalog_snapshot();
|
||||
assert_eq!(catalog.sources.len(), BUILTIN_PROFILE_RESOURCES.len());
|
||||
assert_eq!(catalog.entrypoints.len() + 1, catalog.sources.len());
|
||||
assert_eq!(
|
||||
catalog.entrypoints.get(BUILTIN_DEFAULT_PROFILE),
|
||||
Some(&"profiles/default.dcdl".to_owned())
|
||||
);
|
||||
assert!(catalog.digest().starts_with("sha256:"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_profile_evaluates_from_the_shared_resource_graph() {
|
||||
let value = resolve_builtin_profile_artifact(BUILTIN_DEFAULT_PROFILE)
|
||||
.expect("evaluate built-in default")
|
||||
.expect("default exists");
|
||||
assert_eq!(value["slug"], "default");
|
||||
assert_eq!(value["feature"]["task"]["enabled"], true);
|
||||
assert_eq!(value["feature"]["sub_worker"]["enabled"], true);
|
||||
assert_eq!(value["feature"]["memory"]["enabled"], false);
|
||||
assert_eq!(value["feature"]["ticket"]["enabled"], false);
|
||||
assert_eq!(value["feature"]["worker"]["enabled"], false);
|
||||
assert_eq!(value["feature"]["manage_workdir"]["enabled"], false);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn imports_cannot_escape_the_builtin_resource_catalog() {
|
||||
assert_eq!(
|
||||
resolve_import_path("profiles/default.dcdl", "./base.dcdl").as_deref(),
|
||||
Some("profiles/base.dcdl")
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_import_path("profiles/default.dcdl", "../outside.dcdl"),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_import_path("profiles/default.dcdl", "/outside.dcdl"),
|
||||
None
|
||||
);
|
||||
}
|
||||
}
|
||||
+189
-87
@@ -18,10 +18,11 @@ use crate::model::{AuthRef, ModelManifest, ReasoningControl};
|
||||
use crate::plugin::PluginConfig;
|
||||
use crate::{
|
||||
CompactionConfig, EngineManifest, FeatureConfig, FeatureFlagConfig, FileUploadLimits,
|
||||
McpConfig, McpEnvValue, McpStdioCwdPolicy, MemoryConfig, MemoryFeatureConfig,
|
||||
MergeRequestFeatureConfig, ScopeConfig, SessionConfig, SkillsConfig, TicketFeatureConfig,
|
||||
ToolOutputLimits, ToolPermissionConfig, ToolPermissionRule, WebConfig, WorkerFeatureConfig,
|
||||
WorkerManifest, WorkerMeta,
|
||||
McpConfig, McpEnvValue, McpStdioCwdPolicy, MemoryConsolidationProfileConfig,
|
||||
MemoryExtractionProfileConfig, MemoryFeatureProfileConfig, MemoryResidentProfileConfig,
|
||||
MergeRequestFeatureConfig, ResolvedMemoryFeatureConfig, ScopeConfig, SessionConfig,
|
||||
SkillsConfig, TicketFeatureConfig, ToolOutputLimits, ToolPermissionConfig, ToolPermissionRule,
|
||||
WebConfig, WorkerFeatureConfig, WorkerManifest, WorkerMeta,
|
||||
};
|
||||
|
||||
/// Partial-form Worker manifest. Every field is optional; one or more
|
||||
@@ -67,9 +68,6 @@ pub struct WorkerManifestConfig {
|
||||
/// First-class web tool opt-in. See [`WebConfig`].
|
||||
#[serde(default)]
|
||||
pub web: Option<WebConfig>,
|
||||
/// Memory subsystem opt-in. See [`MemoryConfig`].
|
||||
#[serde(default)]
|
||||
pub memory: Option<MemoryConfig>,
|
||||
/// External Agent Skills directories. See [`crate::SkillsConfig`].
|
||||
#[serde(default)]
|
||||
pub skills: Option<SkillsConfig>,
|
||||
@@ -92,6 +90,8 @@ pub struct FeatureConfigPartial {
|
||||
#[serde(default)]
|
||||
pub worker: Option<WorkerFeatureConfigPartial>,
|
||||
#[serde(default)]
|
||||
pub workspace_worker_discovery: Option<FeatureFlagConfigPartial>,
|
||||
#[serde(default)]
|
||||
pub objective: Option<FeatureFlagConfigPartial>,
|
||||
#[serde(default)]
|
||||
pub manage_workdir: Option<FeatureFlagConfigPartial>,
|
||||
@@ -119,6 +119,11 @@ impl FeatureConfigPartial {
|
||||
),
|
||||
flow: merge_option(self.flow, other.flow, FeatureFlagConfigPartial::merge),
|
||||
worker: merge_option(self.worker, other.worker, WorkerFeatureConfigPartial::merge),
|
||||
workspace_worker_discovery: merge_option(
|
||||
self.workspace_worker_discovery,
|
||||
other.workspace_worker_discovery,
|
||||
FeatureFlagConfigPartial::merge,
|
||||
),
|
||||
objective: merge_option(
|
||||
self.objective,
|
||||
other.objective,
|
||||
@@ -186,18 +191,86 @@ impl From<WorkerFeatureConfigPartial> for WorkerFeatureConfig {
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct MemoryFeatureConfigPartial {
|
||||
#[serde(default)]
|
||||
pub enabled: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub staging: Option<bool>,
|
||||
pub staging_tools: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub resident: Option<MemoryResidentProfileConfigPartial>,
|
||||
#[serde(default)]
|
||||
pub extraction: Option<MemoryExtractionProfileConfigPartial>,
|
||||
#[serde(default)]
|
||||
pub consolidation: Option<MemoryConsolidationProfileConfigPartial>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct MemoryResidentProfileConfigPartial {
|
||||
#[serde(default)]
|
||||
pub inject_summary: Option<bool>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct MemoryExtractionProfileConfigPartial {
|
||||
#[serde(default)]
|
||||
pub enabled: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub model: Option<ModelManifest>,
|
||||
#[serde(default)]
|
||||
pub threshold: Option<u64>,
|
||||
#[serde(default)]
|
||||
pub worker_max_turns: Option<u32>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct MemoryConsolidationProfileConfigPartial {
|
||||
#[serde(default)]
|
||||
pub request_enabled: Option<bool>,
|
||||
}
|
||||
|
||||
impl MemoryFeatureConfigPartial {
|
||||
fn merge(self, other: Self) -> Self {
|
||||
Self {
|
||||
enabled: other.enabled.or(self.enabled),
|
||||
staging: other.staging.or(self.staging),
|
||||
staging_tools: other.staging_tools.or(self.staging_tools),
|
||||
resident: merge_option(
|
||||
self.resident,
|
||||
other.resident,
|
||||
MemoryResidentProfileConfigPartial::merge,
|
||||
),
|
||||
extraction: merge_option(
|
||||
self.extraction,
|
||||
other.extraction,
|
||||
MemoryExtractionProfileConfigPartial::merge,
|
||||
),
|
||||
consolidation: merge_option(
|
||||
self.consolidation,
|
||||
other.consolidation,
|
||||
MemoryConsolidationProfileConfigPartial::merge,
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl MemoryResidentProfileConfigPartial {
|
||||
fn merge(self, other: Self) -> Self {
|
||||
Self {
|
||||
inject_summary: other.inject_summary.or(self.inject_summary),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl MemoryExtractionProfileConfigPartial {
|
||||
fn merge(self, other: Self) -> Self {
|
||||
Self {
|
||||
enabled: other.enabled.or(self.enabled),
|
||||
model: other.model.or(self.model),
|
||||
threshold: other.threshold.or(self.threshold),
|
||||
worker_max_turns: other.worker_max_turns.or(self.worker_max_turns),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -246,13 +319,21 @@ impl MergeRequestFeatureConfigPartial {
|
||||
}
|
||||
}
|
||||
|
||||
impl MemoryConsolidationProfileConfigPartial {
|
||||
fn merge(self, other: Self) -> Self {
|
||||
Self {
|
||||
request_enabled: other.request_enabled.or(self.request_enabled),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<FeatureConfigPartial> for FeatureConfig {
|
||||
fn from(value: FeatureConfigPartial) -> Self {
|
||||
Self {
|
||||
task: value.task.map(FeatureFlagConfig::from).unwrap_or_default(),
|
||||
memory: value
|
||||
.memory
|
||||
.map(MemoryFeatureConfig::from)
|
||||
.map(ResolvedMemoryFeatureConfig::from)
|
||||
.unwrap_or_default(),
|
||||
web: value.web.map(FeatureFlagConfig::from).unwrap_or_default(),
|
||||
image: value.image.map(FeatureFlagConfig::from).unwrap_or_default(),
|
||||
@@ -265,6 +346,10 @@ impl From<FeatureConfigPartial> for FeatureConfig {
|
||||
.worker
|
||||
.map(WorkerFeatureConfig::from)
|
||||
.unwrap_or_default(),
|
||||
workspace_worker_discovery: value
|
||||
.workspace_worker_discovery
|
||||
.map(FeatureFlagConfig::from)
|
||||
.unwrap_or_default(),
|
||||
objective: value
|
||||
.objective
|
||||
.map(FeatureFlagConfig::from)
|
||||
@@ -318,20 +403,52 @@ impl From<WorkerFeatureConfig> for WorkerFeatureConfigPartial {
|
||||
}
|
||||
}
|
||||
|
||||
impl From<MemoryFeatureConfigPartial> for MemoryFeatureConfig {
|
||||
impl From<MemoryFeatureConfigPartial> for ResolvedMemoryFeatureConfig {
|
||||
fn from(value: MemoryFeatureConfigPartial) -> Self {
|
||||
let resident = value.resident.unwrap_or_default();
|
||||
let extraction = value.extraction.unwrap_or_default();
|
||||
let consolidation = value.consolidation.unwrap_or_default();
|
||||
Self {
|
||||
enabled: value.enabled.unwrap_or_default(),
|
||||
staging: value.staging.unwrap_or_default(),
|
||||
profile: MemoryFeatureProfileConfig {
|
||||
enabled: value.enabled.unwrap_or_default(),
|
||||
staging_tools: value.staging_tools.unwrap_or_default(),
|
||||
resident: MemoryResidentProfileConfig {
|
||||
inject_summary: resident.inject_summary.unwrap_or(true),
|
||||
},
|
||||
extraction: MemoryExtractionProfileConfig {
|
||||
enabled: extraction.enabled.unwrap_or(true),
|
||||
model: extraction.model,
|
||||
threshold: extraction.threshold.or(Some(50_000)),
|
||||
worker_max_turns: extraction
|
||||
.worker_max_turns
|
||||
.or(defaults::MEMORY_EXTRACT_WORKER_MAX_TURNS),
|
||||
},
|
||||
consolidation: MemoryConsolidationProfileConfig {
|
||||
request_enabled: consolidation.request_enabled.unwrap_or(true),
|
||||
},
|
||||
},
|
||||
workspace_settings: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<MemoryFeatureConfig> for MemoryFeatureConfigPartial {
|
||||
fn from(value: MemoryFeatureConfig) -> Self {
|
||||
impl From<ResolvedMemoryFeatureConfig> for MemoryFeatureConfigPartial {
|
||||
fn from(value: ResolvedMemoryFeatureConfig) -> Self {
|
||||
Self {
|
||||
enabled: Some(value.enabled),
|
||||
staging: Some(value.staging),
|
||||
enabled: Some(value.profile.enabled),
|
||||
staging_tools: Some(value.profile.staging_tools),
|
||||
resident: Some(MemoryResidentProfileConfigPartial {
|
||||
inject_summary: Some(value.profile.resident.inject_summary),
|
||||
}),
|
||||
extraction: Some(MemoryExtractionProfileConfigPartial {
|
||||
enabled: Some(value.profile.extraction.enabled),
|
||||
model: value.profile.extraction.model,
|
||||
threshold: value.profile.extraction.threshold,
|
||||
worker_max_turns: value.profile.extraction.worker_max_turns,
|
||||
}),
|
||||
consolidation: Some(MemoryConsolidationProfileConfigPartial {
|
||||
request_enabled: Some(value.profile.consolidation.request_enabled),
|
||||
}),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -394,6 +511,7 @@ impl From<FeatureConfig> for FeatureConfigPartial {
|
||||
sub_worker: Some(value.sub_worker.into()),
|
||||
flow: Some(value.flow.into()),
|
||||
worker: Some(value.worker.into()),
|
||||
workspace_worker_discovery: Some(value.workspace_worker_discovery.into()),
|
||||
objective: Some(value.objective.into()),
|
||||
manage_workdir: Some(value.manage_workdir.into()),
|
||||
ticket: Some(value.ticket.into()),
|
||||
@@ -531,13 +649,9 @@ pub(crate) fn reject_removed_manifest_fields(s: &str) -> Result<(), toml::de::Er
|
||||
(removed; use compaction.prune_protected_tokens)",
|
||||
));
|
||||
}
|
||||
if value
|
||||
.get("memory")
|
||||
.and_then(toml::Value::as_table)
|
||||
.is_some_and(|table| table.contains_key("extract_worker_max_input_tokens"))
|
||||
{
|
||||
if value.get("memory").is_some() {
|
||||
return Err(toml::de::Error::custom(
|
||||
"unknown field in manifest: memory.extract_worker_max_input_tokens (removed)",
|
||||
"unknown field in manifest: memory (removed; configure feature.memory)",
|
||||
));
|
||||
}
|
||||
if value
|
||||
@@ -566,15 +680,16 @@ impl WorkerManifestConfig {
|
||||
})
|
||||
}
|
||||
|
||||
/// Base config populated with the in-code defaults listed in
|
||||
/// [`crate::defaults`]. Profile and one-file Manifest resolvers start
|
||||
/// from this layer so every per-field default lives at exactly one
|
||||
/// call site (the `defaults` module).
|
||||
/// Base config populated with the in-code per-field defaults listed in
|
||||
/// [`crate::defaults`]. This is not a selectable Profile and does not
|
||||
/// enable a launch capability surface. Profile and one-file Manifest
|
||||
/// resolvers start from this layer so every per-field default lives at
|
||||
/// exactly one call site (the `defaults` module).
|
||||
///
|
||||
/// `TryFrom<WorkerManifestConfig>` also reads the same constants as a
|
||||
/// belt-and-suspenders fallback, so a manually-constructed config
|
||||
/// that skips this layer still resolves to the same values.
|
||||
pub fn builtin_defaults() -> Self {
|
||||
pub fn resolution_defaults() -> Self {
|
||||
Self {
|
||||
engine: EngineManifestConfig {
|
||||
tool_output: ToolOutputLimitsPartial {
|
||||
@@ -620,11 +735,6 @@ impl WorkerManifestConfig {
|
||||
for rule in &mut self.delegation_scope.deny {
|
||||
rule.target = join_if_relative(base, &rule.target);
|
||||
}
|
||||
if let Some(ref mut memory) = self.memory
|
||||
&& let Some(ref mut root) = memory.workspace_root
|
||||
{
|
||||
*root = join_if_relative(base, root);
|
||||
}
|
||||
if let Some(ref mut compaction) = self.compaction
|
||||
&& let Some(ref mut cp) = compaction.model
|
||||
{
|
||||
@@ -669,7 +779,6 @@ impl WorkerManifestConfig {
|
||||
CompactionConfigPartial::merge,
|
||||
),
|
||||
web: merge_option(self.web, upper.web, WebConfig::merge),
|
||||
memory: merge_option(self.memory, upper.memory, MemoryConfig::merge),
|
||||
skills: merge_option(self.skills, upper.skills, SkillsConfig::merge),
|
||||
}
|
||||
}
|
||||
@@ -741,32 +850,6 @@ impl crate::WebFetchConfig {
|
||||
}
|
||||
}
|
||||
|
||||
impl MemoryConfig {
|
||||
fn merge(self, upper: Self) -> Self {
|
||||
Self {
|
||||
workspace_root: upper.workspace_root.or(self.workspace_root),
|
||||
query_result_limit: upper.query_result_limit.or(self.query_result_limit),
|
||||
query_excerpt_lines: upper.query_excerpt_lines.or(self.query_excerpt_lines),
|
||||
inject_summary: upper.inject_summary.or(self.inject_summary),
|
||||
workspace_id: upper.workspace_id.or(self.workspace_id),
|
||||
settings_revision: upper.settings_revision.or(self.settings_revision),
|
||||
language: upper.language.or(self.language),
|
||||
extract_model: upper.extract_model.or(self.extract_model),
|
||||
extract_threshold: upper.extract_threshold.or(self.extract_threshold),
|
||||
extract_worker_max_turns: upper
|
||||
.extract_worker_max_turns
|
||||
.or(self.extract_worker_max_turns),
|
||||
consolidation_model: upper.consolidation_model.or(self.consolidation_model),
|
||||
consolidation_threshold_files: upper
|
||||
.consolidation_threshold_files
|
||||
.or(self.consolidation_threshold_files),
|
||||
consolidation_threshold_bytes: upper
|
||||
.consolidation_threshold_bytes
|
||||
.or(self.consolidation_threshold_bytes),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl WorkerMetaConfig {
|
||||
fn merge(self, upper: Self) -> Self {
|
||||
Self {
|
||||
@@ -1210,7 +1293,6 @@ impl TryFrom<WorkerManifestConfig> for WorkerManifest {
|
||||
mcp: cfg.mcp,
|
||||
compaction,
|
||||
web: cfg.web,
|
||||
memory: cfg.memory,
|
||||
skills: cfg.skills,
|
||||
profile: None,
|
||||
})
|
||||
@@ -1258,7 +1340,6 @@ mod tests {
|
||||
session: None,
|
||||
compaction: None,
|
||||
web: None,
|
||||
memory: None,
|
||||
skills: None,
|
||||
}
|
||||
}
|
||||
@@ -1833,29 +1914,50 @@ prune_protected_turns = 3
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_toml_rejects_removed_extract_worker_max_input_tokens_field() {
|
||||
let bad = r#"
|
||||
[memory]
|
||||
extract_worker_max_input_tokens = 30000
|
||||
"#;
|
||||
let err = WorkerManifestConfig::from_toml(bad).unwrap_err();
|
||||
assert!(
|
||||
err.to_string()
|
||||
.contains("memory.extract_worker_max_input_tokens"),
|
||||
"unexpected error: {err}"
|
||||
);
|
||||
fn from_toml_accepts_memory_extraction_settings_only_under_feature_memory() {
|
||||
let cfg = WorkerManifestConfig::from_toml(
|
||||
r#"
|
||||
[feature.memory]
|
||||
enabled = true
|
||||
staging_tools = false
|
||||
|
||||
[feature.memory.resident]
|
||||
inject_summary = false
|
||||
|
||||
[feature.memory.extraction]
|
||||
enabled = true
|
||||
threshold = 42000
|
||||
worker_max_turns = 2
|
||||
|
||||
[feature.memory.consolidation]
|
||||
request_enabled = false
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
let memory = cfg.feature.memory.unwrap();
|
||||
assert_eq!(memory.enabled, Some(true));
|
||||
assert_eq!(memory.staging_tools, Some(false));
|
||||
assert_eq!(memory.resident.unwrap().inject_summary, Some(false));
|
||||
assert_eq!(memory.consolidation.unwrap().request_enabled, Some(false));
|
||||
let extraction = memory.extraction.unwrap();
|
||||
assert_eq!(extraction.enabled, Some(true));
|
||||
assert_eq!(extraction.threshold, Some(42_000));
|
||||
assert_eq!(extraction.worker_max_turns, Some(2));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_toml_accepts_extract_worker_max_turns() {
|
||||
let cfg = WorkerManifestConfig::from_toml(
|
||||
fn from_toml_rejects_legacy_top_level_memory_authority() {
|
||||
let err = WorkerManifestConfig::from_toml(
|
||||
r#"
|
||||
[memory]
|
||||
extract_worker_max_turns = 2
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(cfg.memory.unwrap().extract_worker_max_turns, Some(2));
|
||||
.unwrap_err();
|
||||
assert!(
|
||||
err.to_string().contains("memory"),
|
||||
"unexpected error: {err}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -1935,7 +2037,7 @@ worker_max_turns = 7
|
||||
fn feature_flags_default_disabled_in_resolved_manifest() {
|
||||
let manifest: WorkerManifest = minimal_valid().try_into().unwrap();
|
||||
assert!(!manifest.feature.task.enabled);
|
||||
assert!(!manifest.feature.memory.enabled);
|
||||
assert!(!manifest.feature.memory.profile.enabled);
|
||||
assert!(!manifest.feature.web.enabled);
|
||||
assert!(!manifest.feature.sub_worker.enabled);
|
||||
assert!(!manifest.feature.objective.enabled);
|
||||
@@ -1973,7 +2075,7 @@ enabled = false
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
let manifest: WorkerManifest = WorkerManifestConfig::builtin_defaults()
|
||||
let manifest: WorkerManifest = WorkerManifestConfig::resolution_defaults()
|
||||
.merge(cfg)
|
||||
.merge(WorkerManifestConfig {
|
||||
worker: WorkerMetaConfig {
|
||||
@@ -2012,8 +2114,8 @@ enabled = false
|
||||
}
|
||||
);
|
||||
assert!(!manifest.feature.orchestration.enabled);
|
||||
assert!(!manifest.feature.memory.enabled);
|
||||
assert!(!manifest.feature.memory.staging);
|
||||
assert!(!manifest.feature.memory.profile.enabled);
|
||||
assert!(!manifest.feature.memory.profile.staging_tools);
|
||||
assert!(!manifest.feature.objective.enabled);
|
||||
}
|
||||
|
||||
@@ -2061,7 +2163,7 @@ readiness_check = true
|
||||
enabled = true
|
||||
|
||||
[feature.memory]
|
||||
staging = true
|
||||
staging_tools = true
|
||||
|
||||
[feature.manage_workdir]
|
||||
enabled = true
|
||||
@@ -2074,7 +2176,7 @@ enabled = true
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
let manifest: WorkerManifest = WorkerManifestConfig::builtin_defaults()
|
||||
let manifest: WorkerManifest = WorkerManifestConfig::resolution_defaults()
|
||||
.merge(base)
|
||||
.merge(upper)
|
||||
.merge(WorkerManifestConfig {
|
||||
@@ -2098,8 +2200,8 @@ enabled = true
|
||||
})
|
||||
.try_into()
|
||||
.unwrap();
|
||||
assert!(manifest.feature.memory.enabled);
|
||||
assert!(manifest.feature.memory.staging);
|
||||
assert!(manifest.feature.memory.profile.enabled);
|
||||
assert!(manifest.feature.memory.profile.staging_tools);
|
||||
assert!(manifest.feature.manage_workdir.enabled);
|
||||
assert!(manifest.feature.ticket.enabled);
|
||||
assert!(!manifest.feature.ticket.authoring);
|
||||
@@ -2137,7 +2239,7 @@ permission = "write"
|
||||
|
||||
#[test]
|
||||
fn builtin_defaults_populates_worker_limit_defaults() {
|
||||
let cfg = WorkerManifestConfig::builtin_defaults();
|
||||
let cfg = WorkerManifestConfig::resolution_defaults();
|
||||
assert_eq!(
|
||||
cfg.engine.tool_output.default_max_bytes,
|
||||
Some(defaults::TOOL_OUTPUT_MAX_BYTES)
|
||||
@@ -2172,7 +2274,7 @@ permission = "write"
|
||||
},
|
||||
..Default::default()
|
||||
};
|
||||
let merged = WorkerManifestConfig::builtin_defaults().merge(overlay);
|
||||
let merged = WorkerManifestConfig::resolution_defaults().merge(overlay);
|
||||
let manifest: WorkerManifest = merged.try_into().unwrap();
|
||||
assert_eq!(
|
||||
manifest.engine.tool_output.default_max_bytes,
|
||||
|
||||
@@ -93,5 +93,5 @@ pub const COMPACT_RESULT_CONTEXT_MAX_TOKENS: u64 = 60_000;
|
||||
pub const COMPACT_DEFAULT_REFERENCE_COUNT: usize = 5;
|
||||
|
||||
/// Optional maximum extract-worker tool-loop depth. `None` means unlimited.
|
||||
/// See [`crate::MemoryConfig::extract_worker_max_turns`].
|
||||
/// See [`crate::MemoryExtractionProfileConfig::worker_max_turns`].
|
||||
pub const MEMORY_EXTRACT_WORKER_MAX_TURNS: Option<u32> = Some(8);
|
||||
|
||||
+522
-151
@@ -1,3 +1,4 @@
|
||||
mod builtin_profile;
|
||||
mod config;
|
||||
pub mod defaults;
|
||||
mod model;
|
||||
@@ -7,6 +8,11 @@ pub mod plugin;
|
||||
mod profile;
|
||||
mod scope;
|
||||
|
||||
pub use builtin_profile::{
|
||||
BUILTIN_DEFAULT_PROFILE, BUILTIN_PROFILE_CATALOG_ID, BUILTIN_PROFILE_RESOURCES,
|
||||
BuiltinProfileCatalogSnapshot, BuiltinProfileImport, BuiltinProfileResource,
|
||||
builtin_profile_catalog_snapshot, builtin_profile_entrypoints,
|
||||
};
|
||||
pub use config::{
|
||||
CompactionConfigPartial, EngineManifestConfig, FileUploadLimitsPartial,
|
||||
PermissionConfigPartial, ResolveError, SessionConfigPartial, ToolOutputLimitsPartial,
|
||||
@@ -17,10 +23,11 @@ pub use model::{
|
||||
};
|
||||
pub use paths::user_profiles_path;
|
||||
pub use profile::{
|
||||
ProfileDiscovery, ProfileError, ProfileManifestSnapshot, ProfileMetadata, ProfileRegistry,
|
||||
ProfileRegistryEntry, ProfileRegistrySource, ProfileResolveOptions, ProfileResolver,
|
||||
ProfileSelector, ProfileSource, ResolvedProfile, resolve_profile_artifact,
|
||||
resolve_profile_artifact_value,
|
||||
ProfileDiscovery, ProfileError, ProfileExecutionTarget, ProfileManifestSnapshot,
|
||||
ProfileMetadata, ProfileRegistry, ProfileRegistryEntry, ProfileRegistrySource,
|
||||
ProfileResolveOptions, ProfileResolver, ProfileSelector, ProfileSource, ResolvedProfile,
|
||||
WorkspaceAuthorityRequirement, resolve_profile_artifact, resolve_profile_artifact_value,
|
||||
validate_profile_execution_target,
|
||||
};
|
||||
pub use protocol::{Permission, ScopeRule};
|
||||
pub use scope::{DelegationScope, Scope, ScopeError, SharedScope};
|
||||
@@ -40,6 +47,7 @@ use serde::{Deserialize, Serialize};
|
||||
/// part of the manifest — it is the process's `std::env::current_dir()`
|
||||
/// at construction time.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct WorkerManifest {
|
||||
pub worker: WorkerMeta,
|
||||
pub model: ModelManifest,
|
||||
@@ -73,11 +81,6 @@ pub struct WorkerManifest {
|
||||
pub mcp: McpConfig,
|
||||
#[serde(default)]
|
||||
pub compaction: Option<CompactionConfig>,
|
||||
/// Memory subsystem configuration. Presence of `[memory]` configures memory
|
||||
/// storage, extraction, consolidation, and resident injection, but memory
|
||||
/// tools are surfaced only when `[feature.memory].enabled = true`.
|
||||
#[serde(default)]
|
||||
pub memory: Option<MemoryConfig>,
|
||||
/// First-class web tools configuration. Network access remains fail-closed
|
||||
/// under this config; WebSearch/WebFetch schemas are surfaced only when
|
||||
/// `[feature.web].enabled = true`.
|
||||
@@ -102,12 +105,12 @@ pub struct WorkerManifest {
|
||||
/// profile/config data only: they do not carry runtime Worker names, sockets,
|
||||
/// sessions, secrets, or resolved host state. Tool registration still applies
|
||||
/// the normal scope, host-authority, backend, memory, and network checks.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
|
||||
pub struct FeatureConfig {
|
||||
#[serde(default)]
|
||||
pub task: FeatureFlagConfig,
|
||||
#[serde(default)]
|
||||
pub memory: MemoryFeatureConfig,
|
||||
pub memory: ResolvedMemoryFeatureConfig,
|
||||
#[serde(default)]
|
||||
pub web: FeatureFlagConfig,
|
||||
#[serde(default)]
|
||||
@@ -118,6 +121,10 @@ pub struct FeatureConfig {
|
||||
pub flow: FeatureFlagConfig,
|
||||
#[serde(default)]
|
||||
pub worker: WorkerFeatureConfig,
|
||||
/// Privileged read-only discovery of visible Workspace Workers. Backend
|
||||
/// source proof remains required for every listing operation.
|
||||
#[serde(default)]
|
||||
pub workspace_worker_discovery: FeatureFlagConfig,
|
||||
#[serde(default)]
|
||||
pub objective: FeatureFlagConfig,
|
||||
#[serde(default)]
|
||||
@@ -136,12 +143,13 @@ impl Default for FeatureConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
task: FeatureFlagConfig::disabled(),
|
||||
memory: MemoryFeatureConfig::disabled(),
|
||||
memory: ResolvedMemoryFeatureConfig::default(),
|
||||
web: FeatureFlagConfig::disabled(),
|
||||
image: FeatureFlagConfig::disabled(),
|
||||
sub_worker: FeatureFlagConfig::disabled(),
|
||||
flow: FeatureFlagConfig::disabled(),
|
||||
worker: WorkerFeatureConfig::disabled(),
|
||||
workspace_worker_discovery: FeatureFlagConfig::disabled(),
|
||||
objective: FeatureFlagConfig::disabled(),
|
||||
manage_workdir: FeatureFlagConfig::disabled(),
|
||||
ticket: TicketFeatureConfig::default(),
|
||||
@@ -210,34 +218,139 @@ const fn default_true() -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub struct MemoryFeatureConfig {
|
||||
#[serde(default)]
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
|
||||
#[serde(default, deny_unknown_fields)]
|
||||
pub struct MemoryFeatureProfileConfig {
|
||||
pub enabled: bool,
|
||||
/// Exposes Memory staging queue tools in addition to normal Memory CRUD/query tools.
|
||||
#[serde(default)]
|
||||
pub staging: bool,
|
||||
pub staging_tools: bool,
|
||||
pub resident: MemoryResidentProfileConfig,
|
||||
pub extraction: MemoryExtractionProfileConfig,
|
||||
pub consolidation: MemoryConsolidationProfileConfig,
|
||||
}
|
||||
|
||||
impl MemoryFeatureConfig {
|
||||
pub const fn disabled() -> Self {
|
||||
Self {
|
||||
enabled: false,
|
||||
staging: false,
|
||||
}
|
||||
impl MemoryFeatureProfileConfig {
|
||||
pub fn disabled() -> Self {
|
||||
Self::default()
|
||||
}
|
||||
|
||||
pub const fn enabled() -> Self {
|
||||
pub fn enabled() -> Self {
|
||||
Self {
|
||||
enabled: true,
|
||||
staging: false,
|
||||
..Self::default()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for MemoryFeatureConfig {
|
||||
impl Default for MemoryFeatureProfileConfig {
|
||||
fn default() -> Self {
|
||||
Self::disabled()
|
||||
Self {
|
||||
enabled: false,
|
||||
staging_tools: false,
|
||||
resident: MemoryResidentProfileConfig::default(),
|
||||
extraction: MemoryExtractionProfileConfig::default(),
|
||||
consolidation: MemoryConsolidationProfileConfig::default(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
|
||||
#[serde(default, deny_unknown_fields)]
|
||||
pub struct MemoryResidentProfileConfig {
|
||||
pub inject_summary: bool,
|
||||
}
|
||||
|
||||
impl Default for MemoryResidentProfileConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
inject_summary: true,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
|
||||
#[serde(default, deny_unknown_fields)]
|
||||
pub struct MemoryExtractionProfileConfig {
|
||||
pub enabled: bool,
|
||||
pub model: Option<ModelManifest>,
|
||||
pub threshold: Option<u64>,
|
||||
pub worker_max_turns: Option<u32>,
|
||||
}
|
||||
|
||||
impl Default for MemoryExtractionProfileConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
enabled: true,
|
||||
model: None,
|
||||
threshold: Some(50_000),
|
||||
worker_max_turns: defaults::MEMORY_EXTRACT_WORKER_MAX_TURNS,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
#[serde(default, deny_unknown_fields)]
|
||||
pub struct MemoryConsolidationProfileConfig {
|
||||
pub request_enabled: bool,
|
||||
}
|
||||
|
||||
impl Default for MemoryConsolidationProfileConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
request_enabled: true,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Immutable Memory execution configuration persisted in a resolved Worker Manifest.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)]
|
||||
#[serde(default, deny_unknown_fields)]
|
||||
pub struct ResolvedMemoryFeatureConfig {
|
||||
pub profile: MemoryFeatureProfileConfig,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub workspace_settings: Option<WorkspaceMemorySettingsSnapshot>,
|
||||
}
|
||||
|
||||
impl ResolvedMemoryFeatureConfig {
|
||||
pub fn enabled(&self) -> bool {
|
||||
self.profile.enabled
|
||||
}
|
||||
|
||||
pub fn bind_workspace_settings(
|
||||
&mut self,
|
||||
settings: WorkspaceMemorySettingsSnapshot,
|
||||
) -> Result<(), &'static str> {
|
||||
if !self.profile.enabled {
|
||||
if self.workspace_settings.is_some() {
|
||||
return Err("disabled Memory feature must not carry Workspace settings");
|
||||
}
|
||||
return Ok(());
|
||||
}
|
||||
if self.workspace_settings.is_some() {
|
||||
return Err("memory Workspace settings are already bound");
|
||||
}
|
||||
self.workspace_settings = Some(settings);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn workspace_settings(&self) -> Option<WorkspaceMemorySettingsSnapshot> {
|
||||
self.workspace_settings.clone()
|
||||
}
|
||||
|
||||
pub fn validate_execution(&self) -> Result<(), &'static str> {
|
||||
if self.profile.enabled && self.workspace_settings.is_none() {
|
||||
return Err("enabled Memory feature requires trusted Workspace settings");
|
||||
}
|
||||
if !self.profile.enabled && self.workspace_settings.is_some() {
|
||||
return Err("disabled Memory feature must not carry Workspace settings");
|
||||
}
|
||||
if let Some(settings) = &self.workspace_settings
|
||||
&& (settings.settings_revision == 0
|
||||
|| !is_normalized_workspace_memory_language(&settings.language))
|
||||
{
|
||||
return Err("Memory Workspace settings snapshot metadata is invalid");
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -472,98 +585,6 @@ pub struct WorkspaceMemorySettingsSnapshot {
|
||||
pub language: String,
|
||||
}
|
||||
|
||||
/// Memory subsystem configuration. Presence in the manifest enables
|
||||
/// memory; `workspace_root` pins the memory workspace explicitly. When it
|
||||
/// is absent, memory resolution searches upward from the Worker's pwd for a
|
||||
/// `.yoi/memory` marker rather than treating `.yoi` project records alone
|
||||
/// as a memory root.
|
||||
///
|
||||
/// All fields are `Option`; defaults are applied at the consumer
|
||||
/// (`.unwrap_or(defaults::...)`). This keeps cascade `merge` simple
|
||||
/// (`upper.x.or(self.x)`) without a separate partial/resolved split.
|
||||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||||
pub struct MemoryConfig {
|
||||
/// Override for the memory workspace root. When `None`, consumers resolve
|
||||
/// the root from their default path and ancestor `.yoi/memory` markers.
|
||||
/// When set, must be an absolute path.
|
||||
#[serde(default)]
|
||||
pub workspace_root: Option<PathBuf>,
|
||||
/// Maximum number of records returned by `MemoryQuery` /
|
||||
/// `MemoryQuery` per call. `None` ⇒ tool default (20).
|
||||
#[serde(default)]
|
||||
pub query_result_limit: Option<usize>,
|
||||
/// Lines of context before and after each match in query excerpts.
|
||||
/// Ignored when the request omits `query`. `None` ⇒ tool default (3).
|
||||
#[serde(default)]
|
||||
pub query_excerpt_lines: Option<usize>,
|
||||
/// Whether the body of `memory/summary.md` is exposed in the resident
|
||||
/// system-prompt section. `None` ⇒ enabled.
|
||||
#[serde(default)]
|
||||
pub inject_summary: Option<bool>,
|
||||
/// Workspace that owns the bound Memory settings revision.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub workspace_id: Option<String>,
|
||||
/// Monotonic revision of the bound Workspace Memory settings.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub settings_revision: Option<u64>,
|
||||
/// Language from the bound Workspace Memory settings revision.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub language: Option<String>,
|
||||
/// Optional model for the extract worker. When `None`,
|
||||
/// the main engine model is cloned via `clone_boxed()`. Lightweight
|
||||
/// reasoning-capable models (Haiku / 4o-mini / Flash class) are
|
||||
/// recommended.
|
||||
#[serde(default)]
|
||||
pub extract_model: Option<ModelManifest>,
|
||||
/// Cumulative input-token threshold (since the last extract pointer)
|
||||
/// that triggers an extract run. `None` disables the extract trigger
|
||||
/// entirely; memory tools and resident injection still work, only
|
||||
/// the auto-extract trigger is dormant.
|
||||
#[serde(default)]
|
||||
pub extract_threshold: Option<u64>,
|
||||
/// Optional maximum extract-worker tool-loop depth. `None` leaves
|
||||
/// the worker unlimited; the default bounds runaway short-context
|
||||
/// loops. Falls through to
|
||||
/// [`defaults::MEMORY_EXTRACT_WORKER_MAX_TURNS`] when unset.
|
||||
#[serde(default)]
|
||||
pub extract_worker_max_turns: Option<u32>,
|
||||
/// Optional model for the consolidation worker. When
|
||||
/// `None`, the main engine model is cloned via `clone_boxed()`.
|
||||
/// Reasoning-class models are recommended.
|
||||
#[serde(default)]
|
||||
pub consolidation_model: Option<ModelManifest>,
|
||||
/// Consolidation trigger: file-count threshold of `_staging/`. The
|
||||
/// consolidation run fires when the staging directory has at least
|
||||
/// this many entries. Either threshold reaching its limit fires
|
||||
/// consolidation (logical OR). `None` for both thresholds ⇒
|
||||
/// consolidation disabled.
|
||||
#[serde(default)]
|
||||
pub consolidation_threshold_files: Option<usize>,
|
||||
/// Consolidation trigger: byte-size threshold across all `_staging/`
|
||||
/// entries. Either threshold reaching its limit fires consolidation.
|
||||
/// `None` for both thresholds ⇒ consolidation disabled.
|
||||
#[serde(default)]
|
||||
pub consolidation_threshold_bytes: Option<u64>,
|
||||
}
|
||||
|
||||
impl MemoryConfig {
|
||||
/// Replace any untrusted manifest values with a trusted Workspace snapshot.
|
||||
pub fn bind_workspace_settings(&mut self, snapshot: &WorkspaceMemorySettingsSnapshot) {
|
||||
self.workspace_id = Some(snapshot.workspace_id.clone());
|
||||
self.settings_revision = Some(snapshot.settings_revision);
|
||||
self.language = Some(snapshot.language.clone());
|
||||
}
|
||||
|
||||
/// Return the complete bound Workspace settings snapshot, if every field is present.
|
||||
pub fn workspace_settings(&self) -> Option<WorkspaceMemorySettingsSnapshot> {
|
||||
Some(WorkspaceMemorySettingsSnapshot {
|
||||
workspace_id: self.workspace_id.clone()?,
|
||||
settings_revision: self.settings_revision?,
|
||||
language: self.language.clone()?,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Worker metadata.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct WorkerMeta {
|
||||
@@ -919,6 +940,12 @@ impl Default for CompactionConfig {
|
||||
}
|
||||
|
||||
impl WorkerManifest {
|
||||
pub fn requires_persisted_execution_snapshot(&self) -> bool {
|
||||
self.profile.is_some()
|
||||
|| self.plugins.has_resolved_plan()
|
||||
|| self.feature.memory.workspace_settings.is_some()
|
||||
}
|
||||
|
||||
/// Parse a manifest from a TOML string.
|
||||
pub fn from_toml(s: &str) -> Result<Self, toml::de::Error> {
|
||||
config::reject_removed_manifest_fields(s)?;
|
||||
@@ -929,6 +956,212 @@ impl WorkerManifest {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Default, Deserialize)]
|
||||
#[serde(default, deny_unknown_fields)]
|
||||
struct LegacyMemoryFeatureConfig {
|
||||
enabled: bool,
|
||||
staging: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Default, Deserialize)]
|
||||
#[serde(default, deny_unknown_fields)]
|
||||
struct LegacyMemoryConfig {
|
||||
#[serde(rename = "workspace_root")]
|
||||
_workspace_root: Option<PathBuf>,
|
||||
#[serde(rename = "query_result_limit")]
|
||||
_query_result_limit: Option<usize>,
|
||||
#[serde(rename = "query_excerpt_lines")]
|
||||
_query_excerpt_lines: Option<usize>,
|
||||
inject_summary: Option<bool>,
|
||||
workspace_id: Option<String>,
|
||||
settings_revision: Option<u64>,
|
||||
language: Option<String>,
|
||||
extract_model: Option<ModelManifest>,
|
||||
extract_threshold: Option<u64>,
|
||||
extract_worker_max_turns: Option<u32>,
|
||||
consolidation_model: Option<ModelManifest>,
|
||||
consolidation_threshold_files: Option<usize>,
|
||||
consolidation_threshold_bytes: Option<u64>,
|
||||
}
|
||||
|
||||
const RESOLVED_MANIFEST_SNAPSHOT_SCHEMA_VERSION: u64 = 2;
|
||||
|
||||
/// Serialize a resolved Worker Manifest for durable Worker-specific storage.
|
||||
pub fn write_persisted_worker_manifest_snapshot(
|
||||
manifest: &WorkerManifest,
|
||||
) -> Result<serde_json::Value, serde_json::Error> {
|
||||
Ok(serde_json::json!({
|
||||
"schema_version": RESOLVED_MANIFEST_SNAPSHOT_SCHEMA_VERSION,
|
||||
"manifest": serde_json::to_value(manifest)?,
|
||||
}))
|
||||
}
|
||||
|
||||
/// Read a durable resolved Worker Manifest through the versioned compatibility
|
||||
/// boundary. Runtime code must not deserialize persisted snapshots directly.
|
||||
pub fn read_persisted_worker_manifest_snapshot(
|
||||
snapshot: serde_json::Value,
|
||||
) -> Result<WorkerManifest, serde_json::Error> {
|
||||
let object = snapshot.as_object().ok_or_else(|| {
|
||||
serde_json::Error::io(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
"resolved Worker manifest snapshot must be an object",
|
||||
))
|
||||
})?;
|
||||
if let Some(version) = object.get("schema_version") {
|
||||
let version = version.as_u64().ok_or_else(|| {
|
||||
serde_json::Error::io(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
"resolved Worker manifest snapshot schema_version must be an integer",
|
||||
))
|
||||
})?;
|
||||
if version != RESOLVED_MANIFEST_SNAPSHOT_SCHEMA_VERSION {
|
||||
return Err(serde_json::Error::io(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
format!("unsupported resolved Worker manifest snapshot schema version {version}"),
|
||||
)));
|
||||
}
|
||||
if object.len() != 2 {
|
||||
return Err(serde_json::Error::io(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
"resolved Worker manifest snapshot contains unknown fields",
|
||||
)));
|
||||
}
|
||||
let manifest = object.get("manifest").cloned().ok_or_else(|| {
|
||||
serde_json::Error::io(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
"resolved Worker manifest snapshot is missing manifest",
|
||||
))
|
||||
})?;
|
||||
if manifest
|
||||
.as_object()
|
||||
.is_some_and(|manifest| manifest.contains_key("memory"))
|
||||
{
|
||||
return Err(serde_json::Error::io(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
"current resolved Worker manifest contains removed top-level memory authority",
|
||||
)));
|
||||
}
|
||||
return validate_persisted_worker_manifest(serde_json::from_value(manifest)?);
|
||||
}
|
||||
|
||||
migrate_legacy_resolved_manifest_snapshot(snapshot)
|
||||
}
|
||||
|
||||
fn validate_persisted_worker_manifest(
|
||||
manifest: WorkerManifest,
|
||||
) -> Result<WorkerManifest, serde_json::Error> {
|
||||
manifest
|
||||
.feature
|
||||
.memory
|
||||
.validate_execution()
|
||||
.map_err(|message| {
|
||||
serde_json::Error::io(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
message,
|
||||
))
|
||||
})?;
|
||||
Ok(manifest)
|
||||
}
|
||||
|
||||
fn migrate_legacy_resolved_manifest_snapshot(
|
||||
mut snapshot: serde_json::Value,
|
||||
) -> Result<WorkerManifest, serde_json::Error> {
|
||||
let root = snapshot.as_object_mut().ok_or_else(|| {
|
||||
serde_json::Error::io(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
"legacy resolved Worker manifest snapshot must be an object",
|
||||
))
|
||||
})?;
|
||||
let legacy_memory = root.remove("memory");
|
||||
let feature = root
|
||||
.entry("feature")
|
||||
.or_insert_with(|| serde_json::json!({}))
|
||||
.as_object_mut()
|
||||
.ok_or_else(|| {
|
||||
serde_json::Error::io(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
"legacy resolved Worker manifest feature must be an object",
|
||||
))
|
||||
})?;
|
||||
let legacy_feature_memory: LegacyMemoryFeatureConfig = serde_json::from_value(
|
||||
feature
|
||||
.remove("memory")
|
||||
.unwrap_or_else(|| serde_json::json!({})),
|
||||
)?;
|
||||
let enabled = legacy_feature_memory.enabled;
|
||||
let staging_tools = legacy_feature_memory.staging;
|
||||
|
||||
let legacy_memory: LegacyMemoryConfig =
|
||||
serde_json::from_value(legacy_memory.unwrap_or_else(|| serde_json::json!({})))?;
|
||||
let mut workspace_settings = match (
|
||||
legacy_memory.workspace_id,
|
||||
legacy_memory.settings_revision,
|
||||
legacy_memory.language,
|
||||
) {
|
||||
(Some(workspace_id), Some(settings_revision), Some(language)) => Some(serde_json::json!({
|
||||
"workspace_id": workspace_id,
|
||||
"settings_revision": settings_revision,
|
||||
"language": language,
|
||||
})),
|
||||
(None, None, None) => None,
|
||||
_ => {
|
||||
return Err(serde_json::Error::io(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
"legacy resolved Worker manifest contains a partial Memory settings snapshot",
|
||||
)));
|
||||
}
|
||||
};
|
||||
if !enabled {
|
||||
workspace_settings = None;
|
||||
}
|
||||
let extraction_enabled = legacy_memory.extract_threshold.is_some();
|
||||
if legacy_memory.consolidation_model.is_some() {
|
||||
return Err(serde_json::Error::io(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
"legacy resolved Worker manifest uses a Worker-owned consolidation model that cannot be migrated to Backend authority",
|
||||
)));
|
||||
}
|
||||
let consolidation_enabled = match (
|
||||
legacy_memory.consolidation_threshold_files,
|
||||
legacy_memory.consolidation_threshold_bytes,
|
||||
) {
|
||||
(None, None) => false,
|
||||
(Some(5), Some(50_000)) => true,
|
||||
_ => {
|
||||
return Err(serde_json::Error::io(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
"legacy resolved Worker manifest uses custom consolidation thresholds that cannot be migrated to Backend policy",
|
||||
)));
|
||||
}
|
||||
};
|
||||
let mut resolved = serde_json::json!({
|
||||
"profile": {
|
||||
"enabled": enabled,
|
||||
"staging_tools": staging_tools,
|
||||
"resident": {
|
||||
"inject_summary": legacy_memory.inject_summary.unwrap_or(true),
|
||||
},
|
||||
"extraction": {
|
||||
"enabled": extraction_enabled,
|
||||
"model": serde_json::to_value(legacy_memory.extract_model)?,
|
||||
"threshold": legacy_memory.extract_threshold,
|
||||
"worker_max_turns": legacy_memory.extract_worker_max_turns,
|
||||
},
|
||||
"consolidation": {
|
||||
"request_enabled": consolidation_enabled,
|
||||
},
|
||||
},
|
||||
});
|
||||
if let Some(workspace_settings) = workspace_settings {
|
||||
resolved
|
||||
.as_object_mut()
|
||||
.expect("resolved Memory config is an object")
|
||||
.insert("workspace_settings".to_string(), workspace_settings);
|
||||
}
|
||||
feature.insert("memory".to_string(), resolved);
|
||||
validate_persisted_worker_manifest(serde_json::from_value(snapshot)?)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
@@ -1234,36 +1467,182 @@ model_id = "claude-sonnet-4-20250514"
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn omitted_memory_is_none() {
|
||||
fn omitted_memory_feature_is_disabled() {
|
||||
let manifest = WorkerManifest::from_toml(MINIMAL_REQUIRED).unwrap();
|
||||
assert!(manifest.memory.is_none());
|
||||
assert!(!manifest.feature.memory.profile.enabled);
|
||||
assert!(manifest.feature.memory.workspace_settings.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn empty_memory_section_enables_with_default_root() {
|
||||
let toml = format!("{MINIMAL_REQUIRED}\n[memory]\n");
|
||||
fn resolved_memory_feature_requires_nested_profile_and_trusted_snapshot() {
|
||||
let toml = format!(
|
||||
"{MINIMAL_REQUIRED}\n\
|
||||
[feature.memory.profile]\n\
|
||||
enabled = true\n\
|
||||
staging_tools = false\n\n\
|
||||
[feature.memory.profile.resident]\n\
|
||||
inject_summary = false\n\n\
|
||||
[feature.memory.profile.extraction]\n\
|
||||
enabled = true\n\
|
||||
threshold = 42000\n\
|
||||
worker_max_turns = 2\n\n\
|
||||
[feature.memory.workspace_settings]\n\
|
||||
workspace_id = \"workspace-1\"\n\
|
||||
settings_revision = 7\n\
|
||||
language = \"日本語\"\n"
|
||||
);
|
||||
let manifest = WorkerManifest::from_toml(&toml).unwrap();
|
||||
let mem = manifest.memory.expect("memory section parsed");
|
||||
assert!(mem.workspace_root.is_none());
|
||||
assert_eq!(mem.inject_summary, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn memory_section_with_inject_summary_false() {
|
||||
let toml = format!("{MINIMAL_REQUIRED}\n[memory]\ninject_summary = false\n");
|
||||
let manifest = WorkerManifest::from_toml(&toml).unwrap();
|
||||
let mem = manifest.memory.unwrap();
|
||||
assert_eq!(mem.inject_summary, Some(false));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn memory_section_with_explicit_root() {
|
||||
let toml = format!("{MINIMAL_REQUIRED}\n[memory]\nworkspace_root = \"/some/where\"\n");
|
||||
let manifest = WorkerManifest::from_toml(&toml).unwrap();
|
||||
let mem = manifest.memory.unwrap();
|
||||
assert!(manifest.feature.memory.profile.enabled);
|
||||
assert!(!manifest.feature.memory.profile.resident.inject_summary);
|
||||
assert_eq!(
|
||||
mem.workspace_root.unwrap(),
|
||||
std::path::PathBuf::from("/some/where")
|
||||
manifest.feature.memory.profile.extraction.threshold,
|
||||
Some(42_000)
|
||||
);
|
||||
assert_eq!(
|
||||
manifest
|
||||
.feature
|
||||
.memory
|
||||
.workspace_settings()
|
||||
.unwrap()
|
||||
.language,
|
||||
"日本語"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolved_memory_execution_validation_fails_closed() {
|
||||
let snapshot = WorkspaceMemorySettingsSnapshot {
|
||||
workspace_id: "workspace-1".to_string(),
|
||||
settings_revision: 1,
|
||||
language: "English".to_string(),
|
||||
};
|
||||
let mut enabled = ResolvedMemoryFeatureConfig::default();
|
||||
enabled.profile.enabled = true;
|
||||
assert!(enabled.validate_execution().is_err());
|
||||
enabled.bind_workspace_settings(snapshot.clone()).unwrap();
|
||||
assert!(enabled.validate_execution().is_ok());
|
||||
|
||||
let mut disabled = ResolvedMemoryFeatureConfig::default();
|
||||
disabled.workspace_settings = Some(snapshot.clone());
|
||||
assert!(disabled.validate_execution().is_err());
|
||||
assert!(disabled.bind_workspace_settings(snapshot).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn current_manifest_rejects_legacy_top_level_memory_authority() {
|
||||
let toml = format!("{MINIMAL_REQUIRED}\n[memory]\nlanguage = \"Japanese\"\n");
|
||||
assert!(WorkerManifest::from_toml(&toml).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn persisted_manifest_adapter_migrates_legacy_memory_authority() {
|
||||
let mut manifest =
|
||||
serde_json::to_value(WorkerManifest::from_toml(MINIMAL_REQUIRED).unwrap()).unwrap();
|
||||
manifest["feature"]["memory"] = serde_json::json!({
|
||||
"enabled": true,
|
||||
"staging": true,
|
||||
});
|
||||
manifest["memory"] = serde_json::json!({
|
||||
"workspace_root": "/discarded",
|
||||
"query_result_limit": 999,
|
||||
"inject_summary": false,
|
||||
"workspace_id": "workspace-1",
|
||||
"settings_revision": 9,
|
||||
"language": "Français",
|
||||
"extract_threshold": 1234,
|
||||
"extract_worker_max_turns": 3,
|
||||
"consolidation_threshold_files": 5,
|
||||
"consolidation_threshold_bytes": 50000,
|
||||
});
|
||||
|
||||
let migrated = read_persisted_worker_manifest_snapshot(manifest).unwrap();
|
||||
assert!(migrated.feature.memory.profile.enabled);
|
||||
assert!(migrated.feature.memory.profile.staging_tools);
|
||||
assert!(!migrated.feature.memory.profile.resident.inject_summary);
|
||||
assert_eq!(
|
||||
migrated.feature.memory.profile.extraction.threshold,
|
||||
Some(1234)
|
||||
);
|
||||
assert!(
|
||||
migrated
|
||||
.feature
|
||||
.memory
|
||||
.profile
|
||||
.consolidation
|
||||
.request_enabled
|
||||
);
|
||||
assert_eq!(
|
||||
migrated
|
||||
.feature
|
||||
.memory
|
||||
.workspace_settings()
|
||||
.unwrap()
|
||||
.language,
|
||||
"Français"
|
||||
);
|
||||
let current = write_persisted_worker_manifest_snapshot(&migrated).unwrap();
|
||||
assert_eq!(current["schema_version"], 2);
|
||||
assert!(current["manifest"].get("memory").is_none());
|
||||
|
||||
let mut disabled =
|
||||
serde_json::to_value(WorkerManifest::from_toml(MINIMAL_REQUIRED).unwrap()).unwrap();
|
||||
disabled["feature"]["memory"] = serde_json::json!({ "enabled": false });
|
||||
disabled["memory"] = serde_json::json!({
|
||||
"workspace_id": "workspace-1",
|
||||
"settings_revision": 9,
|
||||
"language": "Français",
|
||||
});
|
||||
let disabled = read_persisted_worker_manifest_snapshot(disabled).unwrap();
|
||||
assert!(!disabled.feature.memory.profile.enabled);
|
||||
assert!(disabled.feature.memory.workspace_settings.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn persisted_manifest_adapter_rejects_mixed_or_future_authority() {
|
||||
let manifest =
|
||||
serde_json::to_value(WorkerManifest::from_toml(MINIMAL_REQUIRED).unwrap()).unwrap();
|
||||
let mut mixed = manifest.clone();
|
||||
mixed["feature"]["memory"] = serde_json::json!({ "enabled": true, "profile": {} });
|
||||
mixed["memory"] = serde_json::json!({});
|
||||
assert!(read_persisted_worker_manifest_snapshot(mixed).is_err());
|
||||
|
||||
let mut custom_policy =
|
||||
serde_json::to_value(WorkerManifest::from_toml(MINIMAL_REQUIRED).unwrap()).unwrap();
|
||||
custom_policy["feature"]["memory"] = serde_json::json!({ "enabled": true });
|
||||
custom_policy["memory"] = serde_json::json!({
|
||||
"workspace_id": "workspace-1",
|
||||
"settings_revision": 1,
|
||||
"language": "English",
|
||||
"consolidation_threshold_files": 99,
|
||||
"consolidation_threshold_bytes": 50000,
|
||||
});
|
||||
assert!(read_persisted_worker_manifest_snapshot(custom_policy).is_err());
|
||||
|
||||
let current = WorkerManifest::from_toml(MINIMAL_REQUIRED).unwrap();
|
||||
let mut current = write_persisted_worker_manifest_snapshot(¤t).unwrap();
|
||||
current["manifest"]["memory"] = serde_json::json!({
|
||||
"workspace_id": "workspace-1",
|
||||
"settings_revision": 1,
|
||||
"language": "English",
|
||||
});
|
||||
assert!(read_persisted_worker_manifest_snapshot(current).is_err());
|
||||
|
||||
let mut missing_settings = WorkerManifest::from_toml(MINIMAL_REQUIRED).unwrap();
|
||||
missing_settings.feature.memory.profile.enabled = true;
|
||||
let missing_settings = write_persisted_worker_manifest_snapshot(&missing_settings).unwrap();
|
||||
assert!(read_persisted_worker_manifest_snapshot(missing_settings).is_err());
|
||||
|
||||
let mut malformed_legacy = manifest.clone();
|
||||
malformed_legacy["feature"]["memory"] = serde_json::json!({ "enabled": "yes" });
|
||||
malformed_legacy["memory"] = serde_json::json!({ "unknown": true });
|
||||
assert!(read_persisted_worker_manifest_snapshot(malformed_legacy).is_err());
|
||||
|
||||
assert!(
|
||||
read_persisted_worker_manifest_snapshot(serde_json::json!({
|
||||
"schema_version": 3,
|
||||
"manifest": manifest,
|
||||
}))
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
|
||||
@@ -1279,14 +1658,6 @@ model_id = "claude-sonnet-4-20250514"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn memory_section_with_language() {
|
||||
let toml = format!("{MINIMAL_REQUIRED}\n[memory]\nlanguage = \"Japanese\"\n");
|
||||
let manifest = WorkerManifest::from_toml(&toml).unwrap();
|
||||
let mem = manifest.memory.unwrap();
|
||||
assert_eq!(mem.language.as_deref(), Some("Japanese"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reject_unknown_scheme() {
|
||||
let toml =
|
||||
|
||||
+303
-316
@@ -6,62 +6,28 @@
|
||||
//! from launch context.
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::BTreeMap;
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::fmt;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use crate::builtin_profile::{
|
||||
BUILTIN_DEFAULT_PROFILE, builtin_profile_catalog_snapshot, builtin_profile_entrypoints,
|
||||
resolve_builtin_profile_artifact,
|
||||
};
|
||||
use crate::config::{
|
||||
CompactionConfigPartial, FeatureConfigPartial, PermissionConfigPartial, SessionConfigPartial,
|
||||
};
|
||||
use crate::model::{AuthRef, ModelManifest};
|
||||
use crate::plugin::PluginConfig;
|
||||
use crate::{
|
||||
EngineManifestConfig, McpConfig, McpStdioCwdPolicy, MemoryConfig, Permission, ResolveError,
|
||||
ScopeConfig, ScopeRule, SkillsConfig, WebConfig, WorkerManifest, WorkerManifestConfig,
|
||||
WorkerMetaConfig, paths,
|
||||
EngineManifestConfig, McpConfig, McpStdioCwdPolicy, Permission, ResolveError, ScopeConfig,
|
||||
ScopeRule, SkillsConfig, WebConfig, WorkerManifest, WorkerManifestConfig, WorkerMetaConfig,
|
||||
paths,
|
||||
};
|
||||
|
||||
const PROFILE_FORMAT_V1: &str = "yoi.profile.v1";
|
||||
const BUILTIN_MODEL_CATALOG: &str = include_str!("../../../resources/models/builtin.toml");
|
||||
|
||||
struct BuiltinProfile {
|
||||
name: &'static str,
|
||||
label: &'static str,
|
||||
description: &'static str,
|
||||
}
|
||||
|
||||
const BUILTIN_PROFILES: &[BuiltinProfile] = &[
|
||||
BuiltinProfile {
|
||||
name: "companion",
|
||||
label: "builtin:companion",
|
||||
description: "Bundled Companion role profile",
|
||||
},
|
||||
BuiltinProfile {
|
||||
name: "intake",
|
||||
label: "builtin:intake",
|
||||
description: "Bundled Intake role profile",
|
||||
},
|
||||
BuiltinProfile {
|
||||
name: "orchestrator",
|
||||
label: "builtin:orchestrator",
|
||||
description: "Bundled Orchestrator role profile",
|
||||
},
|
||||
BuiltinProfile {
|
||||
name: "coder",
|
||||
label: "builtin:coder",
|
||||
description: "Bundled Coder role profile",
|
||||
},
|
||||
BuiltinProfile {
|
||||
name: "reviewer",
|
||||
label: "builtin:reviewer",
|
||||
description: "Bundled Reviewer role profile",
|
||||
},
|
||||
BuiltinProfile {
|
||||
name: "memory-consolidation",
|
||||
label: "builtin:memory-consolidation",
|
||||
description: "Bundled Memory staging consolidation profile",
|
||||
},
|
||||
];
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum ProfileRegistrySource {
|
||||
@@ -159,6 +125,108 @@ impl ProfileSelector {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum ProfileExecutionTarget {
|
||||
Workspace,
|
||||
Standalone,
|
||||
}
|
||||
|
||||
impl fmt::Display for ProfileExecutionTarget {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
Self::Workspace => formatter.write_str("workspace"),
|
||||
Self::Standalone => formatter.write_str("standalone"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
|
||||
pub enum WorkspaceAuthorityRequirement {
|
||||
Flow,
|
||||
ManageWorkdir,
|
||||
Memory,
|
||||
MergeRequest,
|
||||
Objective,
|
||||
Orchestration,
|
||||
Plugins,
|
||||
Ticket,
|
||||
Worker,
|
||||
}
|
||||
|
||||
impl fmt::Display for WorkspaceAuthorityRequirement {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
Self::Flow => formatter.write_str("feature.flow"),
|
||||
Self::ManageWorkdir => formatter.write_str("feature.manage_workdir"),
|
||||
Self::Memory => formatter.write_str("feature.memory"),
|
||||
Self::MergeRequest => formatter.write_str("feature.merge_request"),
|
||||
Self::Objective => formatter.write_str("feature.objective"),
|
||||
Self::Orchestration => formatter.write_str("feature.orchestration"),
|
||||
Self::Plugins => formatter.write_str("feature.plugins or plugin packages"),
|
||||
Self::Ticket => formatter.write_str("feature.ticket"),
|
||||
Self::Worker => formatter.write_str("feature.worker"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn validate_profile_execution_target(
|
||||
manifest: &WorkerManifest,
|
||||
target: ProfileExecutionTarget,
|
||||
) -> Result<(), ProfileError> {
|
||||
if target == ProfileExecutionTarget::Workspace {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let feature = &manifest.feature;
|
||||
let mut requirements = BTreeSet::new();
|
||||
if feature.flow.enabled {
|
||||
requirements.insert(WorkspaceAuthorityRequirement::Flow);
|
||||
}
|
||||
if feature.manage_workdir.enabled {
|
||||
requirements.insert(WorkspaceAuthorityRequirement::ManageWorkdir);
|
||||
}
|
||||
if feature.memory.profile.enabled || feature.memory.profile.staging_tools {
|
||||
requirements.insert(WorkspaceAuthorityRequirement::Memory);
|
||||
}
|
||||
if feature.merge_request.show
|
||||
|| feature.merge_request.open
|
||||
|| feature.merge_request.review
|
||||
|| feature.merge_request.readiness_check
|
||||
|| feature.merge_request.complete
|
||||
{
|
||||
requirements.insert(WorkspaceAuthorityRequirement::MergeRequest);
|
||||
}
|
||||
if feature.objective.enabled {
|
||||
requirements.insert(WorkspaceAuthorityRequirement::Objective);
|
||||
}
|
||||
if feature.orchestration.enabled {
|
||||
requirements.insert(WorkspaceAuthorityRequirement::Orchestration);
|
||||
}
|
||||
if feature.plugins.enabled || !manifest.plugins.is_empty() {
|
||||
requirements.insert(WorkspaceAuthorityRequirement::Plugins);
|
||||
}
|
||||
if feature.ticket.enabled
|
||||
|| feature.ticket.authoring
|
||||
|| feature.ticket.thread
|
||||
|| feature.ticket.intake
|
||||
|| feature.ticket.workflow
|
||||
{
|
||||
requirements.insert(WorkspaceAuthorityRequirement::Ticket);
|
||||
}
|
||||
if feature.worker.enabled {
|
||||
requirements.insert(WorkspaceAuthorityRequirement::Worker);
|
||||
}
|
||||
|
||||
if requirements.is_empty() {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(ProfileError::UnsupportedExecutionTarget {
|
||||
target,
|
||||
requirements: requirements.into_iter().collect(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(tag = "kind", rename_all = "snake_case")]
|
||||
pub enum ProfileSource {
|
||||
@@ -217,13 +285,14 @@ impl ProfileRegistryEntry {
|
||||
source: ProfileRegistrySource,
|
||||
name: &'static str,
|
||||
label: &'static str,
|
||||
provenance: String,
|
||||
description: Option<String>,
|
||||
) -> Self {
|
||||
Self {
|
||||
source,
|
||||
name: name.to_string(),
|
||||
path: None,
|
||||
provenance: label.to_string(),
|
||||
provenance,
|
||||
description,
|
||||
is_default: false,
|
||||
artifact: ProfileRegistryArtifact::Builtin { label },
|
||||
@@ -321,12 +390,16 @@ pub struct ProfileDiscovery {
|
||||
}
|
||||
|
||||
impl ProfileDiscovery {
|
||||
pub fn for_cwd(_cwd: &Path) -> Self {
|
||||
pub fn user_settings() -> Self {
|
||||
Self {
|
||||
user_config: paths::user_profiles_path(),
|
||||
project_config: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn for_cwd(_cwd: &Path) -> Self {
|
||||
Self::user_settings()
|
||||
}
|
||||
pub fn with_sources(user_config: Option<PathBuf>, project_config: Option<PathBuf>) -> Self {
|
||||
Self {
|
||||
user_config,
|
||||
@@ -412,15 +485,22 @@ impl ProfileResolver {
|
||||
options,
|
||||
),
|
||||
ProfileSelector::Named { .. } | ProfileSelector::Default => {
|
||||
let cwd = std::env::current_dir().map_err(|source| ProfileError::CommandIo {
|
||||
path: PathBuf::from("."),
|
||||
source,
|
||||
})?;
|
||||
let registry = ProfileDiscovery::for_cwd(&cwd).discover()?;
|
||||
let registry = ProfileDiscovery::user_settings().discover()?;
|
||||
self.resolve_from_registry(selector, ®istry, options)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn resolve_for_target(
|
||||
&self,
|
||||
selector: &ProfileSelector,
|
||||
options: ProfileResolveOptions,
|
||||
target: ProfileExecutionTarget,
|
||||
) -> Result<ResolvedProfile, ProfileError> {
|
||||
let resolved = self.resolve(selector, options)?;
|
||||
validate_profile_execution_target(&resolved.manifest, target)?;
|
||||
Ok(resolved)
|
||||
}
|
||||
/// Resolve a registry/default selector against an already-discovered
|
||||
/// registry. Callers such as SubWorkerSpawn use this to bind discovery to the
|
||||
/// Worker's cwd instead of the process current directory.
|
||||
@@ -503,7 +583,7 @@ impl ProfileResolver {
|
||||
.as_deref()
|
||||
.unwrap_or_else(|| Path::new(".")),
|
||||
)?;
|
||||
let raw_artifact = builtin_profile_artifact(label).ok_or_else(|| {
|
||||
let raw_artifact = resolve_builtin_profile_artifact(label)?.ok_or_else(|| {
|
||||
ProfileError::InvalidProfile(format!("unknown builtin profile artifact `{label}`"))
|
||||
})?;
|
||||
resolve_profile_value(
|
||||
@@ -562,10 +642,10 @@ fn resolve_profile_value(
|
||||
mcp: profile.mcp,
|
||||
compaction,
|
||||
web: profile.web,
|
||||
memory: profile.memory.map(Into::into),
|
||||
skills: profile.skills,
|
||||
};
|
||||
let config = WorkerManifestConfig::builtin_defaults().merge(config.resolve_paths(profile_dir));
|
||||
let config =
|
||||
WorkerManifestConfig::resolution_defaults().merge(config.resolve_paths(profile_dir));
|
||||
let mut manifest = WorkerManifest::try_from(config).map_err(ProfileError::ManifestResolve)?;
|
||||
manifest.profile = Some(ProfileManifestSnapshot {
|
||||
source: source.clone(),
|
||||
@@ -582,51 +662,6 @@ fn resolve_profile_value(
|
||||
})
|
||||
}
|
||||
|
||||
#[derive(Debug, Default, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct ProfileMemoryConfig {
|
||||
#[serde(default)]
|
||||
workspace_root: Option<PathBuf>,
|
||||
#[serde(default)]
|
||||
query_result_limit: Option<usize>,
|
||||
#[serde(default)]
|
||||
query_excerpt_lines: Option<usize>,
|
||||
#[serde(default)]
|
||||
inject_summary: Option<bool>,
|
||||
#[serde(default)]
|
||||
extract_model: Option<ModelManifest>,
|
||||
#[serde(default)]
|
||||
extract_threshold: Option<u64>,
|
||||
#[serde(default)]
|
||||
extract_worker_max_turns: Option<u32>,
|
||||
#[serde(default)]
|
||||
consolidation_model: Option<ModelManifest>,
|
||||
#[serde(default)]
|
||||
consolidation_threshold_files: Option<usize>,
|
||||
#[serde(default)]
|
||||
consolidation_threshold_bytes: Option<u64>,
|
||||
}
|
||||
|
||||
impl From<ProfileMemoryConfig> for MemoryConfig {
|
||||
fn from(profile: ProfileMemoryConfig) -> Self {
|
||||
Self {
|
||||
workspace_root: profile.workspace_root,
|
||||
query_result_limit: profile.query_result_limit,
|
||||
query_excerpt_lines: profile.query_excerpt_lines,
|
||||
inject_summary: profile.inject_summary,
|
||||
workspace_id: None,
|
||||
settings_revision: None,
|
||||
language: None,
|
||||
extract_model: profile.extract_model,
|
||||
extract_threshold: profile.extract_threshold,
|
||||
extract_worker_max_turns: profile.extract_worker_max_turns,
|
||||
consolidation_model: profile.consolidation_model,
|
||||
consolidation_threshold_files: profile.consolidation_threshold_files,
|
||||
consolidation_threshold_bytes: profile.consolidation_threshold_bytes,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Default, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct ProfileConfig {
|
||||
@@ -657,8 +692,6 @@ struct ProfileConfig {
|
||||
#[serde(default)]
|
||||
web: Option<WebConfig>,
|
||||
#[serde(default)]
|
||||
memory: Option<ProfileMemoryConfig>,
|
||||
#[serde(default)]
|
||||
skills: Option<SkillsConfig>,
|
||||
}
|
||||
|
||||
@@ -759,14 +792,30 @@ fn load_profile_registry_file(
|
||||
}
|
||||
|
||||
fn add_builtin_profiles(registry: &mut ProfileRegistry) {
|
||||
for profile in BUILTIN_PROFILES {
|
||||
let catalog = builtin_profile_catalog_snapshot();
|
||||
let digest = catalog.digest();
|
||||
for profile in builtin_profile_entrypoints() {
|
||||
let label = profile
|
||||
.selector
|
||||
.expect("built-in Profile entrypoint must have a selector");
|
||||
let name = label
|
||||
.strip_prefix("builtin:")
|
||||
.expect("built-in Profile selector must be source-qualified");
|
||||
registry.push_entry(ProfileRegistryEntry::embedded(
|
||||
ProfileRegistrySource::Builtin,
|
||||
profile.name,
|
||||
profile.label,
|
||||
name,
|
||||
label,
|
||||
format!("{}#{digest}", profile.path),
|
||||
Some(profile.description.into()),
|
||||
));
|
||||
}
|
||||
registry.set_default(ProfileDefault {
|
||||
source: Some(ProfileRegistrySource::Builtin),
|
||||
name: BUILTIN_DEFAULT_PROFILE
|
||||
.strip_prefix("builtin:")
|
||||
.expect("built-in default selector must be source-qualified")
|
||||
.to_owned(),
|
||||
});
|
||||
}
|
||||
|
||||
fn parse_profile_ref(raw: &str) -> (Option<ProfileRegistrySource>, String) {
|
||||
@@ -804,201 +853,6 @@ fn read_profile_artifact_file(path: &Path) -> Result<serde_json::Value, ProfileE
|
||||
}
|
||||
}
|
||||
|
||||
fn builtin_profile_artifact(label: &str) -> Option<serde_json::Value> {
|
||||
let mut value = builtin_base_profile_artifact();
|
||||
match label {
|
||||
"builtin:companion" | "companion" => {
|
||||
apply_role_profile(
|
||||
&mut value,
|
||||
"companion",
|
||||
"Workspace companion profile.",
|
||||
"workspace_write",
|
||||
true,
|
||||
true,
|
||||
true,
|
||||
true,
|
||||
);
|
||||
Some(value)
|
||||
}
|
||||
"builtin:intake" | "intake" => {
|
||||
apply_role_profile(
|
||||
&mut value,
|
||||
"intake",
|
||||
"Ticket intake profile.",
|
||||
"workspace_write",
|
||||
true,
|
||||
true,
|
||||
true,
|
||||
false,
|
||||
);
|
||||
Some(value)
|
||||
}
|
||||
"builtin:orchestrator" | "orchestrator" => {
|
||||
apply_role_profile(
|
||||
&mut value,
|
||||
"orchestrator",
|
||||
"Ticket orchestrator profile.",
|
||||
"workspace_write",
|
||||
true,
|
||||
true,
|
||||
true,
|
||||
false,
|
||||
);
|
||||
Some(value)
|
||||
}
|
||||
"builtin:coder" | "coder" => {
|
||||
apply_role_profile(
|
||||
&mut value,
|
||||
"coder",
|
||||
"Ticket implementation coder profile.",
|
||||
"workspace_write",
|
||||
true,
|
||||
true,
|
||||
true,
|
||||
true,
|
||||
);
|
||||
Some(value)
|
||||
}
|
||||
"builtin:reviewer" | "reviewer" => {
|
||||
apply_role_profile(
|
||||
&mut value,
|
||||
"reviewer",
|
||||
"Ticket review profile.",
|
||||
"workspace_read",
|
||||
true,
|
||||
true,
|
||||
true,
|
||||
false,
|
||||
);
|
||||
Some(value)
|
||||
}
|
||||
"builtin:memory-consolidation" | "memory-consolidation" => {
|
||||
value["slug"] = serde_json::Value::String("memory-consolidation".to_string());
|
||||
value["description"] =
|
||||
serde_json::Value::String("Memory staging consolidation profile.".to_string());
|
||||
value["feature"]["task"] = serde_json::json!({ "enabled": false });
|
||||
value["feature"]["memory"] = serde_json::json!({ "enabled": true, "staging": true });
|
||||
value["feature"]["web"] = serde_json::json!({ "enabled": false });
|
||||
value["feature"]["sub_worker"] = serde_json::json!({ "enabled": false });
|
||||
value["feature"]["objective"] = serde_json::json!({ "enabled": false });
|
||||
value["feature"]["ticket"] = serde_json::json!({ "enabled": false, "thread": false });
|
||||
Some(value)
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn builtin_base_profile_artifact() -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"slug": "default",
|
||||
"description": "Default Yoi coding profile.",
|
||||
"model": { "ref": "codex-oauth/gpt-5.5" },
|
||||
"session": { "record_event_trace": true },
|
||||
"engine": { "reasoning": "high" },
|
||||
"compaction": {
|
||||
"kind": "tokens",
|
||||
"threshold": 240000,
|
||||
"request_threshold": 270000,
|
||||
"worker_context_max_tokens": 100000
|
||||
},
|
||||
"feature": {
|
||||
"task": { "enabled": true },
|
||||
"memory": { "enabled": true },
|
||||
"web": { "enabled": true },
|
||||
"image": { "enabled": true },
|
||||
"sub_worker": { "enabled": true },
|
||||
"worker": { "enabled": false },
|
||||
"objective": { "enabled": true },
|
||||
"ticket": { "enabled": true, "authoring": true, "thread": true }
|
||||
},
|
||||
"memory": {
|
||||
"extract_threshold": 50000,
|
||||
"consolidation_threshold_files": 5,
|
||||
"consolidation_threshold_bytes": 50000
|
||||
},
|
||||
"web": {
|
||||
"enabled": true,
|
||||
"search": {
|
||||
"provider": "brave",
|
||||
"api_key_secret": "web/brave/default"
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn apply_role_profile(
|
||||
value: &mut serde_json::Value,
|
||||
slug: &str,
|
||||
description: &str,
|
||||
_scope: &str,
|
||||
task: bool,
|
||||
memory: bool,
|
||||
web: bool,
|
||||
sub_worker: bool,
|
||||
) {
|
||||
value["slug"] = serde_json::Value::String(slug.to_string());
|
||||
value["description"] = serde_json::Value::String(description.to_string());
|
||||
value["feature"]["task"] = serde_json::json!({ "enabled": task });
|
||||
value["feature"]["memory"] = serde_json::json!({ "enabled": memory });
|
||||
value["feature"]["web"] = serde_json::json!({ "enabled": web });
|
||||
value["feature"]["image"] = serde_json::json!({ "enabled": true });
|
||||
value["feature"]["sub_worker"] = serde_json::json!({ "enabled": sub_worker });
|
||||
value["feature"]["flow"] = serde_json::json!({ "enabled": slug == "coder" });
|
||||
value["feature"]["worker"] = serde_json::json!({
|
||||
"enabled": matches!(slug, "companion" | "orchestrator"),
|
||||
"direct_spawn": slug != "orchestrator"
|
||||
});
|
||||
value["feature"]["manage_workdir"] = serde_json::json!({
|
||||
"enabled": matches!(slug, "companion" | "orchestrator")
|
||||
});
|
||||
value["feature"]["orchestration"] = serde_json::json!({ "enabled": slug == "orchestrator" });
|
||||
let ticket = match slug {
|
||||
"companion" => serde_json::json!({ "enabled": true, "authoring": true, "thread": true }),
|
||||
"intake" => {
|
||||
serde_json::json!({ "enabled": true, "authoring": true, "thread": true, "intake": true })
|
||||
}
|
||||
"orchestrator" => {
|
||||
serde_json::json!({ "enabled": true, "thread": true, "workflow": true })
|
||||
}
|
||||
"coder" => serde_json::json!({ "enabled": true, "thread": true }),
|
||||
"reviewer" => serde_json::json!({ "enabled": true, "thread": true }),
|
||||
_ => serde_json::json!({ "enabled": true, "authoring": true, "thread": true }),
|
||||
};
|
||||
value["feature"]["ticket"] = ticket;
|
||||
let merge_request = match slug {
|
||||
"coder" => serde_json::json!({
|
||||
"show": true,
|
||||
"open": true,
|
||||
"review": false,
|
||||
"readiness_check": false,
|
||||
"complete": false
|
||||
}),
|
||||
"reviewer" => serde_json::json!({
|
||||
"show": true,
|
||||
"open": false,
|
||||
"review": true,
|
||||
"readiness_check": false,
|
||||
"complete": false
|
||||
}),
|
||||
"orchestrator" => serde_json::json!({
|
||||
"show": true,
|
||||
"open": false,
|
||||
"review": false,
|
||||
"readiness_check": true,
|
||||
"complete": true
|
||||
}),
|
||||
_ => serde_json::json!({
|
||||
"show": false,
|
||||
"open": false,
|
||||
"review": false,
|
||||
"readiness_check": false,
|
||||
"complete": false
|
||||
}),
|
||||
};
|
||||
value["feature"]["merge_request"] = merge_request;
|
||||
}
|
||||
|
||||
fn reject_manifest_shaped_profile(value: &serde_json::Value) -> Result<(), ProfileError> {
|
||||
let Some(map) = value.as_object() else {
|
||||
return Err(ProfileError::InvalidProfile(
|
||||
@@ -1038,12 +892,6 @@ fn validate_profile_paths(profile: &ProfileConfig) -> Result<(), ProfileError> {
|
||||
.map_err(|source| ProfileError::ProfileDeserialize { source })?;
|
||||
reject_absolute_auth_file(&model.auth, "compaction.model.auth.file")?;
|
||||
}
|
||||
if let Some(memory) = &profile.memory
|
||||
&& let Some(root) = &memory.workspace_root
|
||||
&& root.is_absolute()
|
||||
{
|
||||
return Err(ProfileError::InvalidProfile("field `memory.workspace_root` is a resolved path and is not allowed in reusable Profiles".into()));
|
||||
}
|
||||
if let Some(skills) = &profile.skills {
|
||||
for dir in &skills.directories {
|
||||
if dir.is_absolute() {
|
||||
@@ -1288,6 +1136,13 @@ pub enum ProfileError {
|
||||
#[source]
|
||||
source: toml::de::Error,
|
||||
},
|
||||
#[error("failed to evaluate built-in Profile `{selector}`: {message}")]
|
||||
BuiltinProfileEvaluation { selector: String, message: String },
|
||||
#[error("Profile requires unsupported {target} launch authorities: {requirements:?}")]
|
||||
UnsupportedExecutionTarget {
|
||||
target: ProfileExecutionTarget,
|
||||
requirements: Vec<WorkspaceAuthorityRequirement>,
|
||||
},
|
||||
#[error("no default profile is configured")]
|
||||
NoDefaultProfile,
|
||||
#[error("profile resolution requires an explicit runtime Worker name")]
|
||||
@@ -1341,18 +1196,21 @@ mod tests {
|
||||
);
|
||||
}
|
||||
#[test]
|
||||
fn builtin_profiles_do_not_define_an_implicit_default() {
|
||||
fn builtin_default_is_explicit_registry_authority() {
|
||||
let registry = ProfileDiscovery::with_sources(None, None)
|
||||
.discover()
|
||||
.unwrap();
|
||||
assert!(matches!(
|
||||
registry.default_entry(),
|
||||
Err(ProfileError::NoDefaultProfile)
|
||||
));
|
||||
assert!(matches!(
|
||||
registry.select(&ProfileSelector::Default),
|
||||
Err(ProfileError::NoDefaultProfile)
|
||||
));
|
||||
let default = registry.default_entry().unwrap();
|
||||
assert_eq!(default.source, ProfileRegistrySource::Builtin);
|
||||
assert_eq!(default.name, "default");
|
||||
assert_eq!(default.qualified_name(), BUILTIN_DEFAULT_PROFILE);
|
||||
assert!(default.is_default);
|
||||
assert!(
|
||||
default
|
||||
.provenance
|
||||
.starts_with("profiles/default.dcdl#sha256:")
|
||||
);
|
||||
assert_eq!(registry.select(&ProfileSelector::Default).unwrap(), default);
|
||||
}
|
||||
#[test]
|
||||
fn builtin_role_profiles_are_registered_and_resolve() {
|
||||
@@ -1387,7 +1245,9 @@ mod tests {
|
||||
("settings_revision", serde_json::json!(2)),
|
||||
("language", serde_json::json!("Japanese")),
|
||||
] {
|
||||
let artifact = serde_json::json!({ "memory": { (field): value } });
|
||||
let artifact = serde_json::json!({
|
||||
"feature": { "memory": { (field): value } }
|
||||
});
|
||||
let error = resolve_profile_artifact_value(
|
||||
artifact,
|
||||
ProfileSource::Registry {
|
||||
@@ -1408,7 +1268,108 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builtin_companion_can_manage_workdirs() {
|
||||
fn builtin_default_resolves_as_a_standalone_local_capability_profile() {
|
||||
let tmp = TempDir::new().unwrap();
|
||||
let resolved = ProfileResolver::new()
|
||||
.with_workspace_base(tmp.path())
|
||||
.resolve_for_target(
|
||||
&ProfileSelector::source_named(ProfileRegistrySource::Builtin, "default"),
|
||||
ProfileResolveOptions::with_worker_name("standalone-worker"),
|
||||
ProfileExecutionTarget::Standalone,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert!(matches!(
|
||||
&resolved.source,
|
||||
ProfileSource::Registry {
|
||||
source: ProfileRegistrySource::Builtin,
|
||||
name,
|
||||
path: None,
|
||||
provenance: Some(provenance),
|
||||
..
|
||||
} if name == "default" && provenance.starts_with("profiles/default.dcdl#sha256:")
|
||||
));
|
||||
assert!(resolved.manifest.feature.task.enabled);
|
||||
assert!(resolved.manifest.feature.web.enabled);
|
||||
assert!(resolved.manifest.feature.image.enabled);
|
||||
assert!(resolved.manifest.feature.sub_worker.enabled);
|
||||
assert!(resolved.manifest.scope.allow.iter().any(|rule| {
|
||||
rule.permission == protocol::Permission::Write && rule.target == tmp.path()
|
||||
}));
|
||||
assert!(resolved.manifest.delegation_scope.allow.iter().any(|rule| {
|
||||
rule.permission == protocol::Permission::Write && rule.target == tmp.path()
|
||||
}));
|
||||
assert!(!resolved.manifest.feature.memory.profile.enabled);
|
||||
assert!(!resolved.manifest.feature.ticket.enabled);
|
||||
assert!(!resolved.manifest.feature.objective.enabled);
|
||||
assert!(!resolved.manifest.feature.flow.enabled);
|
||||
assert!(!resolved.manifest.feature.worker.enabled);
|
||||
assert!(!resolved.manifest.feature.manage_workdir.enabled);
|
||||
assert!(!resolved.manifest.feature.plugins.enabled);
|
||||
assert!(resolved.manifest.plugins.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn standalone_rejects_profiles_that_require_workspace_authority() {
|
||||
let tmp = TempDir::new().unwrap();
|
||||
let error = ProfileResolver::new()
|
||||
.with_workspace_base(tmp.path())
|
||||
.resolve_for_target(
|
||||
&ProfileSelector::source_named(ProfileRegistrySource::Builtin, "coder"),
|
||||
ProfileResolveOptions::with_worker_name("standalone-worker"),
|
||||
ProfileExecutionTarget::Standalone,
|
||||
)
|
||||
.unwrap_err();
|
||||
let diagnostic = error.to_string();
|
||||
|
||||
let ProfileError::UnsupportedExecutionTarget {
|
||||
target,
|
||||
requirements,
|
||||
} = error
|
||||
else {
|
||||
panic!("unexpected error: {error}");
|
||||
};
|
||||
assert_eq!(target, ProfileExecutionTarget::Standalone);
|
||||
assert!(requirements.contains(&WorkspaceAuthorityRequirement::Memory));
|
||||
assert!(requirements.contains(&WorkspaceAuthorityRequirement::MergeRequest));
|
||||
assert!(requirements.contains(&WorkspaceAuthorityRequirement::Ticket));
|
||||
assert!(!diagnostic.contains(tmp.path().to_string_lossy().as_ref()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn repository_markers_do_not_change_builtin_profile_authority() {
|
||||
let tmp = TempDir::new().unwrap();
|
||||
let nested = tmp.path().join("repository/nested");
|
||||
std::fs::create_dir_all(&nested).unwrap();
|
||||
std::fs::create_dir_all(tmp.path().join("repository/.yoi")).unwrap();
|
||||
std::fs::write(
|
||||
tmp.path().join("repository/.yoi/profiles.toml"),
|
||||
"default = { source = 'project', name = 'shadow' }\n",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let discovery = ProfileDiscovery::for_cwd(&nested);
|
||||
assert_eq!(discovery.user_config, paths::user_profiles_path());
|
||||
assert!(discovery.project_config.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builtin_coder_uses_sub_worker_control_without_worker_control() {
|
||||
let tmp = TempDir::new().unwrap();
|
||||
let resolved = ProfileResolver::new()
|
||||
.with_workspace_base(tmp.path())
|
||||
.resolve(
|
||||
&ProfileSelector::source_named(ProfileRegistrySource::Builtin, "coder"),
|
||||
ProfileResolveOptions::with_worker_name("coder-worker"),
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert!(resolved.manifest.feature.sub_worker.enabled);
|
||||
assert!(!resolved.manifest.feature.worker.enabled);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builtin_companion_combines_runtime_and_sub_worker_control_with_discovery() {
|
||||
let tmp = TempDir::new().unwrap();
|
||||
let resolved = ProfileResolver::new()
|
||||
.with_workspace_base(tmp.path())
|
||||
@@ -1419,6 +1380,32 @@ mod tests {
|
||||
.unwrap();
|
||||
|
||||
assert!(resolved.manifest.feature.manage_workdir.enabled);
|
||||
assert!(resolved.manifest.feature.sub_worker.enabled);
|
||||
assert!(resolved.manifest.feature.worker.enabled);
|
||||
assert!(!resolved.manifest.feature.worker.direct_spawn);
|
||||
assert!(resolved.manifest.feature.workspace_worker_discovery.enabled);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builtin_orchestrator_keeps_cleanup_tool_providers_enabled() {
|
||||
let tmp = TempDir::new().unwrap();
|
||||
let resolved = ProfileResolver::new()
|
||||
.with_workspace_base(tmp.path())
|
||||
.resolve(
|
||||
&ProfileSelector::source_named(ProfileRegistrySource::Builtin, "orchestrator"),
|
||||
ProfileResolveOptions::with_worker_name("orchestrator-worker"),
|
||||
)
|
||||
.unwrap();
|
||||
let feature = resolved.manifest.feature;
|
||||
|
||||
assert!(feature.worker.enabled);
|
||||
assert!(!feature.worker.direct_spawn);
|
||||
assert!(feature.manage_workdir.enabled);
|
||||
assert!(feature.merge_request.show);
|
||||
assert!(feature.merge_request.readiness_check);
|
||||
assert!(feature.merge_request.complete);
|
||||
assert!(!feature.merge_request.open);
|
||||
assert!(!feature.merge_request.review);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -1591,7 +1578,7 @@ enabled = false
|
||||
.unwrap();
|
||||
assert_eq!(resolved.manifest.worker.name, "runtime-worker");
|
||||
assert!(resolved.manifest.feature.task.enabled);
|
||||
assert!(!resolved.manifest.feature.memory.enabled);
|
||||
assert!(!resolved.manifest.feature.memory.profile.enabled);
|
||||
assert!(resolved.manifest.feature.web.enabled);
|
||||
assert!(resolved.manifest.feature.sub_worker.enabled);
|
||||
assert!(resolved.manifest.feature.ticket.enabled);
|
||||
|
||||
@@ -152,13 +152,10 @@ pub enum MemoryStagingAffectedMemoryOperation {
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct MemoryConsolidateStagingOperation {
|
||||
#[serde(default)]
|
||||
pub force: bool,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub threshold_files: Option<usize>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub threshold_bytes: Option<u64>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
@@ -450,10 +447,21 @@ mod tests {
|
||||
use super::*;
|
||||
use crate::extract::{CandidateKind, ExtractedCandidate};
|
||||
|
||||
#[test]
|
||||
fn consolidation_operation_rejects_caller_owned_thresholds() {
|
||||
let error =
|
||||
serde_json::from_value::<MemoryConsolidateStagingOperation>(serde_json::json!({
|
||||
"force": false,
|
||||
"threshold_files": 1,
|
||||
}))
|
||||
.unwrap_err();
|
||||
assert!(error.to_string().contains("threshold_files"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn staging_list_read_close_records_reason_and_deletes_candidate() {
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
let layout = WorkspaceLayout::resolve(&manifest::MemoryConfig::default(), temp.path());
|
||||
let layout = WorkspaceLayout::resolve(temp.path());
|
||||
let source = SourceRef {
|
||||
segment_id: "segment-1".into(),
|
||||
range: [0, 1],
|
||||
|
||||
@@ -21,8 +21,7 @@ pub struct StagingEntry {
|
||||
pub id: Uuid,
|
||||
pub path: PathBuf,
|
||||
pub record: StagingRecord,
|
||||
/// このファイルのバイト長。閾値判定 (`consolidation_threshold_bytes`)
|
||||
/// に使う。
|
||||
/// このファイルのバイト長。Backendのconsolidation閾値判定に使用する。
|
||||
pub bytes: u64,
|
||||
}
|
||||
|
||||
|
||||
@@ -74,6 +74,7 @@ impl ExtractedPayload {
|
||||
|
||||
/// Bounded evidence snippet copied into a flat staging record.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct StagingEvidence {
|
||||
pub id: String,
|
||||
pub kind: EvidenceKind,
|
||||
@@ -89,6 +90,7 @@ pub struct StagingEvidence {
|
||||
|
||||
/// One flat staging record. One record is one consolidation decision unit.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct StagingRecord {
|
||||
pub schema_version: u32,
|
||||
pub id: String,
|
||||
|
||||
@@ -22,6 +22,7 @@ impl<'de> Deserialize<'de> for SourceRef {
|
||||
D: serde::Deserializer<'de>,
|
||||
{
|
||||
#[derive(Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct RawSourceRef {
|
||||
#[serde(default)]
|
||||
segment_id: Option<String>,
|
||||
@@ -83,6 +84,7 @@ pub enum EvidenceOriginKind {
|
||||
/// Bounded origin snapshot attached to extraction evidence. This is audit
|
||||
/// metadata only and cannot authorize Workspace operations.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct EvidenceOrigin {
|
||||
pub kind: EvidenceOriginKind,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
@@ -107,6 +109,7 @@ pub struct EvidenceOrigin {
|
||||
/// ranges, and short labels/summaries. It must not carry raw message bodies or
|
||||
/// full tool result content.
|
||||
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct SourceEvidenceRef {
|
||||
/// Stable session id when the anchor crosses or disambiguates segments.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
|
||||
@@ -70,24 +70,12 @@ impl WorkspaceLayout {
|
||||
Self { root: root.into() }
|
||||
}
|
||||
|
||||
/// Resolve a layout from a `MemoryConfig`.
|
||||
/// Resolve a layout from the nearest Memory marker.
|
||||
///
|
||||
/// An explicit `memory.workspace_root` is honored exactly. Without an
|
||||
/// explicit root, resolution searches `default_root` and its ancestors for
|
||||
/// the nearest `.yoi/memory` directory. This keeps child worktrees that
|
||||
/// contain `.yoi` project records such as tickets from
|
||||
/// becoming independent memory roots merely because they contain `.yoi`.
|
||||
///
|
||||
/// If no memory marker exists, this falls back to `default_root` because
|
||||
/// existing call sites require a concrete layout. That fallback is a
|
||||
/// no-marker compatibility path, not a `.yoi` marker interpretation; it
|
||||
/// must not be used as evidence that `.yoi` alone enables repo-local
|
||||
/// memory.
|
||||
pub fn resolve(cfg: &manifest::MemoryConfig, default_root: &Path) -> Self {
|
||||
if let Some(root) = &cfg.workspace_root {
|
||||
return Self::new(root.clone());
|
||||
}
|
||||
|
||||
/// Resolution searches `default_root` and its ancestors for the nearest
|
||||
/// `.yoi/memory` directory. This legacy local-storage helper owns its path
|
||||
/// policy directly; resolved Worker Manifests do not carry storage paths.
|
||||
pub fn resolve(default_root: &Path) -> Self {
|
||||
let root =
|
||||
find_memory_marker_root(default_root).unwrap_or_else(|| default_root.to_path_buf());
|
||||
Self::new(root)
|
||||
@@ -335,16 +323,6 @@ mod tests {
|
||||
assert!(matches!(err, LintError::InvalidPath(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolve_uses_workspace_root_when_set() {
|
||||
let cfg = manifest::MemoryConfig {
|
||||
workspace_root: Some(PathBuf::from("/explicit")),
|
||||
..Default::default()
|
||||
};
|
||||
let layout = WorkspaceLayout::resolve(&cfg, Path::new("/fallback"));
|
||||
assert_eq!(layout.root(), Path::new("/explicit"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolve_selects_nearest_ancestor_memory_marker_when_workspace_root_missing() {
|
||||
let tmp = TempDir::new().unwrap();
|
||||
@@ -353,8 +331,7 @@ mod tests {
|
||||
std::fs::create_dir_all(workspace.join(".yoi/memory")).unwrap();
|
||||
std::fs::create_dir_all(&child).unwrap();
|
||||
|
||||
let cfg = manifest::MemoryConfig::default();
|
||||
let layout = WorkspaceLayout::resolve(&cfg, &child);
|
||||
let layout = WorkspaceLayout::resolve(&child);
|
||||
assert_eq!(layout.root(), workspace.as_path());
|
||||
}
|
||||
|
||||
@@ -366,8 +343,7 @@ mod tests {
|
||||
std::fs::create_dir_all(workspace.join(".yoi/memory")).unwrap();
|
||||
std::fs::create_dir_all(child.join(".yoi/tickets")).unwrap();
|
||||
|
||||
let cfg = manifest::MemoryConfig::default();
|
||||
let layout = WorkspaceLayout::resolve(&cfg, &child);
|
||||
let layout = WorkspaceLayout::resolve(&child);
|
||||
assert_eq!(layout.root(), workspace.as_path());
|
||||
}
|
||||
|
||||
@@ -381,8 +357,7 @@ mod tests {
|
||||
|
||||
assert_eq!(find_memory_marker_root(&child), None);
|
||||
|
||||
let cfg = manifest::MemoryConfig::default();
|
||||
let layout = WorkspaceLayout::resolve(&cfg, &child);
|
||||
let layout = WorkspaceLayout::resolve(&child);
|
||||
assert_eq!(layout.root(), child.as_path());
|
||||
}
|
||||
}
|
||||
|
||||
+45
-239
@@ -9,7 +9,6 @@ use thiserror::Error;
|
||||
use uuid::Uuid;
|
||||
|
||||
const SCHEMA_VERSION: i64 = 12;
|
||||
const PREVIOUS_SCHEMA_VERSION: i64 = 11;
|
||||
const MAX_BODY_BYTES: usize = 16 * 1024;
|
||||
const DOMAIN_TABLES: [&str; 5] = [
|
||||
"merge_requests",
|
||||
@@ -37,7 +36,7 @@ impl MergeRequestState {
|
||||
|
||||
fn parse(v: &str) -> Result<Self, MergeRequestError> {
|
||||
match v {
|
||||
"draft" | "open" => Ok(Self::Open),
|
||||
"open" => Ok(Self::Open),
|
||||
"merged" => Ok(Self::Merged),
|
||||
"closed" => Ok(Self::Closed),
|
||||
_ => Err(MergeRequestError::Corrupt(format!("unknown state `{v}`"))),
|
||||
@@ -274,6 +273,12 @@ pub struct RegisterReviewerChildSession {
|
||||
pub reviewer_profile: String,
|
||||
pub now: DateTime<Utc>,
|
||||
}
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct ReviewSubmissionAuthorization {
|
||||
pub workspace_id: String,
|
||||
pub subject_ref: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SubmitMergeRequestReview {
|
||||
pub ticket_id: String,
|
||||
@@ -535,6 +540,34 @@ impl MergeRequestStore {
|
||||
t.commit()?;
|
||||
Ok(RequestedMergeRequestReview { request_event: e })
|
||||
}
|
||||
pub fn authorize_review_submission(
|
||||
&self,
|
||||
ticket_id: &str,
|
||||
capability_token: &str,
|
||||
) -> Result<ReviewSubmissionAuthorization, MergeRequestError> {
|
||||
let connection = self.lock()?;
|
||||
connection
|
||||
.query_row(
|
||||
"SELECT g.workspace_id,g.subject_ref
|
||||
FROM merge_request_review_grants g
|
||||
JOIN merge_request_ticket_relations rel
|
||||
ON rel.workspace_id=g.workspace_id AND rel.merge_request_id=g.merge_request_id
|
||||
JOIN merge_requests mr
|
||||
ON mr.workspace_id=g.workspace_id AND mr.merge_request_id=g.merge_request_id
|
||||
WHERE g.capability_token=?1 AND rel.ticket_id=?2
|
||||
AND g.status='issued' AND mr.state='open'",
|
||||
params![capability_token, ticket_id],
|
||||
|row| {
|
||||
Ok(ReviewSubmissionAuthorization {
|
||||
workspace_id: row.get(0)?,
|
||||
subject_ref: row.get(1)?,
|
||||
})
|
||||
},
|
||||
)
|
||||
.optional()?
|
||||
.ok_or_else(|| MergeRequestError::Unauthorized("review grant invalid".into()))
|
||||
}
|
||||
|
||||
pub fn submit_review(
|
||||
&self,
|
||||
i: SubmitMergeRequestReview,
|
||||
@@ -1321,14 +1354,9 @@ pub fn migrate(c: &Connection) -> Result<(), MergeRequestError> {
|
||||
match schema_state(c)? {
|
||||
SchemaState::Fresh => fresh(c),
|
||||
SchemaState::Current(SCHEMA_VERSION) => verify(c),
|
||||
SchemaState::Current(PREVIOUS_SCHEMA_VERSION) => from_v11(c, PreviousSchemaMarker::Current),
|
||||
SchemaState::Legacy(PREVIOUS_SCHEMA_VERSION) => from_v11(c, PreviousSchemaMarker::Legacy),
|
||||
SchemaState::Current(v) => Err(MergeRequestError::Operation(format!(
|
||||
"unsupported schema {v}"
|
||||
))),
|
||||
SchemaState::Legacy(v) => Err(MergeRequestError::Operation(format!(
|
||||
"unsupported legacy schema {v}"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1336,26 +1364,14 @@ pub fn migrate(c: &Connection) -> Result<(), MergeRequestError> {
|
||||
enum SchemaState {
|
||||
Fresh,
|
||||
Current(i64),
|
||||
Legacy(i64),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum PreviousSchemaMarker {
|
||||
Current,
|
||||
Legacy,
|
||||
}
|
||||
|
||||
fn schema_state(c: &Connection) -> Result<SchemaState, MergeRequestError> {
|
||||
let (current, legacy): (bool, bool) = c.query_row(
|
||||
"SELECT EXISTS(SELECT 1 FROM sqlite_master WHERE type='table' AND name='merge_request_schema'),EXISTS(SELECT 1 FROM sqlite_master WHERE type='table' AND name='merge_request_schema_migrations')",
|
||||
let current: bool = c.query_row(
|
||||
"SELECT EXISTS(SELECT 1 FROM sqlite_master WHERE type='table' AND name='merge_request_schema')",
|
||||
[],
|
||||
|r| Ok((r.get(0)?, r.get(1)?)),
|
||||
|r| r.get(0),
|
||||
)?;
|
||||
if current && legacy {
|
||||
return Err(MergeRequestError::Corrupt(
|
||||
"both current and legacy schema markers exist".into(),
|
||||
));
|
||||
}
|
||||
if current {
|
||||
let (count, singleton, version): (i64, Option<i64>, Option<i64>) = c.query_row(
|
||||
"SELECT COUNT(*),MIN(singleton),MAX(version) FROM merge_request_schema",
|
||||
@@ -1372,22 +1388,6 @@ fn schema_state(c: &Connection) -> Result<SchemaState, MergeRequestError> {
|
||||
})?;
|
||||
return Ok(SchemaState::Current(version));
|
||||
}
|
||||
if legacy {
|
||||
let (count, version): (i64, Option<i64>) = c.query_row(
|
||||
"SELECT COUNT(*),MAX(version) FROM merge_request_schema_migrations",
|
||||
[],
|
||||
|r| Ok((r.get(0)?, r.get(1)?)),
|
||||
)?;
|
||||
if count != 1 {
|
||||
return Err(MergeRequestError::Corrupt(
|
||||
"legacy schema marker must contain exactly one version".into(),
|
||||
));
|
||||
}
|
||||
let version = version.ok_or_else(|| {
|
||||
MergeRequestError::Corrupt("legacy schema marker version is null".into())
|
||||
})?;
|
||||
return Ok(SchemaState::Legacy(version));
|
||||
}
|
||||
let domain_tables: bool = c.query_row(
|
||||
"SELECT EXISTS(SELECT 1 FROM sqlite_master WHERE type='table' AND name GLOB 'merge_request*')",
|
||||
[],
|
||||
@@ -1402,214 +1402,20 @@ fn schema_state(c: &Connection) -> Result<SchemaState, MergeRequestError> {
|
||||
}
|
||||
fn fresh(c: &Connection) -> Result<(), MergeRequestError> {
|
||||
let t = c.unchecked_transaction()?;
|
||||
tables(&t, true)?;
|
||||
t.execute("INSERT INTO merge_request_schema VALUES(1,12)", [])?;
|
||||
tables(&t)?;
|
||||
t.execute(
|
||||
"INSERT INTO merge_request_schema VALUES(1,?1)",
|
||||
params![SCHEMA_VERSION],
|
||||
)?;
|
||||
fk(&t)?;
|
||||
t.commit()?;
|
||||
Ok(())
|
||||
}
|
||||
fn tables(t: &Transaction<'_>, marker: bool) -> Result<(), MergeRequestError> {
|
||||
if marker {
|
||||
t.execute_batch("CREATE TABLE merge_request_schema(singleton INTEGER PRIMARY KEY CHECK(singleton=1),version INTEGER NOT NULL);")?
|
||||
}
|
||||
fn tables(t: &Transaction<'_>) -> Result<(), MergeRequestError> {
|
||||
t.execute_batch("CREATE TABLE merge_request_schema(singleton INTEGER PRIMARY KEY CHECK(singleton=1),version INTEGER NOT NULL);")?;
|
||||
t.execute_batch("CREATE TABLE merge_requests(workspace_id TEXT NOT NULL,merge_request_id TEXT NOT NULL,repository_id TEXT NOT NULL,state TEXT NOT NULL CHECK(state IN('open','merged','closed')),selector_from TEXT,selector_to TEXT NOT NULL,created_at TEXT NOT NULL,updated_at TEXT NOT NULL,PRIMARY KEY(workspace_id,merge_request_id),FOREIGN KEY(workspace_id,repository_id)REFERENCES repositories(workspace_id,repository_id));CREATE TABLE merge_request_ticket_relations(workspace_id TEXT NOT NULL,merge_request_id TEXT NOT NULL,ticket_id TEXT NOT NULL,relation_kind TEXT NOT NULL CHECK(relation_kind='implements'),created_at TEXT NOT NULL,PRIMARY KEY(workspace_id,merge_request_id,ticket_id),FOREIGN KEY(workspace_id,merge_request_id)REFERENCES merge_requests(workspace_id,merge_request_id)ON DELETE CASCADE,FOREIGN KEY(workspace_id,ticket_id)REFERENCES typed_tickets(workspace_id,ticket_id)ON DELETE CASCADE);CREATE TABLE merge_request_thread_events(workspace_id TEXT NOT NULL,merge_request_id TEXT NOT NULL,event_id TEXT NOT NULL,sequence INTEGER NOT NULL,kind TEXT NOT NULL CHECK(kind IN('review_requested','review','review_revoked','review_cancelled','comment','merge')),payload_json TEXT NOT NULL,operation_id TEXT,created_at TEXT NOT NULL,PRIMARY KEY(workspace_id,merge_request_id,event_id),UNIQUE(workspace_id,merge_request_id,sequence),FOREIGN KEY(workspace_id,merge_request_id)REFERENCES merge_requests(workspace_id,merge_request_id)ON DELETE CASCADE);CREATE UNIQUE INDEX merge_request_merge_operations ON merge_request_thread_events(workspace_id,operation_id)WHERE operation_id IS NOT NULL;CREATE TABLE merge_request_review_grants(workspace_id TEXT NOT NULL,merge_request_id TEXT NOT NULL,request_event_id TEXT NOT NULL,subject_ref TEXT NOT NULL,reviewer_runtime_id TEXT NOT NULL,reviewer_worker_id TEXT NOT NULL,capability_token TEXT PRIMARY KEY,issued_at TEXT NOT NULL,consumed_at TEXT,revoked_at TEXT,status TEXT NOT NULL CHECK(status IN('issued','consumed','revoked')),FOREIGN KEY(workspace_id,merge_request_id,request_event_id)REFERENCES merge_request_thread_events(workspace_id,merge_request_id,event_id)ON DELETE CASCADE);CREATE TABLE merge_request_reviewer_child_sessions(workspace_id TEXT NOT NULL,child_session_id TEXT NOT NULL,parent_runtime_id TEXT NOT NULL,parent_worker_id TEXT NOT NULL,reviewer_profile TEXT NOT NULL,registered_at TEXT NOT NULL,status TEXT NOT NULL CHECK(status IN('active','consumed')),PRIMARY KEY(workspace_id,child_session_id));")?;
|
||||
Ok(())
|
||||
}
|
||||
fn from_v11(
|
||||
c: &Connection,
|
||||
previous_marker: PreviousSchemaMarker,
|
||||
) -> Result<(), MergeRequestError> {
|
||||
let t = c.unchecked_transaction()?;
|
||||
if previous_marker == PreviousSchemaMarker::Legacy {
|
||||
t.execute_batch("CREATE TABLE merge_request_schema(singleton INTEGER PRIMARY KEY CHECK(singleton=1),version INTEGER NOT NULL);")?;
|
||||
t.execute(
|
||||
"INSERT INTO merge_request_schema VALUES(1,?1)",
|
||||
params![PREVIOUS_SCHEMA_VERSION],
|
||||
)?;
|
||||
}
|
||||
t.execute_batch("ALTER TABLE merge_requests RENAME TO merge_requests_v11;ALTER TABLE merge_request_ticket_relations RENAME TO merge_request_ticket_relations_v11;ALTER TABLE merge_request_revisions RENAME TO merge_request_revisions_v11;ALTER TABLE merge_request_revision_paths RENAME TO merge_request_revision_paths_v11;ALTER TABLE merge_request_reviewer_child_sessions RENAME TO merge_request_reviewer_child_sessions_v11;ALTER TABLE merge_request_review_attempts RENAME TO merge_request_review_attempts_v11;ALTER TABLE merge_request_reviews RENAME TO merge_request_reviews_v11;ALTER TABLE merge_request_review_findings RENAME TO merge_request_review_findings_v11;ALTER TABLE merge_request_completion_operations RENAME TO merge_request_completion_operations_v11;")?;
|
||||
tables(&t, false)?;
|
||||
t.execute("INSERT INTO merge_requests SELECT workspace_id,merge_request_id,repository_id,CASE state WHEN 'draft'THEN'open'ELSE state END,NULL,target_ref_selector,created_at,updated_at FROM merge_requests_v11",[])?;
|
||||
t.execute("INSERT INTO merge_request_ticket_relations SELECT * FROM merge_request_ticket_relations_v11",[])?;
|
||||
migrate_events(&t)?;
|
||||
if previous_marker == PreviousSchemaMarker::Legacy {
|
||||
t.execute("DROP TABLE merge_request_schema_migrations", [])?;
|
||||
}
|
||||
t.execute_batch("DROP TABLE merge_request_review_findings_v11;DROP TABLE merge_request_reviews_v11;DROP TABLE merge_request_review_attempts_v11;DROP TABLE merge_request_reviewer_child_sessions_v11;DROP TABLE merge_request_revision_paths_v11;DROP TABLE merge_request_revisions_v11;DROP TABLE merge_request_completion_operations_v11;DROP TABLE merge_request_ticket_relations_v11;DROP TABLE merge_requests_v11;UPDATE merge_request_schema SET version=12 WHERE singleton=1;")?;
|
||||
fk(&t)?;
|
||||
t.commit()?;
|
||||
Ok(())
|
||||
}
|
||||
fn migrate_events(t: &Transaction<'_>) -> Result<(), MergeRequestError> {
|
||||
let attempts = {
|
||||
let mut s=t.prepare("SELECT a.workspace_id,a.attempt_id,a.merge_request_id,a.parent_runtime_id,a.parent_worker_id,a.child_session_id,a.status,a.created_at,a.consumed_at,r.head_commit FROM merge_request_review_attempts_v11 a JOIN merge_request_revisions_v11 r ON r.workspace_id=a.workspace_id AND r.merge_request_id=a.merge_request_id AND r.revision_id=a.revision_id ORDER BY a.created_at")?;
|
||||
s.query_map([], |r| {
|
||||
Ok((
|
||||
r.get::<_, String>(0)?,
|
||||
r.get::<_, String>(1)?,
|
||||
r.get::<_, String>(2)?,
|
||||
r.get::<_, String>(3)?,
|
||||
r.get::<_, String>(4)?,
|
||||
r.get::<_, String>(5)?,
|
||||
r.get::<_, String>(6)?,
|
||||
r.get::<_, String>(7)?,
|
||||
r.get::<_, Option<String>>(8)?,
|
||||
r.get::<_, String>(9)?,
|
||||
))
|
||||
})?
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
};
|
||||
for (ws, a, mr, pr, pw, child, status, created, consumed, subject) in attempts {
|
||||
let req = ReviewRequestedEvent {
|
||||
event_id: format!("migrated-request-{a}"),
|
||||
sequence: next_seq(t, &ws, &mr)?,
|
||||
subject_ref: subject.clone(),
|
||||
requested_by: WorkerIdentity {
|
||||
runtime_id: pr.clone(),
|
||||
worker_id: pw,
|
||||
},
|
||||
reviewer: WorkerIdentity {
|
||||
runtime_id: pr,
|
||||
worker_id: child,
|
||||
},
|
||||
created_at: time(&created)?,
|
||||
};
|
||||
insert_event(t, &ws, &mr, "review_requested", &req, req.created_at, None)?;
|
||||
if status == "submitted" {
|
||||
let(row_dec,row_body,row_at):(String,String,String)=t.query_row("SELECT decision,body,submitted_at FROM merge_request_reviews_v11 WHERE workspace_id=?1 AND attempt_id=?2",params![ws,a],|r|Ok((r.get(0)?,r.get(1)?,r.get(2)?)))?;
|
||||
let findings = {
|
||||
let mut s=t.prepare("SELECT severity,code,path,line,body FROM merge_request_review_findings_v11 WHERE workspace_id=?1 AND attempt_id=?2 ORDER BY ordinal")?;
|
||||
s.query_map(params![ws, a], |r| {
|
||||
Ok(ReviewFinding {
|
||||
severity: match r.get::<_, String>(0)?.as_str() {
|
||||
"blocker" => FindingSeverity::Blocker,
|
||||
"major" => FindingSeverity::Major,
|
||||
"minor" => FindingSeverity::Minor,
|
||||
_ => FindingSeverity::Note,
|
||||
},
|
||||
code: r.get(1)?,
|
||||
path: r.get(2)?,
|
||||
line: r.get(3)?,
|
||||
body: r.get(4)?,
|
||||
})
|
||||
})?
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
};
|
||||
let rev = ReviewEvent {
|
||||
event_id: format!("migrated-review-{a}"),
|
||||
sequence: next_seq(t, &ws, &mr)?,
|
||||
request_event_id: req.event_id,
|
||||
subject_ref: subject,
|
||||
decision: if row_dec == "approve" {
|
||||
ReviewDecision::Approve
|
||||
} else {
|
||||
ReviewDecision::RequestChanges
|
||||
},
|
||||
body: row_body,
|
||||
findings,
|
||||
reviewer: req.reviewer,
|
||||
created_at: time(&row_at)?,
|
||||
};
|
||||
insert_event(t, &ws, &mr, "review", &rev, rev.created_at, None)?
|
||||
} else {
|
||||
let at = consumed.as_deref().unwrap_or(&created);
|
||||
let e = ReviewCancelledEvent {
|
||||
event_id: format!("migrated-cancel-{a}"),
|
||||
sequence: next_seq(t, &ws, &mr)?,
|
||||
request_event_id: req.event_id,
|
||||
subject_ref: subject,
|
||||
reason: format!(
|
||||
"legacy `{status}` review request cancelled because its capability cannot be migrated"
|
||||
),
|
||||
created_at: time(at)?,
|
||||
};
|
||||
insert_event(t, &ws, &mr, "review_cancelled", &e, e.created_at, None)?
|
||||
}
|
||||
}
|
||||
let completed = {
|
||||
let mut q=t.prepare("SELECT c.workspace_id,c.operation_id,c.ticket_id,c.target_commit,c.source_commit,c.result_commit,c.strategy,c.resolution,c.completion_actor_runtime_id,c.completion_actor_worker_id,c.updated_at,rel.merge_request_id FROM merge_request_completion_operations_v11 c JOIN merge_request_ticket_relations_v11 rel ON rel.workspace_id=c.workspace_id AND rel.ticket_id=c.ticket_id WHERE c.status='completed' ORDER BY c.updated_at")?;
|
||||
q.query_map([], |r| {
|
||||
Ok((
|
||||
r.get::<_, String>(0)?,
|
||||
r.get::<_, String>(1)?,
|
||||
r.get::<_, String>(2)?,
|
||||
r.get::<_, Option<String>>(3)?,
|
||||
r.get::<_, Option<String>>(4)?,
|
||||
r.get::<_, Option<String>>(5)?,
|
||||
r.get::<_, Option<String>>(6)?,
|
||||
r.get::<_, Option<String>>(7)?,
|
||||
r.get::<_, Option<String>>(8)?,
|
||||
r.get::<_, Option<String>>(9)?,
|
||||
r.get::<_, String>(10)?,
|
||||
r.get::<_, String>(11)?,
|
||||
))
|
||||
})?
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
};
|
||||
for (
|
||||
ws,
|
||||
op,
|
||||
_ticket,
|
||||
target,
|
||||
source,
|
||||
result,
|
||||
strategy,
|
||||
resolution,
|
||||
runtime,
|
||||
worker,
|
||||
updated,
|
||||
mr,
|
||||
) in completed
|
||||
{
|
||||
let subject = source.ok_or_else(|| {
|
||||
MergeRequestError::Operation(format!("completed operation {op} lacks source evidence"))
|
||||
})?;
|
||||
let approval:Option<String>=t.query_row("SELECT event_id FROM merge_request_thread_events WHERE workspace_id=?1 AND merge_request_id=?2 AND kind='review' AND json_extract(payload_json,'$.subject_ref')=?3 AND json_extract(payload_json,'$.decision')='approve' ORDER BY sequence DESC LIMIT 1",params![ws,mr,subject],|r|r.get(0)).optional()?;
|
||||
let approval = approval.ok_or_else(|| {
|
||||
MergeRequestError::Operation(format!(
|
||||
"completed operation {op} lacks approval evidence"
|
||||
))
|
||||
})?;
|
||||
let e = MergeEvent {
|
||||
event_id: format!("migrated-merge-{op}"),
|
||||
sequence: next_seq(t, &ws, &mr)?,
|
||||
operation_id: op,
|
||||
approval_event_id: approval,
|
||||
approved_source_ref: subject,
|
||||
target_ref_before: target.ok_or_else(|| {
|
||||
MergeRequestError::Operation("completed operation lacks target evidence".into())
|
||||
})?,
|
||||
target_ref_after: result.ok_or_else(|| {
|
||||
MergeRequestError::Operation("completed operation lacks result evidence".into())
|
||||
})?,
|
||||
strategy: if strategy.as_deref() == Some("merge") {
|
||||
MergeStrategy::Merge
|
||||
} else {
|
||||
MergeStrategy::FastForward
|
||||
},
|
||||
resolution: match resolution.as_deref() {
|
||||
Some("clean") => ConflictResolution::Clean,
|
||||
Some("conflicts_resolved") => ConflictResolution::ConflictsResolved,
|
||||
_ => ConflictResolution::None,
|
||||
},
|
||||
merged_by: WorkerIdentity {
|
||||
runtime_id: runtime.unwrap_or_else(|| "legacy".into()),
|
||||
worker_id: worker.unwrap_or_else(|| "legacy".into()),
|
||||
},
|
||||
created_at: time(&updated)?,
|
||||
};
|
||||
insert_event(
|
||||
t,
|
||||
&ws,
|
||||
&mr,
|
||||
"merge",
|
||||
&e,
|
||||
e.created_at,
|
||||
Some(&e.operation_id),
|
||||
)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
fn verify(c: &Connection) -> Result<(), MergeRequestError> {
|
||||
for n in DOMAIN_TABLES {
|
||||
let e: bool = c.query_row(
|
||||
|
||||
@@ -91,6 +91,23 @@ fn approve(s: &MergeRequestStore, subject: &str, token: &str) -> ReviewEvent {
|
||||
})
|
||||
.unwrap()
|
||||
}
|
||||
#[test]
|
||||
fn review_submission_authorization_rejects_invalid_grants_before_side_effects() {
|
||||
let (_d, store) = fixture();
|
||||
open(&store);
|
||||
request(&store, "published-source", "valid-token");
|
||||
|
||||
let invalid = store
|
||||
.authorize_review_submission("T", "invalid-token")
|
||||
.unwrap_err();
|
||||
assert!(matches!(invalid, MergeRequestError::Unauthorized(_)));
|
||||
let authorized = store
|
||||
.authorize_review_submission("T", "valid-token")
|
||||
.unwrap();
|
||||
assert_eq!(authorized.workspace_id, "W");
|
||||
assert_eq!(authorized.subject_ref, "published-source");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn selectors_thread_and_completion_have_no_revision_or_commit_api() {
|
||||
let (d, s) = fixture();
|
||||
@@ -284,21 +301,13 @@ fn review_revocation_invalidates_readiness() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn legacy_v11_migration_preserves_review_events_and_replaces_marker() {
|
||||
fn fresh_schema_uses_version_12_and_reopens_as_current() {
|
||||
let c = Connection::open_in_memory().unwrap();
|
||||
c.execute_batch("CREATE TABLE repositories(workspace_id TEXT,repository_id TEXT,PRIMARY KEY(workspace_id,repository_id));CREATE TABLE typed_tickets(workspace_id TEXT,ticket_id TEXT,PRIMARY KEY(workspace_id,ticket_id));INSERT INTO repositories VALUES('W','R');INSERT INTO typed_tickets VALUES('W','T');CREATE TABLE merge_request_schema_migrations(version INTEGER PRIMARY KEY,applied_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP);INSERT INTO merge_request_schema_migrations(version) VALUES(11);CREATE TABLE merge_requests(workspace_id TEXT,merge_request_id TEXT,repository_id TEXT,state TEXT,target_ref_selector TEXT,current_revision_ordinal INTEGER,current_revision_id TEXT,created_at TEXT,updated_at TEXT,merged_revision_id TEXT,merged_at TEXT);CREATE TABLE merge_request_ticket_relations(workspace_id TEXT,merge_request_id TEXT,ticket_id TEXT,relation_kind TEXT,created_at TEXT);CREATE TABLE merge_request_revisions(workspace_id TEXT,merge_request_id TEXT,revision_id TEXT,ordinal INTEGER,base_commit TEXT,head_commit TEXT,diff_digest TEXT,summary TEXT,assignment_id TEXT,created_at TEXT);CREATE TABLE merge_request_revision_paths(workspace_id TEXT,merge_request_id TEXT,revision_id TEXT,ordinal INTEGER,path TEXT);CREATE TABLE merge_request_reviewer_child_sessions(workspace_id TEXT,child_session_id TEXT,parent_runtime_id TEXT,parent_worker_id TEXT,reviewer_profile TEXT,registered_at TEXT);CREATE TABLE merge_request_review_attempts(workspace_id TEXT,attempt_id TEXT,merge_request_id TEXT,ticket_id TEXT,revision_id TEXT,revision_ordinal INTEGER,parent_assignment_id TEXT,parent_runtime_id TEXT,parent_worker_id TEXT,child_session_id TEXT,reviewer_effective_profile TEXT,capability_token TEXT,status TEXT,created_at TEXT,consumed_at TEXT);CREATE TABLE merge_request_reviews(workspace_id TEXT,attempt_id TEXT,merge_request_id TEXT,revision_id TEXT,decision TEXT,body TEXT,submitted_at TEXT);CREATE TABLE merge_request_review_findings(workspace_id TEXT,attempt_id TEXT,ordinal INTEGER,severity TEXT,code TEXT,path TEXT,line INTEGER,body TEXT);CREATE TABLE merge_request_completion_operations(workspace_id TEXT,operation_id TEXT,ticket_id TEXT,revision_id TEXT,authority_kind TEXT,implementation_assignment_id TEXT,completion_actor_runtime_id TEXT,completion_actor_worker_id TEXT,target_commit TEXT,source_commit TEXT,result_commit TEXT,strategy TEXT,resolution TEXT,fingerprint TEXT,status TEXT,result_ticket_state TEXT,created_at TEXT,updated_at TEXT);INSERT INTO merge_requests VALUES('W','MR','R','open','develop',1,'V','2026-07-26T12:00:00Z','2026-07-26T12:00:00Z',NULL,NULL);INSERT INTO merge_request_ticket_relations VALUES('W','MR','T','implements','2026-07-26T12:00:00Z');INSERT INTO merge_request_revisions VALUES('W','MR','V',1,'base','subject','digest','summary','A','2026-07-26T12:00:00Z');INSERT INTO merge_request_review_attempts VALUES('W','AT','MR','T','V',1,'A','runtime','coder','child','builtin:reviewer','token','submitted','2026-07-26T12:00:00Z','2026-07-26T12:00:01Z');INSERT INTO merge_request_reviews VALUES('W','AT','MR','V','approve','approved','2026-07-26T12:00:01Z');INSERT INTO merge_request_review_attempts VALUES('W','PENDING','MR','T','V',1,'A','runtime','coder','pending-child','builtin:reviewer','pending-token','registered','2026-07-26T12:00:02Z',NULL);").unwrap();
|
||||
c.execute_batch(
|
||||
"CREATE TABLE unrelated_parent(left_id TEXT,right_id TEXT,PRIMARY KEY(left_id,right_id));CREATE TABLE unrelated_child(left_id TEXT REFERENCES unrelated_parent(left_id));",
|
||||
"CREATE TABLE repositories(workspace_id TEXT,repository_id TEXT,PRIMARY KEY(workspace_id,repository_id));CREATE TABLE typed_tickets(workspace_id TEXT,ticket_id TEXT,PRIMARY KEY(workspace_id,ticket_id));",
|
||||
)
|
||||
.unwrap();
|
||||
let unrelated_mismatch = c
|
||||
.query_row("PRAGMA foreign_key_check", [], |_| Ok(()))
|
||||
.unwrap_err();
|
||||
assert!(
|
||||
unrelated_mismatch
|
||||
.to_string()
|
||||
.contains("foreign key mismatch")
|
||||
);
|
||||
|
||||
merge_request::migrate(&c).unwrap();
|
||||
assert_eq!(
|
||||
c.query_row("SELECT version FROM merge_request_schema", [], |r| {
|
||||
@@ -307,66 +316,26 @@ fn legacy_v11_migration_preserves_review_events_and_replaces_marker() {
|
||||
.unwrap(),
|
||||
12
|
||||
);
|
||||
let legacy_marker: bool = c
|
||||
.query_row(
|
||||
"SELECT EXISTS(SELECT 1 FROM sqlite_master WHERE type='table' AND name='merge_request_schema_migrations')",
|
||||
[],
|
||||
|r| r.get(0),
|
||||
)
|
||||
.unwrap();
|
||||
assert!(!legacy_marker);
|
||||
let selector: Option<String> = c
|
||||
.query_row("SELECT selector_from FROM merge_requests", [], |r| r.get(0))
|
||||
.unwrap();
|
||||
assert!(selector.is_none());
|
||||
let kinds: String = c
|
||||
.query_row(
|
||||
"SELECT group_concat(kind,',') FROM merge_request_thread_events ORDER BY sequence",
|
||||
[],
|
||||
|r| r.get(0),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
kinds,
|
||||
"review_requested,review,review_requested,review_cancelled"
|
||||
);
|
||||
let old: bool = c
|
||||
.query_row(
|
||||
"SELECT EXISTS(SELECT 1 FROM sqlite_master WHERE name='merge_request_revisions')",
|
||||
[],
|
||||
|r| r.get(0),
|
||||
)
|
||||
.unwrap();
|
||||
assert!(!old);
|
||||
merge_request::migrate(&c).unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn failed_legacy_v11_migration_rolls_back_marker_bridge() {
|
||||
fn current_schema_validation_rejects_missing_tables() {
|
||||
let c = Connection::open_in_memory().unwrap();
|
||||
c.execute_batch(
|
||||
"CREATE TABLE merge_request_schema_migrations(version INTEGER PRIMARY KEY,applied_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP);INSERT INTO merge_request_schema_migrations(version) VALUES(11);CREATE TABLE merge_requests(merge_request_id TEXT);",
|
||||
"CREATE TABLE repositories(workspace_id TEXT,repository_id TEXT,PRIMARY KEY(workspace_id,repository_id));CREATE TABLE typed_tickets(workspace_id TEXT,ticket_id TEXT,PRIMARY KEY(workspace_id,ticket_id));",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert!(merge_request::migrate(&c).is_err());
|
||||
for table in ["merge_request_schema_migrations", "merge_requests"] {
|
||||
let exists: bool = c
|
||||
.query_row(
|
||||
"SELECT EXISTS(SELECT 1 FROM sqlite_master WHERE type='table' AND name=?1)",
|
||||
[table],
|
||||
|r| r.get(0),
|
||||
)
|
||||
.unwrap();
|
||||
assert!(exists, "{table} was not rolled back");
|
||||
}
|
||||
let current_marker: bool = c
|
||||
.query_row(
|
||||
"SELECT EXISTS(SELECT 1 FROM sqlite_master WHERE type='table' AND name='merge_request_schema')",
|
||||
[],
|
||||
|r| r.get(0),
|
||||
)
|
||||
merge_request::migrate(&c).unwrap();
|
||||
c.execute_batch("DROP TABLE merge_request_review_grants;")
|
||||
.unwrap();
|
||||
assert!(!current_marker);
|
||||
|
||||
let error = merge_request::migrate(&c).unwrap_err();
|
||||
assert!(matches!(
|
||||
error,
|
||||
MergeRequestError::Corrupt(message)
|
||||
if message == "missing `merge_request_review_grants`"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -14,6 +14,7 @@ json-schema = ["dep:schemars"]
|
||||
schemars = { workspace = true, optional = true }
|
||||
serde = { workspace = true, features = ["derive"] }
|
||||
serde_json = { workspace = true }
|
||||
sha2.workspace = true
|
||||
tokio = { workspace = true, features = ["io-util"], optional = true }
|
||||
ts-rs = { version = "12.0.1", optional = true }
|
||||
uuid = { workspace = true, features = ["serde"] }
|
||||
uuid = { workspace = true, features = ["serde", "v7"] }
|
||||
|
||||
@@ -0,0 +1,132 @@
|
||||
use std::{fmt, str::FromStr};
|
||||
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer, de};
|
||||
use sha2::{Digest, Sha256};
|
||||
use uuid::{Uuid, Version};
|
||||
|
||||
/// Stable Worker identity independent of its current Runtime placement or
|
||||
/// conversation Session.
|
||||
///
|
||||
/// Workspace authority allocates this ID for managed Workers. A standalone
|
||||
/// Worker store allocates it locally when no Workspace authority is present.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
|
||||
pub struct WorkerId(Uuid);
|
||||
|
||||
impl WorkerId {
|
||||
pub fn now_v7() -> Self {
|
||||
Self(Uuid::now_v7())
|
||||
}
|
||||
|
||||
/// Converts a legacy Runtime-local numeric id into a syntactically valid
|
||||
/// migration-only UUIDv7 value. New Worker allocation must use `now_v7`.
|
||||
pub fn from_legacy_u64(value: u64) -> Self {
|
||||
let mut bytes = [0_u8; 16];
|
||||
bytes[8..].copy_from_slice(&value.to_be_bytes());
|
||||
bytes[6] = 0x70;
|
||||
bytes[8] = (bytes[8] & 0x3f) | 0x80;
|
||||
Self(Uuid::from_bytes(bytes))
|
||||
}
|
||||
|
||||
pub fn from_legacy_binding(workspace_id: &str, runtime_id: &str, value: u64) -> Self {
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(b"yoi.workspace-worker-id.v1\0");
|
||||
hasher.update(workspace_id.as_bytes());
|
||||
hasher.update([0]);
|
||||
hasher.update(runtime_id.as_bytes());
|
||||
hasher.update([0]);
|
||||
hasher.update(value.to_be_bytes());
|
||||
let digest = hasher.finalize();
|
||||
let mut bytes = [0_u8; 16];
|
||||
bytes.copy_from_slice(&digest[..16]);
|
||||
// Migrated ids sort before normally allocated UUIDv7 values while retaining
|
||||
// deterministic collision-resistant payload bits.
|
||||
bytes[..6].fill(0);
|
||||
bytes[6] = (bytes[6] & 0x0f) | 0x70;
|
||||
bytes[8] = (bytes[8] & 0x3f) | 0x80;
|
||||
Self(Uuid::from_bytes(bytes))
|
||||
}
|
||||
|
||||
pub fn parse(value: &str) -> Option<Self> {
|
||||
let value = Uuid::parse_str(value).ok()?;
|
||||
(value.get_version() == Some(Version::SortRand)).then_some(Self(value))
|
||||
}
|
||||
|
||||
pub const fn as_uuid(self) -> Uuid {
|
||||
self.0
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn short(self) -> String {
|
||||
let simple = self.0.simple().to_string();
|
||||
simple[simple.len() - 12..].to_string()
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for WorkerId {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
self.0.fmt(formatter)
|
||||
}
|
||||
}
|
||||
|
||||
impl FromStr for WorkerId {
|
||||
type Err = WorkerIdParseError;
|
||||
|
||||
fn from_str(value: &str) -> Result<Self, Self::Err> {
|
||||
Self::parse(value).ok_or(WorkerIdParseError)
|
||||
}
|
||||
}
|
||||
|
||||
impl Serialize for WorkerId {
|
||||
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
serializer.serialize_str(&self.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for WorkerId {
|
||||
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
|
||||
where
|
||||
D: Deserializer<'de>,
|
||||
{
|
||||
let value = String::deserialize(deserializer)?;
|
||||
Self::parse(&value).ok_or_else(|| de::Error::custom("Worker id must be a UUIDv7"))
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct WorkerIdParseError;
|
||||
|
||||
impl fmt::Display for WorkerIdParseError {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter.write_str("Worker id must be a UUIDv7")
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for WorkerIdParseError {}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn worker_id_accepts_only_uuid_v7() {
|
||||
let worker_id = WorkerId::now_v7();
|
||||
assert_eq!(WorkerId::parse(&worker_id.to_string()), Some(worker_id));
|
||||
assert!(WorkerId::parse("30").is_none());
|
||||
assert!(WorkerId::parse(&Uuid::nil().to_string()).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn legacy_worker_id_mapping_is_stable() {
|
||||
assert_eq!(
|
||||
WorkerId::from_legacy_binding("workspace", "runtime", 42),
|
||||
WorkerId::from_legacy_binding("workspace", "runtime", 42)
|
||||
);
|
||||
assert_ne!(
|
||||
WorkerId::from_legacy_binding("workspace", "runtime", 42),
|
||||
WorkerId::from_legacy_binding("workspace", "runtime", 43)
|
||||
);
|
||||
}
|
||||
}
|
||||
+331
-29
@@ -1,3 +1,4 @@
|
||||
pub mod identity;
|
||||
#[cfg(feature = "stream")]
|
||||
pub mod stream;
|
||||
pub mod subscription;
|
||||
@@ -8,6 +9,8 @@ use std::path::PathBuf;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
pub use identity::{WorkerId, WorkerIdParseError};
|
||||
|
||||
fn default_true() -> bool {
|
||||
true
|
||||
}
|
||||
@@ -190,6 +193,106 @@ impl WorkerEvent {
|
||||
/// variants — emits an alert and inserts a `[unknown input segment]`
|
||||
/// placeholder into the LLM context so neither user nor LLM is blind to
|
||||
/// the dropped intent.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[cfg_attr(feature = "json-schema", derive(schemars::JsonSchema))]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum PasteArtifactMediaType {
|
||||
TextPlainUtf8,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[cfg_attr(feature = "json-schema", derive(schemars::JsonSchema))]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum PasteArtifactAvailability {
|
||||
Available,
|
||||
Unavailable,
|
||||
IntegrityFailed,
|
||||
}
|
||||
|
||||
impl PasteArtifactMediaType {
|
||||
pub fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::TextPlainUtf8 => "text/plain; charset=utf-8",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl PasteArtifactAvailability {
|
||||
pub fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Available => "available",
|
||||
Self::Unavailable => "unavailable",
|
||||
Self::IntegrityFailed => "integrity_failed",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Session-owned reference to a large pasted-input artifact.
|
||||
///
|
||||
/// The reference contains only bounded integrity and provenance metadata. The
|
||||
/// artifact body remains in session storage and is available to the model only
|
||||
/// through the scoped paste-artifact tools installed by Worker.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[cfg_attr(feature = "json-schema", derive(schemars::JsonSchema))]
|
||||
pub struct PasteArtifactRef {
|
||||
pub artifact_id: String,
|
||||
pub created_at_ms: u64,
|
||||
pub media_type: PasteArtifactMediaType,
|
||||
/// Availability observed when this immutable reference was committed.
|
||||
/// Reads revalidate storage and integrity rather than trusting this field.
|
||||
pub availability: PasteArtifactAvailability,
|
||||
pub byte_len: u64,
|
||||
pub char_count: u64,
|
||||
pub line_count: u64,
|
||||
pub sha256: String,
|
||||
pub source_entry_id: String,
|
||||
}
|
||||
|
||||
/// Availability recorded for an uploaded client-local file.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[cfg_attr(feature = "json-schema", derive(schemars::JsonSchema))]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum UploadedFileAvailability {
|
||||
Available,
|
||||
Unavailable,
|
||||
IntegrityFailed,
|
||||
}
|
||||
|
||||
impl UploadedFileAvailability {
|
||||
pub fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Available => "available",
|
||||
Self::Unavailable => "unavailable",
|
||||
Self::IntegrityFailed => "integrity_failed",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Session-owned immutable reference to a client-local uploaded file.
|
||||
///
|
||||
/// Upload transports return an unbound reference. Worker fills
|
||||
/// `source_entry_id` immediately before the containing user input is committed;
|
||||
/// committed Session Log and public snapshot records therefore always retain
|
||||
/// the durable source-entry identity without storing the file body.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[cfg_attr(feature = "json-schema", derive(schemars::JsonSchema))]
|
||||
pub struct UploadedFileRef {
|
||||
pub artifact_id: String,
|
||||
pub file_name: String,
|
||||
pub media_type: String,
|
||||
pub created_at_ms: u64,
|
||||
pub availability: UploadedFileAvailability,
|
||||
pub byte_len: u64,
|
||||
pub sha256: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub source_entry_id: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[cfg_attr(feature = "json-schema", derive(schemars::JsonSchema))]
|
||||
@@ -207,6 +310,14 @@ pub enum Segment {
|
||||
lines: u32,
|
||||
content: String,
|
||||
},
|
||||
/// Internal reference produced when Worker stores a large `Paste` before
|
||||
/// committing input. Clients may receive this in history/event projections;
|
||||
/// the body is intentionally absent.
|
||||
PasteArtifact { artifact: PasteArtifactRef },
|
||||
/// Client-local file uploaded into the owning Worker session before submit.
|
||||
/// The Session Log stores only this immutable reference, never file bytes or
|
||||
/// the client's local path.
|
||||
UploadedFile { file: UploadedFileRef },
|
||||
/// `@<path>` file-system reference. Worker resolves readable files to
|
||||
/// `[File: <path>]` attachments and readable normal directories to shallow
|
||||
/// `[Dir: <path>]` listings; the flattened user text keeps the literal
|
||||
@@ -247,6 +358,35 @@ impl Segment {
|
||||
match seg {
|
||||
Segment::Text { content } => out.push_str(content),
|
||||
Segment::Paste { content, .. } => out.push_str(content),
|
||||
Segment::PasteArtifact { artifact } => {
|
||||
use std::fmt::Write as _;
|
||||
let _ = write!(
|
||||
out,
|
||||
"[Large paste stored as artifact {}: {} bytes, {} chars, {} lines, {}, {}, created at {} ms, sha256 {}; use SearchInputArtifact and ReadInputArtifact to inspect it]",
|
||||
artifact.artifact_id,
|
||||
artifact.byte_len,
|
||||
artifact.char_count,
|
||||
artifact.line_count,
|
||||
artifact.media_type.as_str(),
|
||||
artifact.availability.as_str(),
|
||||
artifact.created_at_ms,
|
||||
artifact.sha256
|
||||
);
|
||||
}
|
||||
Segment::UploadedFile { file } => {
|
||||
use std::fmt::Write as _;
|
||||
let _ = write!(
|
||||
out,
|
||||
"[Attached file {} stored as input artifact {}: {} bytes, {}, {}, created at {} ms, sha256 {}; use SearchInputArtifact and ReadInputArtifact for supported text content]",
|
||||
file.file_name,
|
||||
file.artifact_id,
|
||||
file.byte_len,
|
||||
file.media_type,
|
||||
file.availability.as_str(),
|
||||
file.created_at_ms,
|
||||
file.sha256
|
||||
);
|
||||
}
|
||||
Segment::FileRef { path } => {
|
||||
out.push('@');
|
||||
out.push_str(path);
|
||||
@@ -340,8 +480,7 @@ pub struct InternalWorkerRef {
|
||||
pub struct InternalWorkerSnapshot {
|
||||
pub worker: InternalWorkerRef,
|
||||
pub revision: u64,
|
||||
#[cfg_attr(feature = "typescript", ts(type = "Array<unknown>"))]
|
||||
pub entries: Vec<serde_json::Value>,
|
||||
pub session: SessionSnapshot,
|
||||
#[serde(default)]
|
||||
pub status: WorkerStatus,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
@@ -364,12 +503,114 @@ pub enum ToolResultDisposition {
|
||||
OutcomeUnknown,
|
||||
}
|
||||
|
||||
/// Canonical, storage-independent projection of committed session history.
|
||||
///
|
||||
/// Worker protocols expose this DTO instead of append-log records. New
|
||||
/// storage variants can therefore be added without teaching every client how
|
||||
/// to replay the durable log format.
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
pub struct SessionSnapshot {
|
||||
pub entries: Vec<SessionSnapshotEntry>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum SessionEntryProvenance {
|
||||
HumanInput,
|
||||
WorkerInput,
|
||||
FlowInstruction,
|
||||
BackendInstruction,
|
||||
ModelOutput,
|
||||
ToolOutput,
|
||||
DerivedSummary,
|
||||
LegacyUnknown,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
pub struct SessionSnapshotEntry {
|
||||
/// Stable identity from durable history metadata, or a deterministic
|
||||
/// identity derived from the legacy segment and log position.
|
||||
pub entry_id: String,
|
||||
/// Timestamp copied from the durable log record that commits this entry.
|
||||
pub timestamp: u64,
|
||||
pub provenance: SessionEntryProvenance,
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
pub derived_from: Vec<String>,
|
||||
#[serde(flatten)]
|
||||
pub data: SessionSnapshotEntryData,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[serde(tag = "kind", rename_all = "snake_case")]
|
||||
pub enum SessionSnapshotEntryData {
|
||||
UserInput {
|
||||
segments: Vec<Segment>,
|
||||
},
|
||||
Message {
|
||||
role: SessionMessageRole,
|
||||
content: Vec<SessionContentPart>,
|
||||
},
|
||||
ToolCall {
|
||||
call_id: String,
|
||||
name: String,
|
||||
arguments: String,
|
||||
},
|
||||
ToolResult {
|
||||
call_id: String,
|
||||
summary: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
content: Option<String>,
|
||||
is_error: bool,
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
attachments: Vec<SessionToolAttachment>,
|
||||
},
|
||||
SystemItem {
|
||||
item_kind: String,
|
||||
content: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
#[cfg_attr(feature = "typescript", ts(type = "unknown"))]
|
||||
data: Option<serde_json::Value>,
|
||||
},
|
||||
RunError {
|
||||
message: String,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum SessionMessageRole {
|
||||
User,
|
||||
Assistant,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[serde(tag = "kind", rename_all = "snake_case")]
|
||||
pub enum SessionContentPart {
|
||||
Text { text: String },
|
||||
Refusal { refusal: String },
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
pub struct SessionToolAttachment {
|
||||
pub media_type: String,
|
||||
/// Base64-encoded durable attachment body. Public snapshots preserve the
|
||||
/// committed multimodal value instead of replacing it with placeholder text.
|
||||
pub data_base64: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[serde(tag = "event", content = "data", rename_all = "snake_case")]
|
||||
pub enum Event {
|
||||
/// A user input message was accepted, persisted as
|
||||
/// `LogEntry::UserInput`, and is about to start a new turn.
|
||||
/// `LogEntry::AnnotatedUserInput`, and is about to start a new turn.
|
||||
/// Broadcast to every subscribed client so TUI / GUI instances show
|
||||
/// the same user line that reconnect snapshots would replay from
|
||||
/// history; clients must not synthesize a separate pending/fake
|
||||
@@ -390,7 +631,7 @@ pub enum Event {
|
||||
/// of parsing free-text prefixes like `[Notification] …` or
|
||||
/// `[File: …]`.
|
||||
///
|
||||
/// One event per `LogEntry::SystemItem` commit. Disk-side and
|
||||
/// One event per `LogEntry::AnnotatedSystemItem` commit. Disk-side and
|
||||
/// wire-side are 1:1.
|
||||
SystemItem {
|
||||
#[cfg_attr(feature = "typescript", ts(type = "unknown"))]
|
||||
@@ -555,8 +796,7 @@ pub enum Event {
|
||||
/// role-specific entry events (`SegmentRotated` / `SystemItem`) —
|
||||
/// there is no generic "every committed entry" broadcast.
|
||||
Snapshot {
|
||||
#[cfg_attr(feature = "typescript", ts(type = "Array<unknown>"))]
|
||||
entries: Vec<serde_json::Value>,
|
||||
session: SessionSnapshot,
|
||||
greeting: Greeting,
|
||||
#[serde(default)]
|
||||
status: WorkerStatus,
|
||||
@@ -589,14 +829,10 @@ pub enum Event {
|
||||
/// Server-side segment log rotated to a fresh `SegmentStart`.
|
||||
///
|
||||
/// Fires on compaction and on auto-fork when the store head drifts
|
||||
/// from the live writer's cached head. Clients drop their derived
|
||||
/// view and reseed from `entry.history` exactly the way they would
|
||||
/// from a connect-time `Snapshot`.
|
||||
///
|
||||
/// Payload is the JSON form of `session_store::LogEntry::SegmentStart`.
|
||||
/// A compaction/fork has replaced the authoritative segment. Clients drop
|
||||
/// their derived view and reseed from the canonical committed snapshot.
|
||||
SegmentRotated {
|
||||
#[cfg_attr(feature = "typescript", ts(type = "unknown"))]
|
||||
entry: serde_json::Value,
|
||||
session: SessionSnapshot,
|
||||
},
|
||||
/// Current Worker controller status. Broadcast on every controller-level
|
||||
/// transition and included in `History` snapshots for late attach.
|
||||
@@ -623,11 +859,10 @@ pub enum Event {
|
||||
head_entries: usize,
|
||||
targets: Vec<RewindTarget>,
|
||||
},
|
||||
/// A rewind has truncated the authoritative session. `entries` is the
|
||||
/// retained session-log prefix clients should use to reseed display state.
|
||||
/// A rewind has truncated the authoritative session. `session` is the
|
||||
/// retained canonical snapshot clients should use to reseed display state.
|
||||
RewindApplied {
|
||||
#[cfg_attr(feature = "typescript", ts(type = "Array<unknown>"))]
|
||||
entries: Vec<serde_json::Value>,
|
||||
session: SessionSnapshot,
|
||||
input: Vec<Segment>,
|
||||
summary: RewindSummary,
|
||||
},
|
||||
@@ -1104,6 +1339,55 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn paste_artifact_segment_roundtrips_without_body() {
|
||||
let artifact = PasteArtifactRef {
|
||||
artifact_id: "019ca7c8-57b6-7f05-8edf-524147aba7b2".to_string(),
|
||||
created_at_ms: 1_700_000_000_000,
|
||||
media_type: PasteArtifactMediaType::TextPlainUtf8,
|
||||
availability: PasteArtifactAvailability::Available,
|
||||
byte_len: 65_536,
|
||||
char_count: 65_530,
|
||||
line_count: 200,
|
||||
sha256: "a".repeat(64),
|
||||
source_entry_id: "entry-1".to_string(),
|
||||
};
|
||||
let segment = Segment::PasteArtifact {
|
||||
artifact: artifact.clone(),
|
||||
};
|
||||
let json = serde_json::to_string(&segment).unwrap();
|
||||
assert!(!json.contains("pasted body"));
|
||||
assert_eq!(serde_json::from_str::<Segment>(&json).unwrap(), segment);
|
||||
let projected = Segment::flatten_to_text(&[segment]);
|
||||
assert!(projected.contains(&artifact.artifact_id));
|
||||
assert!(projected.contains("SearchInputArtifact"));
|
||||
assert!(projected.contains("ReadInputArtifact"));
|
||||
assert!(!projected.contains("pasted body"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn uploaded_file_segment_roundtrips_without_path_or_body() {
|
||||
let file = UploadedFileRef {
|
||||
artifact_id: "019ca7c8-57b6-7f05-8edf-524147aba7b3".to_string(),
|
||||
file_name: "notes.md".to_string(),
|
||||
media_type: "text/markdown".to_string(),
|
||||
created_at_ms: 1_700_000_000_001,
|
||||
availability: UploadedFileAvailability::Available,
|
||||
byte_len: 128,
|
||||
sha256: "b".repeat(64),
|
||||
source_entry_id: Some("entry-2".to_string()),
|
||||
};
|
||||
let segment = Segment::UploadedFile { file: file.clone() };
|
||||
let json = serde_json::to_string(&segment).unwrap();
|
||||
assert!(!json.contains("/home/user/private"));
|
||||
assert!(!json.contains("file body"));
|
||||
assert_eq!(serde_json::from_str::<Segment>(&json).unwrap(), segment);
|
||||
let projected = Segment::flatten_to_text(&[segment]);
|
||||
assert!(projected.contains("notes.md"));
|
||||
assert!(projected.contains(&file.artifact_id));
|
||||
assert!(projected.contains("ReadInputArtifact"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn method_run_flow_segment_roundtrip() {
|
||||
let method = Method::Run {
|
||||
@@ -1440,7 +1724,17 @@ mod tests {
|
||||
#[test]
|
||||
fn event_snapshot_format() {
|
||||
let event = Event::Snapshot {
|
||||
entries: vec![serde_json::json!({"kind": "user_input", "ts": 1, "segments": []})],
|
||||
session: SessionSnapshot {
|
||||
entries: vec![SessionSnapshotEntry {
|
||||
entry_id: "entry-1".into(),
|
||||
timestamp: 1,
|
||||
provenance: SessionEntryProvenance::HumanInput,
|
||||
derived_from: Vec::new(),
|
||||
data: SessionSnapshotEntryData::UserInput {
|
||||
segments: Vec::new(),
|
||||
},
|
||||
}],
|
||||
},
|
||||
greeting: Greeting {
|
||||
worker_name: "test".into(),
|
||||
cwd: "/tmp".into(),
|
||||
@@ -1458,8 +1752,12 @@ mod tests {
|
||||
let json = serde_json::to_string(&event).unwrap();
|
||||
let parsed: serde_json::Value = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(parsed["event"], "snapshot");
|
||||
assert!(parsed["data"]["entries"].is_array());
|
||||
assert_eq!(parsed["data"]["entries"][0]["kind"], "user_input");
|
||||
assert!(parsed["data"]["session"]["entries"].is_array());
|
||||
assert_eq!(
|
||||
parsed["data"]["session"]["entries"][0]["kind"],
|
||||
"user_input"
|
||||
);
|
||||
assert_eq!(parsed["data"]["session"]["entries"][0]["timestamp"], 1);
|
||||
assert_eq!(parsed["data"]["greeting"]["worker_name"], "test");
|
||||
assert_eq!(parsed["data"]["greeting"]["tools"][0], "Read");
|
||||
assert_eq!(parsed["data"]["greeting"]["context_window"], 200_000);
|
||||
@@ -1469,7 +1767,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn event_snapshot_in_flight_roundtrip_and_default() {
|
||||
let inbound = r#"{"event":"snapshot","data":{"entries":[],"greeting":{"worker_name":"test","cwd":"/tmp","provider":"p","model":"m","scope_summary":"s","tools":[]},"status":"running"}}"#;
|
||||
let inbound = r#"{"event":"snapshot","data":{"session":{"entries":[]},"greeting":{"worker_name":"test","cwd":"/tmp","provider":"p","model":"m","scope_summary":"s","tools":[]},"status":"running"}}"#;
|
||||
let decoded: Event = serde_json::from_str(inbound).unwrap();
|
||||
match decoded {
|
||||
Event::Snapshot { in_flight, .. } => assert!(in_flight.is_empty()),
|
||||
@@ -1477,7 +1775,9 @@ mod tests {
|
||||
}
|
||||
|
||||
let event = Event::Snapshot {
|
||||
entries: Vec::new(),
|
||||
session: SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
greeting: Greeting {
|
||||
worker_name: "test".into(),
|
||||
cwd: "/tmp".into(),
|
||||
@@ -1543,15 +1843,17 @@ mod tests {
|
||||
#[test]
|
||||
fn event_segment_rotated_roundtrip() {
|
||||
let event = Event::SegmentRotated {
|
||||
entry: serde_json::json!({"kind": "segment_start", "ts": 1, "history": []}),
|
||||
session: SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
};
|
||||
let json = serde_json::to_string(&event).unwrap();
|
||||
let parsed: serde_json::Value = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(parsed["event"], "segment_rotated");
|
||||
assert_eq!(parsed["data"]["entry"]["kind"], "segment_start");
|
||||
assert!(parsed["data"]["session"]["entries"].is_array());
|
||||
let decoded: Event = serde_json::from_str(&json).unwrap();
|
||||
match decoded {
|
||||
Event::SegmentRotated { entry } => assert_eq!(entry["kind"], "segment_start"),
|
||||
Event::SegmentRotated { session } => assert!(session.entries.is_empty()),
|
||||
other => panic!("expected SegmentRotated, got {other:?}"),
|
||||
}
|
||||
}
|
||||
@@ -1627,8 +1929,8 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn event_snapshot_legacy_without_status_defaults_to_idle() {
|
||||
let json = r#"{"event":"snapshot","data":{"entries":[],"greeting":{"worker_name":"test","cwd":"/tmp","provider":"anthropic","model":"claude","scope_summary":"","tools":[]}}}"#;
|
||||
fn event_snapshot_without_status_defaults_to_idle() {
|
||||
let json = r#"{"event":"snapshot","data":{"session":{"entries":[]},"greeting":{"worker_name":"test","cwd":"/tmp","provider":"anthropic","model":"claude","scope_summary":"","tools":[]}}}"#;
|
||||
let decoded: Event = serde_json::from_str(json).unwrap();
|
||||
match decoded {
|
||||
Event::Snapshot {
|
||||
@@ -2039,11 +2341,11 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn legacy_snapshot_defaults_internal_workers_to_empty() {
|
||||
fn snapshot_defaults_internal_workers_to_empty() {
|
||||
let snapshot: Event = serde_json::from_value(serde_json::json!({
|
||||
"event": "snapshot",
|
||||
"data": {
|
||||
"entries": [],
|
||||
"session": { "entries": [] },
|
||||
"greeting": {
|
||||
"worker_name": "parent",
|
||||
"cwd": ".",
|
||||
|
||||
@@ -170,6 +170,23 @@ fn validate_identifier(
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_repository_key(value: &str) -> Result<(), SubscriptionValidationError> {
|
||||
let bytes = value.as_bytes();
|
||||
if bytes.is_empty()
|
||||
|| bytes.len() > 64
|
||||
|| bytes.first() == Some(&b'-')
|
||||
|| bytes.last() == Some(&b'-')
|
||||
|| !bytes
|
||||
.iter()
|
||||
.all(|byte| byte.is_ascii_lowercase() || byte.is_ascii_digit() || *byte == b'-')
|
||||
{
|
||||
return Err(SubscriptionValidationError::InvalidIdentifier {
|
||||
field: "repository_key",
|
||||
});
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_rejection_message(message: &str) -> Result<(), SubscriptionValidationError> {
|
||||
if message.is_empty() {
|
||||
return Err(SubscriptionValidationError::EmptyRejectionMessage);
|
||||
@@ -540,7 +557,6 @@ pub enum SubscriptionWorkerState {
|
||||
Running,
|
||||
Paused,
|
||||
Stopped,
|
||||
Cancelled,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
@@ -567,7 +583,12 @@ pub struct SubscriptionWorker {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub profile: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
#[cfg_attr(feature = "typescript", ts(skip))]
|
||||
pub repository_id: Option<String>,
|
||||
/// Workspace-facing Repository key. Runtime producers leave this unset and
|
||||
/// Workspace Server projections replace `repository_id` with this field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub repository_key: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub working_directory_id: Option<SubscriptionWorkdirId>,
|
||||
}
|
||||
@@ -584,6 +605,14 @@ impl SubscriptionWorker {
|
||||
if let Some(repository_id) = &self.repository_id {
|
||||
validate_identifier("repository_id", repository_id, MAX_RESOURCE_ID_BYTES)?;
|
||||
}
|
||||
if let Some(repository_key) = &self.repository_key {
|
||||
validate_repository_key(repository_key)?;
|
||||
}
|
||||
if self.repository_id.is_some() && self.repository_key.is_some() {
|
||||
return Err(SubscriptionValidationError::InvalidIdentifier {
|
||||
field: "repository_authority",
|
||||
});
|
||||
}
|
||||
if let Some(working_directory_id) = &self.working_directory_id {
|
||||
working_directory_id.validate()?;
|
||||
}
|
||||
@@ -595,7 +624,13 @@ impl SubscriptionWorker {
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
pub struct SubscriptionWorkdir {
|
||||
pub working_directory_id: SubscriptionWorkdirId,
|
||||
pub repository_id: String,
|
||||
/// Runtime-internal Repository id. Workspace-facing TypeScript contracts
|
||||
/// omit this field and require `repository_key` from the Server projection.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
#[cfg_attr(feature = "typescript", ts(skip))]
|
||||
pub repository_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub repository_key: Option<String>,
|
||||
pub state: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub primary_worker_id: Option<SubscriptionWorkerId>,
|
||||
@@ -604,7 +639,41 @@ pub struct SubscriptionWorkdir {
|
||||
impl SubscriptionWorkdir {
|
||||
pub fn validate(&self) -> Result<(), SubscriptionValidationError> {
|
||||
self.working_directory_id.validate()?;
|
||||
validate_identifier("repository_id", &self.repository_id, MAX_RESOURCE_ID_BYTES)?;
|
||||
match (&self.repository_id, &self.repository_key) {
|
||||
(Some(repository_id), None) => {
|
||||
validate_identifier("repository_id", repository_id, MAX_RESOURCE_ID_BYTES)?;
|
||||
}
|
||||
(None, Some(repository_key)) => validate_repository_key(repository_key)?,
|
||||
_ => {
|
||||
return Err(SubscriptionValidationError::InvalidIdentifier {
|
||||
field: "repository_authority",
|
||||
});
|
||||
}
|
||||
}
|
||||
validate_identifier("workdir_state", &self.state, MAX_RESOURCE_ID_BYTES)?;
|
||||
if let Some(worker_id) = &self.primary_worker_id {
|
||||
worker_id.validate()?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// Workspace-facing Workdir summary. Backend-generated Repository UUIDs never
|
||||
/// enter this DTO; Workspace Server must resolve the required Repository key.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
pub struct WorkspaceSubscriptionWorkdir {
|
||||
pub working_directory_id: SubscriptionWorkdirId,
|
||||
pub repository_key: String,
|
||||
pub state: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub primary_worker_id: Option<SubscriptionWorkerId>,
|
||||
}
|
||||
|
||||
impl WorkspaceSubscriptionWorkdir {
|
||||
pub fn validate(&self) -> Result<(), SubscriptionValidationError> {
|
||||
self.working_directory_id.validate()?;
|
||||
validate_repository_key(&self.repository_key)?;
|
||||
validate_identifier("workdir_state", &self.state, MAX_RESOURCE_ID_BYTES)?;
|
||||
if let Some(worker_id) = &self.primary_worker_id {
|
||||
worker_id.validate()?;
|
||||
@@ -625,7 +694,7 @@ pub enum SubscriptionSnapshot {
|
||||
events: Vec<WorkerProtocolEvent>,
|
||||
},
|
||||
WorkspaceWorkdirs {
|
||||
workdirs: Vec<SubscriptionWorkdir>,
|
||||
workdirs: Vec<WorkspaceSubscriptionWorkdir>,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -693,7 +762,7 @@ pub enum SubscriptionEventPayload {
|
||||
event: WorkerProtocolEvent,
|
||||
},
|
||||
WorkdirUpserted {
|
||||
workdir: SubscriptionWorkdir,
|
||||
workdir: WorkspaceSubscriptionWorkdir,
|
||||
},
|
||||
WorkdirRemoved {
|
||||
working_directory_id: SubscriptionWorkdirId,
|
||||
@@ -811,10 +880,42 @@ mod tests {
|
||||
display_name: Some(format!("Worker {value}")),
|
||||
profile: Some("builtin:coder".to_string()),
|
||||
repository_id: None,
|
||||
repository_key: None,
|
||||
working_directory_id: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_and_workspace_repository_identity_projections_do_not_alias() {
|
||||
let mut runtime_worker = worker("worker-1");
|
||||
runtime_worker.repository_id = Some("01890f47-3c22-7cc0-98c4-dc0c0c07398f".to_string());
|
||||
runtime_worker.validate().unwrap();
|
||||
let runtime_json = serde_json::to_value(&runtime_worker).unwrap();
|
||||
assert_eq!(
|
||||
runtime_json["repository_id"],
|
||||
"01890f47-3c22-7cc0-98c4-dc0c0c07398f"
|
||||
);
|
||||
assert!(runtime_json.get("repository_key").is_none());
|
||||
|
||||
let mut workspace_worker = worker("worker-1");
|
||||
workspace_worker.repository_key = Some("main".to_string());
|
||||
workspace_worker.validate().unwrap();
|
||||
let workspace_json = serde_json::to_value(&workspace_worker).unwrap();
|
||||
assert_eq!(workspace_json["repository_key"], "main");
|
||||
assert!(workspace_json.get("repository_id").is_none());
|
||||
|
||||
let workspace_workdir = WorkspaceSubscriptionWorkdir {
|
||||
working_directory_id: SubscriptionWorkdirId::new("workdir-1").unwrap(),
|
||||
repository_key: "main".to_string(),
|
||||
state: "active".to_string(),
|
||||
primary_worker_id: Some(worker_id("worker-1")),
|
||||
};
|
||||
workspace_workdir.validate().unwrap();
|
||||
let workdir_json = serde_json::to_value(&workspace_workdir).unwrap();
|
||||
assert_eq!(workdir_json["repository_key"], "main");
|
||||
assert!(workdir_json.get("repository_id").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn subscribe_frame_has_stable_versioned_json_shape() {
|
||||
let frame = SubscriptionFrame::new(SubscriptionFramePayload::Request(
|
||||
@@ -1008,6 +1109,25 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn worker_subscription_state_has_exactly_four_lifecycle_values() {
|
||||
for (state, wire) in [
|
||||
(SubscriptionWorkerState::Idle, "idle"),
|
||||
(SubscriptionWorkerState::Running, "running"),
|
||||
(SubscriptionWorkerState::Paused, "paused"),
|
||||
(SubscriptionWorkerState::Stopped, "stopped"),
|
||||
] {
|
||||
assert_eq!(
|
||||
serde_json::to_value(state).unwrap(),
|
||||
serde_json::json!(wire)
|
||||
);
|
||||
}
|
||||
assert!(
|
||||
serde_json::from_value::<SubscriptionWorkerState>(serde_json::json!("cancelled"))
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn client_selector_has_no_workspace_scope_field() {
|
||||
let json = serde_json::to_value(EventSubscriptionSelector::WorkspaceWorkers).unwrap();
|
||||
|
||||
@@ -7,15 +7,19 @@ use crate::{
|
||||
CommandStreamSlice, CompactionLifecycle, CompactionLifecycleState, CompletionEntry,
|
||||
CompletionKind, ErrorCode, Event, Greeting, InFlightBlock, InFlightSnapshot,
|
||||
InFlightToolCallState, InternalWorkerKind, InternalWorkerRef, InternalWorkerSnapshot,
|
||||
InvokeKind, MemoryWorkerEvent, Method, Permission, RewindSummary, RewindTarget, RewindTargetId,
|
||||
RunResult, ScopeRule, Segment, ToolResultDisposition, TurnResult, WorkerEvent, WorkerStatus,
|
||||
InvokeKind, MemoryWorkerEvent, Method, PasteArtifactAvailability, PasteArtifactMediaType,
|
||||
PasteArtifactRef, Permission, RewindSummary, RewindTarget, RewindTargetId, RunResult,
|
||||
ScopeRule, Segment, SessionContentPart, SessionEntryProvenance, SessionMessageRole,
|
||||
SessionSnapshot, SessionSnapshotEntry, SessionSnapshotEntryData, SessionToolAttachment,
|
||||
ToolResultDisposition, TurnResult, UploadedFileAvailability, UploadedFileRef, WorkerEvent,
|
||||
WorkerStatus,
|
||||
subscription::{
|
||||
EventSubscriptionSelector, SubscriptionEvent, SubscriptionEventPayload, SubscriptionFrame,
|
||||
SubscriptionFramePayload, SubscriptionId, SubscriptionRejectionCode, SubscriptionRequest,
|
||||
SubscriptionRequestId, SubscriptionResponse, SubscriptionSnapshot,
|
||||
SubscriptionTerminationCode, SubscriptionWorkdir, SubscriptionWorkdirId,
|
||||
SubscriptionWorker, SubscriptionWorkerId, SubscriptionWorkerIds,
|
||||
SubscriptionWorkerProtocolMethod, SubscriptionWorkerState,
|
||||
SubscriptionTerminationCode, SubscriptionWorkdirId, SubscriptionWorker,
|
||||
SubscriptionWorkerId, SubscriptionWorkerIds, SubscriptionWorkerProtocolMethod,
|
||||
SubscriptionWorkerState, WorkspaceSubscriptionWorkdir,
|
||||
},
|
||||
};
|
||||
|
||||
@@ -56,6 +60,8 @@ pub fn generated_protocol_types() -> String {
|
||||
push_decl::<CommandEvent>(&cfg, &mut output);
|
||||
push_decl::<CompactionLifecycleState>(&cfg, &mut output);
|
||||
push_decl::<CompactionLifecycle>(&cfg, &mut output);
|
||||
push_decl::<UploadedFileAvailability>(&cfg, &mut output);
|
||||
push_decl::<UploadedFileRef>(&cfg, &mut output);
|
||||
push_decl::<ScopeRule>(&cfg, &mut output);
|
||||
push_decl::<CompletionEntry>(&cfg, &mut output);
|
||||
push_decl::<RewindTargetId>(&cfg, &mut output);
|
||||
@@ -63,12 +69,22 @@ pub fn generated_protocol_types() -> String {
|
||||
push_decl::<RewindSummary>(&cfg, &mut output);
|
||||
push_decl::<InFlightBlock>(&cfg, &mut output);
|
||||
push_decl::<InFlightSnapshot>(&cfg, &mut output);
|
||||
push_decl::<SessionEntryProvenance>(&cfg, &mut output);
|
||||
push_decl::<SessionMessageRole>(&cfg, &mut output);
|
||||
push_decl::<SessionContentPart>(&cfg, &mut output);
|
||||
push_decl::<SessionToolAttachment>(&cfg, &mut output);
|
||||
push_decl::<SessionSnapshotEntryData>(&cfg, &mut output);
|
||||
push_decl::<SessionSnapshotEntry>(&cfg, &mut output);
|
||||
push_decl::<SessionSnapshot>(&cfg, &mut output);
|
||||
push_decl::<InternalWorkerKind>(&cfg, &mut output);
|
||||
push_decl::<InternalWorkerRef>(&cfg, &mut output);
|
||||
push_decl::<InternalWorkerSnapshot>(&cfg, &mut output);
|
||||
push_decl::<Greeting>(&cfg, &mut output);
|
||||
push_decl::<Alert>(&cfg, &mut output);
|
||||
push_decl::<MemoryWorkerEvent>(&cfg, &mut output);
|
||||
push_decl::<PasteArtifactMediaType>(&cfg, &mut output);
|
||||
push_decl::<PasteArtifactAvailability>(&cfg, &mut output);
|
||||
push_decl::<PasteArtifactRef>(&cfg, &mut output);
|
||||
push_decl::<Segment>(&cfg, &mut output);
|
||||
push_decl::<WorkerEvent>(&cfg, &mut output);
|
||||
push_decl::<SubscriptionRequestId>(&cfg, &mut output);
|
||||
@@ -79,7 +95,7 @@ pub fn generated_protocol_types() -> String {
|
||||
push_decl::<SubscriptionWorkerState>(&cfg, &mut output);
|
||||
push_decl::<EventSubscriptionSelector>(&cfg, &mut output);
|
||||
push_decl::<SubscriptionWorker>(&cfg, &mut output);
|
||||
push_decl::<SubscriptionWorkdir>(&cfg, &mut output);
|
||||
push_decl::<WorkspaceSubscriptionWorkdir>(&cfg, &mut output);
|
||||
push_decl::<SubscriptionSnapshot>(&cfg, &mut output);
|
||||
push_decl::<SubscriptionEventPayload>(&cfg, &mut output);
|
||||
push_decl::<SubscriptionRejectionCode>(&cfg, &mut output);
|
||||
@@ -123,6 +139,14 @@ fn export_decl(decl: &str) -> String {
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn workspace_typescript_omits_runtime_repository_ids() {
|
||||
let generated = generated_protocol_types();
|
||||
assert!(!generated.contains("repository_id?:"), "{generated}");
|
||||
assert!(!generated.contains("repository_id:"), "{generated}");
|
||||
assert!(generated.contains("repository_key"), "{generated}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generated_protocol_types_are_current() {
|
||||
let expected = generated_protocol_types();
|
||||
|
||||
@@ -8,12 +8,17 @@ license.workspace = true
|
||||
[dependencies]
|
||||
base64.workspace = true
|
||||
agen = { workspace = true }
|
||||
fs4.workspace = true
|
||||
serde = { workspace = true, features = ["derive"] }
|
||||
serde_json = { workspace = true }
|
||||
sha2.workspace = true
|
||||
uuid = { workspace = true, features = ["v7", "serde"] }
|
||||
thiserror = { workspace = true }
|
||||
protocol = { workspace = true }
|
||||
tracing.workspace = true
|
||||
unicode-normalization = "0.1.25"
|
||||
unicode-properties = { version = "0.1.4", features = ["general-category"] }
|
||||
unicode-security = "0.1.2"
|
||||
|
||||
[dev-dependencies]
|
||||
async-trait = { workspace = true }
|
||||
|
||||
@@ -16,9 +16,18 @@
|
||||
//! enumerable by the picker.
|
||||
|
||||
use crate::event_trace::TraceEntry;
|
||||
use crate::paste_artifact::{read_from_dir, write_to_dir};
|
||||
use crate::segment_log::LogEntry;
|
||||
use crate::store::{Store, StoreError};
|
||||
use crate::{SegmentId, SessionId};
|
||||
use crate::uploaded_file::{
|
||||
bind_uploaded_file, clear_uploaded_file_binding, copy_committed_uploaded_files,
|
||||
delete_uncommitted_uploaded_files, delete_uploaded_file, list_uploaded_file_refs,
|
||||
read_uploaded_file, read_uploaded_file_by_id, write_uploaded_file,
|
||||
};
|
||||
use crate::{
|
||||
PasteArtifactLimits, SegmentId, SessionId, UploadedFileLimits, UploadedFileUploadContext,
|
||||
};
|
||||
use protocol::{PasteArtifactRef, UploadedFileRef};
|
||||
use std::fs;
|
||||
use std::io::{Read, Seek, SeekFrom, Write};
|
||||
use std::path::{Path, PathBuf};
|
||||
@@ -109,6 +118,50 @@ impl FsStore {
|
||||
.join(format!("{segment_id}.trace.jsonl"))
|
||||
}
|
||||
|
||||
fn paste_artifact_dir(&self, session_id: SessionId) -> PathBuf {
|
||||
self.session_dir(session_id).join("artifacts").join("paste")
|
||||
}
|
||||
|
||||
fn uploaded_file_is_referenced(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
artifact_id: &str,
|
||||
) -> Result<bool, StoreError> {
|
||||
fn segments_contain(segments: &[protocol::Segment], artifact_id: &str) -> bool {
|
||||
segments.iter().any(|segment| {
|
||||
matches!(
|
||||
segment,
|
||||
protocol::Segment::UploadedFile { file }
|
||||
if file.artifact_id == artifact_id
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
for segment_id in self.list_segments(session_id)? {
|
||||
for entry in self.read_all(session_id, segment_id)? {
|
||||
let referenced = match entry {
|
||||
LogEntry::AnnotatedUserInput { segments, .. } => {
|
||||
segments_contain(&segments, artifact_id)
|
||||
}
|
||||
LogEntry::InputSegmentsCheckpoint { user_segments, .. } => user_segments
|
||||
.iter()
|
||||
.any(|segments| segments_contain(segments, artifact_id)),
|
||||
_ => false,
|
||||
};
|
||||
if referenced {
|
||||
return Ok(true);
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(false)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
fn paste_artifact_path(&self, session_id: SessionId, artifact_id: &str) -> PathBuf {
|
||||
self.paste_artifact_dir(session_id)
|
||||
.join(format!("{artifact_id}.json"))
|
||||
}
|
||||
|
||||
fn append_line(&self, path: &Path, line: &str) -> Result<(), StoreError> {
|
||||
let _guard = self
|
||||
.append_lock
|
||||
@@ -350,6 +403,171 @@ impl Store for FsStore {
|
||||
Ok(complete.lines().filter(|l| !l.trim().is_empty()).count())
|
||||
}
|
||||
|
||||
fn write_paste_artifact(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
source_entry_id: &str,
|
||||
content: &str,
|
||||
limits: PasteArtifactLimits,
|
||||
) -> Result<PasteArtifactRef, StoreError> {
|
||||
let _guard = self
|
||||
.append_lock
|
||||
.lock()
|
||||
.map_err(|_| std::io::Error::other("session store append lock was poisoned"))?;
|
||||
write_to_dir(
|
||||
&self.paste_artifact_dir(session_id),
|
||||
source_entry_id,
|
||||
content,
|
||||
limits,
|
||||
)
|
||||
}
|
||||
|
||||
fn read_paste_artifact(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
artifact_id: &str,
|
||||
) -> Result<(PasteArtifactRef, String), StoreError> {
|
||||
read_from_dir(&self.paste_artifact_dir(session_id), artifact_id)
|
||||
}
|
||||
|
||||
fn write_uploaded_file(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
file_name: &str,
|
||||
media_type: &str,
|
||||
content: &[u8],
|
||||
limits: UploadedFileLimits,
|
||||
) -> Result<UploadedFileRef, StoreError> {
|
||||
let _guard = self
|
||||
.append_lock
|
||||
.lock()
|
||||
.map_err(|_| std::io::Error::other("session store append lock was poisoned"))?;
|
||||
write_uploaded_file(
|
||||
&self.paste_artifact_dir(session_id),
|
||||
file_name,
|
||||
media_type,
|
||||
content,
|
||||
None,
|
||||
limits,
|
||||
)
|
||||
}
|
||||
|
||||
fn write_uploaded_file_with_context(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
file_name: &str,
|
||||
media_type: &str,
|
||||
content: &[u8],
|
||||
context: &UploadedFileUploadContext,
|
||||
limits: UploadedFileLimits,
|
||||
) -> Result<UploadedFileRef, StoreError> {
|
||||
let _guard = self
|
||||
.append_lock
|
||||
.lock()
|
||||
.map_err(|_| std::io::Error::other("session store append lock was poisoned"))?;
|
||||
write_uploaded_file(
|
||||
&self.paste_artifact_dir(session_id),
|
||||
file_name,
|
||||
media_type,
|
||||
content,
|
||||
Some(context),
|
||||
limits,
|
||||
)
|
||||
}
|
||||
|
||||
fn read_uploaded_file(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
reference: &UploadedFileRef,
|
||||
) -> Result<Vec<u8>, StoreError> {
|
||||
read_uploaded_file(&self.paste_artifact_dir(session_id), reference)
|
||||
}
|
||||
|
||||
fn read_uploaded_file_by_id(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
artifact_id: &str,
|
||||
) -> Result<(UploadedFileRef, Vec<u8>), StoreError> {
|
||||
read_uploaded_file_by_id(&self.paste_artifact_dir(session_id), artifact_id)
|
||||
}
|
||||
|
||||
fn bind_uploaded_file(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
reference: &UploadedFileRef,
|
||||
source_entry_id: &str,
|
||||
) -> Result<UploadedFileRef, StoreError> {
|
||||
let _guard = self
|
||||
.append_lock
|
||||
.lock()
|
||||
.map_err(|_| std::io::Error::other("session store append lock was poisoned"))?;
|
||||
let dir = self.paste_artifact_dir(session_id);
|
||||
match bind_uploaded_file(&dir, reference, source_entry_id) {
|
||||
Err(StoreError::ArtifactAlreadyCommitted) => {
|
||||
let (stored, _) = read_uploaded_file_by_id(&dir, &reference.artifact_id)?;
|
||||
let previous_source = stored
|
||||
.source_entry_id
|
||||
.ok_or(StoreError::ArtifactIntegrityMismatch)?;
|
||||
if self.uploaded_file_is_referenced(session_id, &reference.artifact_id)? {
|
||||
return Err(StoreError::ArtifactAlreadyCommitted);
|
||||
}
|
||||
clear_uploaded_file_binding(&dir, &reference.artifact_id, &previous_source)?;
|
||||
bind_uploaded_file(&dir, reference, source_entry_id)
|
||||
}
|
||||
result => result,
|
||||
}
|
||||
}
|
||||
|
||||
fn delete_uploaded_file(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
artifact_id: &str,
|
||||
) -> Result<bool, StoreError> {
|
||||
let _guard = self
|
||||
.append_lock
|
||||
.lock()
|
||||
.map_err(|_| std::io::Error::other("session store append lock was poisoned"))?;
|
||||
delete_uploaded_file(&self.paste_artifact_dir(session_id), artifact_id)
|
||||
}
|
||||
|
||||
fn delete_uncommitted_uploaded_files(&self, session_id: SessionId) -> Result<u64, StoreError> {
|
||||
let _guard = self
|
||||
.append_lock
|
||||
.lock()
|
||||
.map_err(|_| std::io::Error::other("session store append lock was poisoned"))?;
|
||||
let dir = self.paste_artifact_dir(session_id);
|
||||
let mut removed = delete_uncommitted_uploaded_files(&dir)?;
|
||||
for reference in list_uploaded_file_refs(&dir)? {
|
||||
let Some(source_entry_id) = reference.source_entry_id.as_deref() else {
|
||||
continue;
|
||||
};
|
||||
if !self.uploaded_file_is_referenced(session_id, &reference.artifact_id)? {
|
||||
clear_uploaded_file_binding(&dir, &reference.artifact_id, source_entry_id)?;
|
||||
if delete_uploaded_file(&dir, &reference.artifact_id)? {
|
||||
removed = removed
|
||||
.checked_add(1)
|
||||
.ok_or(StoreError::ArtifactQuotaExceeded)?;
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(removed)
|
||||
}
|
||||
|
||||
fn copy_committed_uploaded_files(
|
||||
&self,
|
||||
source_session_id: SessionId,
|
||||
target_session_id: SessionId,
|
||||
) -> Result<u64, StoreError> {
|
||||
let _guard = self
|
||||
.append_lock
|
||||
.lock()
|
||||
.map_err(|_| std::io::Error::other("session store append lock was poisoned"))?;
|
||||
copy_committed_uploaded_files(
|
||||
&self.paste_artifact_dir(source_session_id),
|
||||
&self.paste_artifact_dir(target_session_id),
|
||||
)
|
||||
}
|
||||
|
||||
fn append_trace(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
@@ -398,4 +616,424 @@ mod tests {
|
||||
store.create_segment(session_id, segment_id, &[]).unwrap();
|
||||
assert!(store.session_modified_at(session_id).unwrap().is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn paste_artifacts_are_atomic_integrity_checked_and_session_scoped() {
|
||||
let tmp = tempfile::TempDir::new().unwrap();
|
||||
let store = FsStore::new(tmp.path()).unwrap();
|
||||
let owner = new_session_id();
|
||||
let other = new_session_id();
|
||||
let content = "αβγ\nsecond line\n";
|
||||
let reference = store
|
||||
.write_paste_artifact(owner, "entry-1", content, PasteArtifactLimits::default())
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(reference.byte_len, content.len() as u64);
|
||||
assert!(reference.created_at_ms > 0);
|
||||
assert_eq!(
|
||||
reference.media_type,
|
||||
protocol::PasteArtifactMediaType::TextPlainUtf8
|
||||
);
|
||||
assert_eq!(
|
||||
reference.availability,
|
||||
protocol::PasteArtifactAvailability::Available
|
||||
);
|
||||
assert_eq!(reference.char_count, content.chars().count() as u64);
|
||||
assert_eq!(reference.source_entry_id, "entry-1");
|
||||
assert_eq!(
|
||||
store
|
||||
.read_paste_artifact(owner, &reference.artifact_id)
|
||||
.unwrap()
|
||||
.1,
|
||||
content
|
||||
);
|
||||
assert!(matches!(
|
||||
store.read_paste_artifact(other, &reference.artifact_id),
|
||||
Err(StoreError::PasteArtifactNotFound(_))
|
||||
));
|
||||
assert!(
|
||||
self::fs::read_dir(store.paste_artifact_dir(owner))
|
||||
.unwrap()
|
||||
.all(|entry| !entry
|
||||
.unwrap()
|
||||
.file_name()
|
||||
.to_string_lossy()
|
||||
.ends_with(".tmp"))
|
||||
);
|
||||
let very_large = "z".repeat(1024 * 1024);
|
||||
let very_large_ref = store
|
||||
.write_paste_artifact(
|
||||
owner,
|
||||
"entry-2",
|
||||
&very_large,
|
||||
PasteArtifactLimits::default(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
store
|
||||
.read_paste_artifact(owner, &very_large_ref.artifact_id)
|
||||
.unwrap()
|
||||
.1,
|
||||
very_large
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn concurrent_paste_writes_atomically_enforce_aggregate_caps() {
|
||||
let tmp = tempfile::TempDir::new().unwrap();
|
||||
let session_id = new_session_id();
|
||||
let barrier = std::sync::Arc::new(std::sync::Barrier::new(3));
|
||||
let limits = PasteArtifactLimits {
|
||||
max_artifact_bytes: 4,
|
||||
max_session_bytes: 8,
|
||||
max_session_artifacts: 1,
|
||||
};
|
||||
let mut handles = Vec::new();
|
||||
for entry_id in ["entry-1", "entry-2"] {
|
||||
let root = tmp.path().to_path_buf();
|
||||
let barrier = barrier.clone();
|
||||
handles.push(std::thread::spawn(move || {
|
||||
let store = FsStore::new(root).unwrap();
|
||||
barrier.wait();
|
||||
store.write_paste_artifact(session_id, entry_id, "1234", limits)
|
||||
}));
|
||||
}
|
||||
barrier.wait();
|
||||
let results = handles
|
||||
.into_iter()
|
||||
.map(|handle| handle.join().unwrap())
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(results.iter().filter(|result| result.is_ok()).count(), 1);
|
||||
assert_eq!(
|
||||
results
|
||||
.iter()
|
||||
.filter(|result| matches!(result, Err(StoreError::PasteArtifactLimit(_))))
|
||||
.count(),
|
||||
1
|
||||
);
|
||||
assert_eq!(
|
||||
std::fs::read_dir(
|
||||
FsStore::new(tmp.path())
|
||||
.unwrap()
|
||||
.paste_artifact_dir(session_id)
|
||||
)
|
||||
.unwrap()
|
||||
.filter_map(Result::ok)
|
||||
.filter(
|
||||
|entry| entry.path().extension().and_then(|value| value.to_str()) == Some("json")
|
||||
)
|
||||
.count(),
|
||||
1
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn uploaded_file_persists_trusted_upload_context_without_projecting_it() {
|
||||
let tmp = tempfile::TempDir::new().unwrap();
|
||||
let store = FsStore::new(tmp.path()).unwrap();
|
||||
let session_id = new_session_id();
|
||||
let context = UploadedFileUploadContext {
|
||||
upload_id: "upload-1".into(),
|
||||
principal_id: "account-1".into(),
|
||||
workspace_id: "workspace-1".into(),
|
||||
runtime_id: "runtime-1".into(),
|
||||
worker_id: "worker-1".into(),
|
||||
};
|
||||
let reference = store
|
||||
.write_uploaded_file_with_context(
|
||||
session_id,
|
||||
"notes.txt",
|
||||
"text/plain",
|
||||
b"hello",
|
||||
&context,
|
||||
UploadedFileLimits::default(),
|
||||
)
|
||||
.unwrap();
|
||||
let raw = fs::read_to_string(
|
||||
store
|
||||
.paste_artifact_dir(session_id)
|
||||
.join(format!("{}.file.json", reference.artifact_id)),
|
||||
)
|
||||
.unwrap();
|
||||
assert!(raw.contains("account-1"));
|
||||
assert!(raw.contains("workspace-1"));
|
||||
assert!(raw.contains("runtime-1"));
|
||||
assert!(raw.contains("worker-1"));
|
||||
assert!(
|
||||
!serde_json::to_string(&reference)
|
||||
.unwrap()
|
||||
.contains("account-1")
|
||||
);
|
||||
|
||||
let replay = store
|
||||
.write_uploaded_file_with_context(
|
||||
session_id,
|
||||
"notes.txt",
|
||||
"text/plain",
|
||||
b"hello",
|
||||
&context,
|
||||
UploadedFileLimits::default(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(replay.artifact_id, reference.artifact_id);
|
||||
assert!(matches!(
|
||||
store.write_uploaded_file_with_context(
|
||||
session_id,
|
||||
"renamed.txt",
|
||||
"text/plain",
|
||||
b"hello",
|
||||
&context,
|
||||
UploadedFileLimits::default(),
|
||||
),
|
||||
Err(StoreError::InvalidUploadedFileName)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn uploaded_file_exact_replay_succeeds_at_session_count_limit() {
|
||||
let tmp = tempfile::TempDir::new().unwrap();
|
||||
let store = FsStore::new(tmp.path()).unwrap();
|
||||
let session_id = new_session_id();
|
||||
let limits = UploadedFileLimits {
|
||||
max_file_bytes: 1,
|
||||
max_session_bytes: crate::DEFAULT_MAX_SESSION_UPLOADED_FILES,
|
||||
};
|
||||
let mut first = None;
|
||||
for index in 0..crate::DEFAULT_MAX_SESSION_UPLOADED_FILES {
|
||||
let reference = store
|
||||
.write_uploaded_file(
|
||||
session_id,
|
||||
&format!("file-{index}.txt"),
|
||||
"text/plain",
|
||||
b"x",
|
||||
limits,
|
||||
)
|
||||
.unwrap();
|
||||
first.get_or_insert(reference);
|
||||
}
|
||||
|
||||
let replay = store
|
||||
.write_uploaded_file(session_id, "file-0.txt", "text/plain", b"x", limits)
|
||||
.unwrap();
|
||||
assert_eq!(replay.artifact_id, first.unwrap().artifact_id);
|
||||
assert!(matches!(
|
||||
store.write_uploaded_file(session_id, "overflow.txt", "text/plain", b"x", limits),
|
||||
Err(StoreError::ArtifactQuotaExceeded)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn uploaded_files_are_session_scoped_integrity_checked_and_removable() {
|
||||
let tmp = tempfile::TempDir::new().unwrap();
|
||||
let store = FsStore::new(tmp.path()).unwrap();
|
||||
let owner = new_session_id();
|
||||
let other = new_session_id();
|
||||
let limits = UploadedFileLimits {
|
||||
max_file_bytes: 16,
|
||||
max_session_bytes: 16,
|
||||
};
|
||||
let reference = store
|
||||
.write_uploaded_file(owner, "notes.txt", "text/plain", b"hello", limits)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(reference.file_name, "notes.txt");
|
||||
assert_eq!(reference.media_type, "text/plain");
|
||||
assert_eq!(reference.byte_len, 5);
|
||||
assert_eq!(reference.source_entry_id, None);
|
||||
assert_eq!(
|
||||
store.read_uploaded_file(owner, &reference).unwrap(),
|
||||
b"hello"
|
||||
);
|
||||
assert!(store.read_uploaded_file(other, &reference).is_err());
|
||||
|
||||
let mut forged = reference.clone();
|
||||
forged.file_name = "other.txt".to_string();
|
||||
assert!(matches!(
|
||||
store.read_uploaded_file(owner, &forged),
|
||||
Err(StoreError::ArtifactIntegrityMismatch)
|
||||
));
|
||||
assert!(
|
||||
store
|
||||
.delete_uploaded_file(owner, &reference.artifact_id)
|
||||
.unwrap()
|
||||
);
|
||||
assert!(
|
||||
!store
|
||||
.delete_uploaded_file(owner, &reference.artifact_id)
|
||||
.unwrap()
|
||||
);
|
||||
assert!(store.read_uploaded_file(owner, &reference).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn uploaded_file_validation_and_shared_quota_fail_closed() {
|
||||
let tmp = tempfile::TempDir::new().unwrap();
|
||||
let store = FsStore::new(tmp.path()).unwrap();
|
||||
let session_id = new_session_id();
|
||||
let limits = UploadedFileLimits {
|
||||
max_file_bytes: 8,
|
||||
max_session_bytes: 8,
|
||||
};
|
||||
assert!(matches!(
|
||||
store.write_uploaded_file(session_id, "../secret", "text/plain", b"x", limits),
|
||||
Err(StoreError::InvalidUploadedFileName)
|
||||
));
|
||||
assert!(matches!(
|
||||
store.write_uploaded_file(session_id, "notes.txt", "not a type", b"x", limits),
|
||||
Err(StoreError::InvalidUploadedFileMediaType)
|
||||
));
|
||||
assert!(matches!(
|
||||
store.write_uploaded_file(
|
||||
session_id,
|
||||
"safe\u{202e}txt.exe",
|
||||
"text/plain",
|
||||
b"x",
|
||||
limits
|
||||
),
|
||||
Err(StoreError::InvalidUploadedFileName)
|
||||
));
|
||||
assert!(matches!(
|
||||
store.write_uploaded_file(session_id, "image.png", "image/png", b"not a png", limits),
|
||||
Err(StoreError::ArtifactIntegrityMismatch)
|
||||
));
|
||||
let pending = store
|
||||
.write_uploaded_file(session_id, "Readme.txt", "text/plain", b"x", limits)
|
||||
.unwrap();
|
||||
let replay = store
|
||||
.write_uploaded_file(session_id, "Readme.txt", "text/plain", b"x", limits)
|
||||
.unwrap();
|
||||
assert_eq!(replay.artifact_id, pending.artifact_id);
|
||||
assert!(matches!(
|
||||
store.write_uploaded_file(session_id, "README.txt", "text/plain", b"changed", limits),
|
||||
Err(StoreError::InvalidUploadedFileName)
|
||||
));
|
||||
assert!(matches!(
|
||||
store.write_uploaded_file(session_id, "README.txt", "text/plain", b"y", limits),
|
||||
Err(StoreError::InvalidUploadedFileName)
|
||||
));
|
||||
store
|
||||
.bind_uploaded_file(session_id, &pending, "entry-from-failed-submit")
|
||||
.unwrap();
|
||||
let bound = store
|
||||
.bind_uploaded_file(session_id, &pending, "entry-upload")
|
||||
.unwrap();
|
||||
store
|
||||
.create_segment(
|
||||
session_id,
|
||||
new_segment_id(),
|
||||
&[LogEntry::InputSegmentsCheckpoint {
|
||||
ts: 1,
|
||||
user_segments: vec![vec![protocol::Segment::UploadedFile {
|
||||
file: bound.clone(),
|
||||
}]],
|
||||
}],
|
||||
)
|
||||
.unwrap();
|
||||
let other = store
|
||||
.write_uploaded_file(session_id, "other.txt", "text/plain", b"z", limits)
|
||||
.unwrap();
|
||||
let stale = store
|
||||
.write_uploaded_file(session_id, "stale.txt", "text/plain", b"s", limits)
|
||||
.unwrap();
|
||||
store
|
||||
.bind_uploaded_file(session_id, &stale, "entry-never-committed")
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
store.delete_uncommitted_uploaded_files(session_id).unwrap(),
|
||||
2
|
||||
);
|
||||
assert!(store.read_uploaded_file(session_id, &other).is_err());
|
||||
assert!(store.read_uploaded_file(session_id, &stale).is_err());
|
||||
assert_eq!(store.read_uploaded_file(session_id, &bound).unwrap(), b"x");
|
||||
let fork_session_id = new_session_id();
|
||||
assert_eq!(
|
||||
store
|
||||
.copy_committed_uploaded_files(session_id, fork_session_id)
|
||||
.unwrap(),
|
||||
1
|
||||
);
|
||||
assert_eq!(
|
||||
store.read_uploaded_file(fork_session_id, &bound).unwrap(),
|
||||
b"x"
|
||||
);
|
||||
store
|
||||
.write_paste_artifact(
|
||||
session_id,
|
||||
"entry-1",
|
||||
"1234",
|
||||
PasteArtifactLimits {
|
||||
max_artifact_bytes: 8,
|
||||
max_session_bytes: 8,
|
||||
max_session_artifacts: 4,
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
assert!(matches!(
|
||||
store.write_uploaded_file(session_id, "notes.txt", "text/plain", b"56789", limits),
|
||||
Err(StoreError::ArtifactQuotaExceeded)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn uploaded_file_names_reject_format_mixed_script_and_confusable_forms() {
|
||||
let tmp = tempfile::TempDir::new().unwrap();
|
||||
let store = FsStore::new(tmp.path()).unwrap();
|
||||
let session_id = new_session_id();
|
||||
let limits = UploadedFileLimits::default();
|
||||
|
||||
for file_name in [
|
||||
"safe\u{00ad}name.txt",
|
||||
"safe\u{061c}name.txt",
|
||||
"safe\u{180e}name.txt",
|
||||
"safe\u{e0001}name.txt",
|
||||
"p\u{0430}ypal.txt",
|
||||
"report.\u{03c1}df",
|
||||
"\u{0440}\u{0430}\u{0443}\u{0440}\u{0430}\u{04cf}.txt",
|
||||
"\u{ff26}\u{ff49}\u{ff4c}\u{ff45}.txt",
|
||||
"re\u{0301}sume\u{0301}.txt",
|
||||
] {
|
||||
assert!(matches!(
|
||||
store.write_uploaded_file(session_id, file_name, "text/plain", b"safe", limits),
|
||||
Err(StoreError::InvalidUploadedFileName)
|
||||
));
|
||||
}
|
||||
|
||||
for file_name in ["notes.txt", "résumé.txt", "日本語.txt", "📎.txt"] {
|
||||
store
|
||||
.write_uploaded_file(session_id, file_name, "text/plain", b"safe", limits)
|
||||
.unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn paste_artifact_limits_and_corruption_fail_closed() {
|
||||
let tmp = tempfile::TempDir::new().unwrap();
|
||||
let store = FsStore::new(tmp.path()).unwrap();
|
||||
let session_id = new_session_id();
|
||||
let limits = PasteArtifactLimits {
|
||||
max_artifact_bytes: 5,
|
||||
max_session_bytes: 8,
|
||||
max_session_artifacts: 2,
|
||||
};
|
||||
let first = store
|
||||
.write_paste_artifact(session_id, "entry-1", "1234", limits)
|
||||
.unwrap();
|
||||
assert!(matches!(
|
||||
store.write_paste_artifact(session_id, "entry-2", "56789", limits),
|
||||
Err(StoreError::PasteArtifactLimit(_))
|
||||
));
|
||||
assert!(matches!(
|
||||
store.write_paste_artifact(session_id, "entry-2", "5678", limits),
|
||||
Ok(_)
|
||||
));
|
||||
std::fs::write(
|
||||
store.paste_artifact_path(session_id, &first.artifact_id),
|
||||
b"{}",
|
||||
)
|
||||
.unwrap();
|
||||
assert!(matches!(
|
||||
store.read_paste_artifact(session_id, &first.artifact_id),
|
||||
Err(StoreError::Serde(_)) | Err(StoreError::PasteArtifactIntegrity(_))
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::{LoggedItem, SessionId};
|
||||
use crate::LoggedItem;
|
||||
|
||||
/// Stable logical identity of one model-visible history entry.
|
||||
///
|
||||
@@ -142,12 +142,15 @@ mod tests {
|
||||
#[test]
|
||||
fn annotated_segment_start_is_restore_visible_without_projecting_metadata() {
|
||||
let session_id = uuid::Uuid::now_v7();
|
||||
let history_entry = legacy_logged_history(LoggedItem::Message {
|
||||
role: LoggedRole::Assistant,
|
||||
content: vec![crate::LoggedContentPart::Text {
|
||||
text: "answer".into(),
|
||||
}],
|
||||
});
|
||||
let history_entry = LoggedHistoryEntry {
|
||||
item: LoggedItem::Message {
|
||||
role: LoggedRole::Assistant,
|
||||
content: vec![crate::LoggedContentPart::Text {
|
||||
text: "answer".into(),
|
||||
}],
|
||||
},
|
||||
metadata: LoggedSessionHistoryMetadata::legacy_unknown(),
|
||||
};
|
||||
let state = crate::collect_state(&[crate::LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1,
|
||||
session_id,
|
||||
@@ -160,21 +163,3 @@ mod tests {
|
||||
assert_eq!(state.history[0].as_text(), Some("answer"));
|
||||
}
|
||||
}
|
||||
|
||||
/// Legacy Session Logs did not persist annotations. Decode helpers explicitly
|
||||
/// create `LegacyUnknown`; they never infer Human/System authority from role or
|
||||
/// plaintext.
|
||||
pub fn legacy_logged_history(item: LoggedItem) -> LoggedHistoryEntry {
|
||||
LoggedHistoryEntry {
|
||||
item,
|
||||
metadata: LoggedSessionHistoryMetadata::legacy_unknown(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn legacy_segment_history(
|
||||
session_id: SessionId,
|
||||
items: impl IntoIterator<Item = LoggedItem>,
|
||||
) -> Vec<LoggedHistoryEntry> {
|
||||
let _ = session_id;
|
||||
items.into_iter().map(legacy_logged_history).collect()
|
||||
}
|
||||
|
||||
@@ -0,0 +1,188 @@
|
||||
//! Versioned decoder for Session schemas that predate canonical annotated history.
|
||||
//!
|
||||
//! These types are intentionally private to `session-store`. Current writers,
|
||||
//! replay, and public projections use [`crate::LogEntry`] exclusively; only the
|
||||
//! Worker Session schema migration is allowed to deserialize these shapes.
|
||||
|
||||
use agen::llm_client::types::RequestConfig;
|
||||
use protocol::Segment;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::{
|
||||
LogEntry, LoggedHistoryEntry, LoggedItem, LoggedSessionHistoryEntryId,
|
||||
LoggedSessionHistoryMetadata, LoggedSessionHistoryOrigin, LoggedSystemHistoryEntry, SegmentId,
|
||||
SegmentOrigin, SessionExtension, SessionId, SystemItem,
|
||||
};
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(tag = "kind", rename_all = "snake_case")]
|
||||
enum LegacyHistoryLogEntry {
|
||||
SegmentStart {
|
||||
ts: u64,
|
||||
session_id: SessionId,
|
||||
system_prompt: Option<String>,
|
||||
config: RequestConfig,
|
||||
history: Vec<LoggedItem>,
|
||||
#[serde(default)]
|
||||
forked_from: Option<SegmentOrigin>,
|
||||
#[serde(default)]
|
||||
compacted_from: Option<SegmentOrigin>,
|
||||
},
|
||||
UserInput {
|
||||
ts: u64,
|
||||
segments: Vec<Segment>,
|
||||
#[serde(default)]
|
||||
extensions: Vec<SessionExtension>,
|
||||
},
|
||||
AssistantItem {
|
||||
ts: u64,
|
||||
item: LoggedItem,
|
||||
},
|
||||
ToolResult {
|
||||
ts: u64,
|
||||
item: LoggedItem,
|
||||
},
|
||||
SystemItem {
|
||||
ts: u64,
|
||||
item: SystemItem,
|
||||
},
|
||||
}
|
||||
|
||||
/// Schema-v1 decoder. Non-history records already had their current shape, so
|
||||
/// they pass through `LogEntry`; legacy history records are converted below.
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
enum LegacySessionLogEntryV1 {
|
||||
History(LegacyHistoryLogEntry),
|
||||
Current(LogEntry),
|
||||
}
|
||||
|
||||
/// Schema v2 retained the v1 history shapes while adding non-history records.
|
||||
/// Keep a distinct type so supported source versions remain explicit rather
|
||||
/// than turning migration compatibility into the current `LogEntry` contract.
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
enum LegacySessionLogEntryV2 {
|
||||
History(LegacyHistoryLogEntry),
|
||||
Current(LogEntry),
|
||||
}
|
||||
|
||||
pub(crate) fn decode_entry(
|
||||
schema_version: u32,
|
||||
line: &str,
|
||||
session_id: SessionId,
|
||||
segment_id: SegmentId,
|
||||
line_index: usize,
|
||||
) -> Result<LogEntry, serde_json::Error> {
|
||||
let entry = match schema_version {
|
||||
1 => match serde_json::from_str::<LegacySessionLogEntryV1>(line)? {
|
||||
LegacySessionLogEntryV1::History(entry) => Entry::History(entry),
|
||||
LegacySessionLogEntryV1::Current(entry) => Entry::Current(entry),
|
||||
},
|
||||
2 => match serde_json::from_str::<LegacySessionLogEntryV2>(line)? {
|
||||
LegacySessionLogEntryV2::History(entry) => Entry::History(entry),
|
||||
LegacySessionLogEntryV2::Current(entry) => Entry::Current(entry),
|
||||
},
|
||||
_ => unreachable!("legacy decoder called for unsupported schema {schema_version}"),
|
||||
};
|
||||
Ok(match entry {
|
||||
Entry::History(entry) => {
|
||||
canonicalize_history_entry(session_id, segment_id, line_index, entry)
|
||||
}
|
||||
Entry::Current(entry) => entry,
|
||||
})
|
||||
}
|
||||
|
||||
enum Entry {
|
||||
History(LegacyHistoryLogEntry),
|
||||
Current(LogEntry),
|
||||
}
|
||||
|
||||
fn legacy_metadata(
|
||||
segment_id: SegmentId,
|
||||
line_index: usize,
|
||||
item_index: usize,
|
||||
) -> LoggedSessionHistoryMetadata {
|
||||
let mut identity = Vec::with_capacity(32);
|
||||
identity.extend_from_slice(segment_id.as_bytes());
|
||||
identity.extend_from_slice(&(line_index as u64).to_be_bytes());
|
||||
identity.extend_from_slice(&(item_index as u64).to_be_bytes());
|
||||
LoggedSessionHistoryMetadata {
|
||||
entry_id: LoggedSessionHistoryEntryId(format!(
|
||||
"l-{}",
|
||||
base64::Engine::encode(&base64::engine::general_purpose::URL_SAFE_NO_PAD, identity)
|
||||
)),
|
||||
origin: LoggedSessionHistoryOrigin::LegacyUnknown,
|
||||
derivation: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn canonicalize_history_entry(
|
||||
_session_id: SessionId,
|
||||
segment_id: SegmentId,
|
||||
line_index: usize,
|
||||
entry: LegacyHistoryLogEntry,
|
||||
) -> LogEntry {
|
||||
match entry {
|
||||
LegacyHistoryLogEntry::SegmentStart {
|
||||
ts,
|
||||
session_id,
|
||||
system_prompt,
|
||||
config,
|
||||
history,
|
||||
forked_from,
|
||||
compacted_from,
|
||||
} => LogEntry::AnnotatedSegmentStart {
|
||||
ts,
|
||||
session_id,
|
||||
system_prompt,
|
||||
config,
|
||||
history: history
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.map(|(item_index, item)| LoggedHistoryEntry {
|
||||
item,
|
||||
metadata: legacy_metadata(segment_id, line_index, item_index),
|
||||
})
|
||||
.collect(),
|
||||
forked_from,
|
||||
compacted_from,
|
||||
},
|
||||
LegacyHistoryLogEntry::UserInput {
|
||||
ts,
|
||||
segments,
|
||||
extensions,
|
||||
} => LogEntry::AnnotatedUserInput {
|
||||
ts,
|
||||
history: vec![LoggedHistoryEntry {
|
||||
item: LoggedItem::from(agen::Item::user_message(Segment::flatten_to_text(
|
||||
&segments,
|
||||
))),
|
||||
metadata: legacy_metadata(segment_id, line_index, 0),
|
||||
}],
|
||||
segments,
|
||||
extensions,
|
||||
},
|
||||
LegacyHistoryLogEntry::AssistantItem { ts, item } => LogEntry::AnnotatedAssistantItem {
|
||||
ts,
|
||||
entry: LoggedHistoryEntry {
|
||||
item,
|
||||
metadata: legacy_metadata(segment_id, line_index, 0),
|
||||
},
|
||||
},
|
||||
LegacyHistoryLogEntry::ToolResult { ts, item } => LogEntry::AnnotatedToolResult {
|
||||
ts,
|
||||
entry: LoggedHistoryEntry {
|
||||
item,
|
||||
metadata: legacy_metadata(segment_id, line_index, 0),
|
||||
},
|
||||
},
|
||||
LegacyHistoryLogEntry::SystemItem { ts, item } => LogEntry::AnnotatedSystemItem {
|
||||
ts,
|
||||
entry: LoggedSystemHistoryEntry {
|
||||
item,
|
||||
metadata: legacy_metadata(segment_id, line_index, 0),
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -26,18 +26,23 @@
|
||||
//! let (session_id, segment_id) = create_segment(&store, SegmentStartState {
|
||||
//! system_prompt: None,
|
||||
//! config: &config,
|
||||
//! history: &[],
|
||||
//! history: Vec::new(),
|
||||
//! user_segments: Vec::new(),
|
||||
//! })?;
|
||||
//! ```
|
||||
|
||||
pub mod event_trace;
|
||||
pub mod fs_store;
|
||||
pub mod history;
|
||||
mod legacy_session_log;
|
||||
pub mod logged_item;
|
||||
mod paste_artifact;
|
||||
pub mod public_snapshot;
|
||||
pub mod segment;
|
||||
pub mod segment_log;
|
||||
pub mod store;
|
||||
pub mod system_item;
|
||||
pub mod uploaded_file;
|
||||
pub mod worker_metadata;
|
||||
pub mod worker_session_store;
|
||||
|
||||
@@ -48,11 +53,12 @@ pub use fs_store::FsStore;
|
||||
pub use history::{
|
||||
LoggedHistoryDerivation, LoggedHistoryEntry, LoggedSessionHistoryEntryId,
|
||||
LoggedSessionHistoryMetadata, LoggedSessionHistoryOrigin, LoggedSystemHistoryEntry,
|
||||
LoggedWorkerSubject, legacy_logged_history, legacy_segment_history,
|
||||
LoggedWorkerSubject,
|
||||
};
|
||||
pub use logged_item::{LoggedContentPart, LoggedItem, LoggedRole, from_logged, to_logged};
|
||||
pub use paste_artifact::PasteArtifactLimits;
|
||||
pub use segment::{
|
||||
SegmentStartState, append_entry, append_system_item, classify_history_item,
|
||||
SegmentStartState, append_entry, append_system_item, classify_logged_history_entry,
|
||||
create_compacted_segment, create_segment, create_segment_with_ids, ensure_head_or_fork, fork,
|
||||
fork_at, restore, restore_by_segment, save_config_changed, save_delta, save_extension,
|
||||
save_run_completed, save_run_errored, save_turn_end, save_usage, save_user_input,
|
||||
@@ -62,6 +68,11 @@ pub use store::{Store, StoreError};
|
||||
pub use system_item::{
|
||||
PromptRenderProvenance, SystemItem, SystemReminder, SystemReminderSource, render_worker_event,
|
||||
};
|
||||
pub use uploaded_file::{
|
||||
DEFAULT_MAX_FILES_PER_SUBMISSION, DEFAULT_MAX_SESSION_ARTIFACT_BYTES,
|
||||
DEFAULT_MAX_SESSION_UPLOADED_FILES, DEFAULT_MAX_UPLOADED_FILE_BYTES, UploadedFileLimits,
|
||||
UploadedFileUploadContext,
|
||||
};
|
||||
pub use worker_metadata::{
|
||||
CombinedStore, FsWorkerStore, WorkerActiveSegmentRef, WorkerAggregateStore, WorkerMetadata,
|
||||
WorkerMetadataStore, WorkerPeer, WorkerReclaimedChild, WorkerSpawnedChild,
|
||||
|
||||
@@ -0,0 +1,205 @@
|
||||
//! Session-owned storage for large pasted-input artifacts.
|
||||
|
||||
use std::fs;
|
||||
use std::io::Write as _;
|
||||
use std::path::Path;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use fs4::fs_std::FileExt;
|
||||
use protocol::{PasteArtifactAvailability, PasteArtifactMediaType, PasteArtifactRef};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
use crate::StoreError;
|
||||
|
||||
/// Bounded storage policy applied before a large paste becomes durable input.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct PasteArtifactLimits {
|
||||
pub max_artifact_bytes: u64,
|
||||
pub max_session_bytes: u64,
|
||||
pub max_session_artifacts: u64,
|
||||
}
|
||||
|
||||
impl Default for PasteArtifactLimits {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
max_artifact_bytes: 8 * 1024 * 1024,
|
||||
max_session_bytes: 64 * 1024 * 1024,
|
||||
max_session_artifacts: 1_024,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Integrity-bearing on-disk record. The body and metadata are committed in one
|
||||
/// atomic file replacement so readers never observe a half-written artifact.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) struct StoredPasteArtifact {
|
||||
pub reference: PasteArtifactRef,
|
||||
pub content: String,
|
||||
}
|
||||
|
||||
pub(crate) fn stored_paste_usage(artifact_dir: &Path) -> Result<(u64, u64), StoreError> {
|
||||
if !artifact_dir.exists() {
|
||||
return Ok((0, 0));
|
||||
}
|
||||
let mut aggregate = 0_u64;
|
||||
let mut artifact_count = 0_u64;
|
||||
for entry in fs::read_dir(artifact_dir)? {
|
||||
let path = entry?.path();
|
||||
let Some(name) = path.file_name().and_then(|value| value.to_str()) else {
|
||||
continue;
|
||||
};
|
||||
if !name.ends_with(".json") || name.ends_with(".file.json") {
|
||||
continue;
|
||||
}
|
||||
let stored: StoredPasteArtifact = serde_json::from_slice(&fs::read(&path)?)?;
|
||||
verify(&stored, &stored.reference.artifact_id)?;
|
||||
artifact_count = artifact_count.checked_add(1).ok_or_else(|| {
|
||||
StoreError::PasteArtifactLimit("session artifact count overflow".to_string())
|
||||
})?;
|
||||
aggregate = aggregate
|
||||
.checked_add(stored.reference.byte_len)
|
||||
.ok_or_else(|| {
|
||||
StoreError::PasteArtifactLimit("session aggregate size overflow".to_string())
|
||||
})?;
|
||||
}
|
||||
Ok((aggregate, artifact_count))
|
||||
}
|
||||
|
||||
pub(crate) fn write_to_dir(
|
||||
artifact_dir: &Path,
|
||||
source_entry_id: &str,
|
||||
content: &str,
|
||||
limits: PasteArtifactLimits,
|
||||
) -> Result<PasteArtifactRef, StoreError> {
|
||||
let byte_len = content.len() as u64;
|
||||
if byte_len > limits.max_artifact_bytes {
|
||||
return Err(StoreError::PasteArtifactLimit(format!(
|
||||
"artifact has {byte_len} bytes; maximum is {}",
|
||||
limits.max_artifact_bytes
|
||||
)));
|
||||
}
|
||||
fs::create_dir_all(artifact_dir)?;
|
||||
let aggregate_lock = fs::OpenOptions::new()
|
||||
.create(true)
|
||||
.read(true)
|
||||
.write(true)
|
||||
.open(artifact_dir.join(".aggregate.lock"))?;
|
||||
FileExt::lock_exclusive(&aggregate_lock)?;
|
||||
let (paste_bytes, artifact_count) = stored_paste_usage(artifact_dir)?;
|
||||
let (uploaded_bytes, uploaded_count) =
|
||||
crate::uploaded_file::stored_uploaded_file_usage(artifact_dir)?;
|
||||
let aggregate = paste_bytes.checked_add(uploaded_bytes).ok_or_else(|| {
|
||||
StoreError::PasteArtifactLimit("session aggregate size overflow".to_string())
|
||||
})?;
|
||||
let artifact_count = artifact_count.checked_add(uploaded_count).ok_or_else(|| {
|
||||
StoreError::PasteArtifactLimit("session artifact count overflow".to_string())
|
||||
})?;
|
||||
let projected = aggregate.checked_add(byte_len).ok_or_else(|| {
|
||||
StoreError::PasteArtifactLimit("session aggregate size overflow".to_string())
|
||||
})?;
|
||||
if projected > limits.max_session_bytes {
|
||||
return Err(StoreError::PasteArtifactLimit(format!(
|
||||
"session artifacts would use {projected} bytes; maximum is {}",
|
||||
limits.max_session_bytes
|
||||
)));
|
||||
}
|
||||
if artifact_count >= limits.max_session_artifacts {
|
||||
return Err(StoreError::PasteArtifactLimit(format!(
|
||||
"session already has {artifact_count} artifacts; maximum is {}",
|
||||
limits.max_session_artifacts
|
||||
)));
|
||||
}
|
||||
|
||||
let artifact_id = uuid::Uuid::now_v7().to_string();
|
||||
let created_at_ms = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map_err(|error| StoreError::PasteArtifactIntegrity(error.to_string()))?
|
||||
.as_millis() as u64;
|
||||
let reference = PasteArtifactRef {
|
||||
artifact_id: artifact_id.clone(),
|
||||
created_at_ms,
|
||||
media_type: PasteArtifactMediaType::TextPlainUtf8,
|
||||
availability: PasteArtifactAvailability::Available,
|
||||
byte_len,
|
||||
char_count: content.chars().count() as u64,
|
||||
line_count: line_count(content),
|
||||
sha256: sha256_hex(content),
|
||||
source_entry_id: source_entry_id.to_string(),
|
||||
};
|
||||
let bytes = serde_json::to_vec(&StoredPasteArtifact {
|
||||
reference: reference.clone(),
|
||||
content: content.to_string(),
|
||||
})?;
|
||||
let target = artifact_dir.join(format!("{artifact_id}.json"));
|
||||
let temporary = artifact_dir.join(format!(".{artifact_id}.tmp"));
|
||||
let mut file = fs::OpenOptions::new()
|
||||
.create_new(true)
|
||||
.write(true)
|
||||
.open(&temporary)?;
|
||||
if let Err(error) = file.write_all(&bytes).and_then(|_| file.sync_all()) {
|
||||
let _ = fs::remove_file(&temporary);
|
||||
return Err(error.into());
|
||||
}
|
||||
if let Err(error) = fs::rename(&temporary, &target) {
|
||||
let _ = fs::remove_file(&temporary);
|
||||
return Err(error.into());
|
||||
}
|
||||
if let Ok(directory) = fs::File::open(artifact_dir) {
|
||||
directory.sync_all()?;
|
||||
}
|
||||
Ok(reference)
|
||||
}
|
||||
|
||||
pub(crate) fn read_from_dir(
|
||||
artifact_dir: &Path,
|
||||
artifact_id: &str,
|
||||
) -> Result<(PasteArtifactRef, String), StoreError> {
|
||||
let parsed = uuid::Uuid::parse_str(artifact_id)
|
||||
.map_err(|_| StoreError::PasteArtifactNotFound(artifact_id.to_string()))?;
|
||||
if parsed.to_string() != artifact_id {
|
||||
return Err(StoreError::PasteArtifactNotFound(artifact_id.to_string()));
|
||||
}
|
||||
let path = artifact_dir.join(format!("{artifact_id}.json"));
|
||||
let bytes = match fs::read(path) {
|
||||
Ok(bytes) => bytes,
|
||||
Err(error) if error.kind() == std::io::ErrorKind::NotFound => {
|
||||
return Err(StoreError::PasteArtifactNotFound(artifact_id.to_string()));
|
||||
}
|
||||
Err(error) => return Err(error.into()),
|
||||
};
|
||||
let stored: StoredPasteArtifact = serde_json::from_slice(&bytes)?;
|
||||
verify(&stored, artifact_id)?;
|
||||
Ok((stored.reference, stored.content))
|
||||
}
|
||||
|
||||
fn verify(stored: &StoredPasteArtifact, artifact_id: &str) -> Result<(), StoreError> {
|
||||
let actual_digest = sha256_hex(&stored.content);
|
||||
if stored.reference.artifact_id != artifact_id
|
||||
|| stored.reference.created_at_ms == 0
|
||||
|| stored.reference.media_type != PasteArtifactMediaType::TextPlainUtf8
|
||||
|| stored.reference.availability != PasteArtifactAvailability::Available
|
||||
|| stored.reference.byte_len != stored.content.len() as u64
|
||||
|| stored.reference.char_count != stored.content.chars().count() as u64
|
||||
|| stored.reference.line_count != line_count(&stored.content)
|
||||
|| stored.reference.sha256 != actual_digest
|
||||
{
|
||||
return Err(StoreError::PasteArtifactIntegrity(artifact_id.to_string()));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn sha256_hex(content: &str) -> String {
|
||||
Sha256::digest(content.as_bytes())
|
||||
.iter()
|
||||
.map(|byte| format!("{byte:02x}"))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn line_count(content: &str) -> u64 {
|
||||
if content.is_empty() {
|
||||
0
|
||||
} else {
|
||||
content.lines().count().max(1) as u64
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,559 @@
|
||||
use base64::{
|
||||
Engine as _,
|
||||
engine::general_purpose::{STANDARD as BASE64, URL_SAFE_NO_PAD},
|
||||
};
|
||||
use protocol::{
|
||||
Segment, SessionContentPart, SessionEntryProvenance, SessionMessageRole, SessionSnapshot,
|
||||
SessionSnapshotEntry, SessionSnapshotEntryData, SessionToolAttachment,
|
||||
};
|
||||
|
||||
use crate::{
|
||||
LogEntry, LoggedContentPart, LoggedHistoryEntry, LoggedItem, LoggedRole,
|
||||
LoggedSessionHistoryOrigin, SessionId, SystemItem,
|
||||
};
|
||||
|
||||
/// Project a complete current-segment log. A valid segment starts with one
|
||||
/// canonical annotated SegmentStart record; malformed partial input uses the
|
||||
/// nil session only to keep the public failure projection deterministic.
|
||||
pub fn project_current_session_snapshot(log: &[LogEntry]) -> SessionSnapshot {
|
||||
let session_id = log.iter().find_map(|entry| match entry {
|
||||
LogEntry::AnnotatedSegmentStart { session_id, .. } => Some(*session_id),
|
||||
_ => None,
|
||||
});
|
||||
project_session_snapshot(session_id.unwrap_or_else(SessionId::nil), log)
|
||||
}
|
||||
|
||||
/// Project the current durable segment into the only public session-history
|
||||
/// representation. Append-log records remain an internal persistence format.
|
||||
pub fn project_session_snapshot(session_id: SessionId, log: &[LogEntry]) -> SessionSnapshot {
|
||||
let mut session_key = session_id;
|
||||
let mut entries = Vec::new();
|
||||
|
||||
for (log_index, record) in log.iter().enumerate() {
|
||||
match record {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts,
|
||||
session_id,
|
||||
history,
|
||||
..
|
||||
} => {
|
||||
session_key = *session_id;
|
||||
entries.clear();
|
||||
extend_history(&mut entries, history, None, *ts);
|
||||
}
|
||||
LogEntry::InputSegmentsCheckpoint { user_segments, .. } => {
|
||||
let mut segments = user_segments.iter();
|
||||
for entry in &mut entries {
|
||||
let is_user = matches!(
|
||||
&entry.data,
|
||||
SessionSnapshotEntryData::UserInput { .. }
|
||||
| SessionSnapshotEntryData::Message {
|
||||
role: SessionMessageRole::User,
|
||||
..
|
||||
}
|
||||
);
|
||||
if is_user && let Some(checkpoint) = segments.next() {
|
||||
entry.data = SessionSnapshotEntryData::UserInput {
|
||||
segments: checkpoint.clone(),
|
||||
};
|
||||
}
|
||||
}
|
||||
}
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts,
|
||||
segments,
|
||||
history,
|
||||
..
|
||||
} => extend_history(&mut entries, history, Some(segments), *ts),
|
||||
LogEntry::AnnotatedAssistantItem { ts, entry }
|
||||
| LogEntry::AnnotatedToolResult { ts, entry } => {
|
||||
if let Some(data) = project_item(&entry.item) {
|
||||
entries.push(history_entry(entry, *ts, data));
|
||||
}
|
||||
}
|
||||
LogEntry::AnnotatedSystemItem { ts, entry } => entries.push(system_entry(
|
||||
&entry.item,
|
||||
entry.metadata.entry_id.0.clone(),
|
||||
*ts,
|
||||
provenance(&entry.metadata.origin),
|
||||
derivation_ids(entry),
|
||||
)),
|
||||
LogEntry::RunErrored { ts, message, .. } => entries.push(legacy_entry(
|
||||
&session_key,
|
||||
log_index,
|
||||
0,
|
||||
*ts,
|
||||
SessionSnapshotEntryData::RunError {
|
||||
message: message.clone(),
|
||||
},
|
||||
)),
|
||||
// Run checkpoints, configuration, usage, and extension state are
|
||||
// controller/storage authority rather than committed conversation.
|
||||
LogEntry::Invoke { .. }
|
||||
| LogEntry::TurnEnd { .. }
|
||||
| LogEntry::RunCompleted { .. }
|
||||
| LogEntry::ActiveRunCheckpoint { .. }
|
||||
| LogEntry::PausedTurnAbandoned { .. }
|
||||
| LogEntry::ConfigChanged { .. }
|
||||
| LogEntry::LlmUsage { .. }
|
||||
| LogEntry::Extension { .. } => {}
|
||||
}
|
||||
}
|
||||
|
||||
SessionSnapshot { entries }
|
||||
}
|
||||
|
||||
fn extend_history(
|
||||
output: &mut Vec<SessionSnapshotEntry>,
|
||||
history: &[LoggedHistoryEntry],
|
||||
input_segments: Option<&Vec<Segment>>,
|
||||
timestamp: u64,
|
||||
) {
|
||||
let mut attached_segments = false;
|
||||
for entry in history {
|
||||
let data = if !attached_segments
|
||||
&& input_segments.is_some()
|
||||
&& matches!(
|
||||
&entry.item,
|
||||
LoggedItem::Message {
|
||||
role: LoggedRole::User,
|
||||
..
|
||||
}
|
||||
) {
|
||||
attached_segments = true;
|
||||
SessionSnapshotEntryData::UserInput {
|
||||
segments: input_segments.cloned().unwrap_or_default(),
|
||||
}
|
||||
} else {
|
||||
let Some(data) = project_item(&entry.item) else {
|
||||
continue;
|
||||
};
|
||||
data
|
||||
};
|
||||
output.push(history_entry(entry, timestamp, data));
|
||||
}
|
||||
}
|
||||
|
||||
fn history_entry(
|
||||
entry: &LoggedHistoryEntry,
|
||||
timestamp: u64,
|
||||
data: SessionSnapshotEntryData,
|
||||
) -> SessionSnapshotEntry {
|
||||
SessionSnapshotEntry {
|
||||
entry_id: entry.metadata.entry_id.0.clone(),
|
||||
timestamp,
|
||||
provenance: provenance(&entry.metadata.origin),
|
||||
derived_from: entry
|
||||
.metadata
|
||||
.derivation
|
||||
.as_ref()
|
||||
.map(|derivation| {
|
||||
derivation
|
||||
.sources
|
||||
.iter()
|
||||
.map(|source| source.0.clone())
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default(),
|
||||
data,
|
||||
}
|
||||
}
|
||||
|
||||
fn derivation_ids(entry: &crate::LoggedSystemHistoryEntry) -> Vec<String> {
|
||||
entry
|
||||
.metadata
|
||||
.derivation
|
||||
.as_ref()
|
||||
.map(|derivation| {
|
||||
derivation
|
||||
.sources
|
||||
.iter()
|
||||
.map(|source| source.0.clone())
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
fn legacy_entry(
|
||||
session_key: &SessionId,
|
||||
log_index: usize,
|
||||
item_index: usize,
|
||||
timestamp: u64,
|
||||
data: SessionSnapshotEntryData,
|
||||
) -> SessionSnapshotEntry {
|
||||
SessionSnapshotEntry {
|
||||
entry_id: legacy_entry_id(session_key, log_index, item_index),
|
||||
timestamp,
|
||||
provenance: SessionEntryProvenance::LegacyUnknown,
|
||||
derived_from: Vec::new(),
|
||||
data,
|
||||
}
|
||||
}
|
||||
|
||||
fn legacy_entry_id(session_key: &SessionId, log_index: usize, item_index: usize) -> String {
|
||||
let mut identity = Vec::with_capacity(32);
|
||||
identity.extend_from_slice(session_key.as_bytes());
|
||||
identity.extend_from_slice(&(log_index as u64).to_be_bytes());
|
||||
identity.extend_from_slice(&(item_index as u64).to_be_bytes());
|
||||
format!("l-{}", URL_SAFE_NO_PAD.encode(identity))
|
||||
}
|
||||
|
||||
fn provenance(origin: &LoggedSessionHistoryOrigin) -> SessionEntryProvenance {
|
||||
match origin {
|
||||
LoggedSessionHistoryOrigin::HumanInput { .. } => SessionEntryProvenance::HumanInput,
|
||||
LoggedSessionHistoryOrigin::WorkerInput { .. } => SessionEntryProvenance::WorkerInput,
|
||||
LoggedSessionHistoryOrigin::FlowInstruction { .. } => {
|
||||
SessionEntryProvenance::FlowInstruction
|
||||
}
|
||||
LoggedSessionHistoryOrigin::BackendInstruction { .. } => {
|
||||
SessionEntryProvenance::BackendInstruction
|
||||
}
|
||||
LoggedSessionHistoryOrigin::ModelOutput { .. } => SessionEntryProvenance::ModelOutput,
|
||||
LoggedSessionHistoryOrigin::ToolOutput { .. } => SessionEntryProvenance::ToolOutput,
|
||||
LoggedSessionHistoryOrigin::DerivedSummary => SessionEntryProvenance::DerivedSummary,
|
||||
LoggedSessionHistoryOrigin::LegacyUnknown => SessionEntryProvenance::LegacyUnknown,
|
||||
}
|
||||
}
|
||||
|
||||
fn project_item(item: &LoggedItem) -> Option<SessionSnapshotEntryData> {
|
||||
match item {
|
||||
LoggedItem::Message { role, content } => {
|
||||
let role = match role {
|
||||
LoggedRole::User => SessionMessageRole::User,
|
||||
LoggedRole::Assistant => SessionMessageRole::Assistant,
|
||||
// System prompts and instruction history never cross the public
|
||||
// snapshot boundary. Typed SystemItems have separate records.
|
||||
LoggedRole::System => return None,
|
||||
};
|
||||
Some(SessionSnapshotEntryData::Message {
|
||||
role,
|
||||
content: content
|
||||
.iter()
|
||||
.map(|part| match part {
|
||||
LoggedContentPart::Text { text } => {
|
||||
SessionContentPart::Text { text: text.clone() }
|
||||
}
|
||||
LoggedContentPart::Refusal { refusal } => SessionContentPart::Refusal {
|
||||
refusal: refusal.clone(),
|
||||
},
|
||||
})
|
||||
.collect(),
|
||||
})
|
||||
}
|
||||
LoggedItem::ToolCall {
|
||||
call_id,
|
||||
name,
|
||||
arguments,
|
||||
} => Some(SessionSnapshotEntryData::ToolCall {
|
||||
call_id: call_id.clone(),
|
||||
name: name.clone(),
|
||||
arguments: arguments.clone(),
|
||||
}),
|
||||
LoggedItem::ToolResult {
|
||||
call_id,
|
||||
summary,
|
||||
content,
|
||||
is_error,
|
||||
attachments,
|
||||
..
|
||||
} => Some(SessionSnapshotEntryData::ToolResult {
|
||||
call_id: call_id.clone(),
|
||||
summary: summary.clone(),
|
||||
content: content.clone(),
|
||||
is_error: *is_error,
|
||||
attachments: attachments
|
||||
.iter()
|
||||
.map(|attachment| match attachment {
|
||||
crate::logged_item::LoggedAttachment::Image { mime_type, data } => {
|
||||
SessionToolAttachment {
|
||||
media_type: mime_type.clone(),
|
||||
data_base64: BASE64.encode(data),
|
||||
}
|
||||
}
|
||||
})
|
||||
.collect(),
|
||||
}),
|
||||
// Hidden model reasoning is never observable.
|
||||
LoggedItem::Reasoning { .. } => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn system_entry(
|
||||
item: &SystemItem,
|
||||
entry_id: String,
|
||||
timestamp: u64,
|
||||
provenance: SessionEntryProvenance,
|
||||
derived_from: Vec<String>,
|
||||
) -> SessionSnapshotEntry {
|
||||
let mut data = serde_json::to_value(item).ok();
|
||||
if let Some(serde_json::Value::Object(object)) = data.as_mut() {
|
||||
object.remove("prompt_provenance");
|
||||
}
|
||||
let item_kind = data
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("kind"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("system_item")
|
||||
.to_owned();
|
||||
SessionSnapshotEntry {
|
||||
entry_id,
|
||||
timestamp,
|
||||
provenance,
|
||||
derived_from,
|
||||
data: SessionSnapshotEntryData::SystemItem {
|
||||
item_kind,
|
||||
content: item.history_text(),
|
||||
data,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use agen::llm_client::RequestConfig;
|
||||
|
||||
use super::*;
|
||||
use crate::{
|
||||
LoggedHistoryDerivation, LoggedSessionHistoryEntryId, LoggedSessionHistoryMetadata,
|
||||
LoggedWorkerSubject,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn current_projection_is_stable_and_hides_reasoning_and_system_prompts() {
|
||||
let session_id = crate::new_session_id();
|
||||
let log = vec![LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1,
|
||||
session_id,
|
||||
system_prompt: None,
|
||||
config: RequestConfig::default(),
|
||||
history: vec![
|
||||
LoggedItem::Message {
|
||||
role: LoggedRole::System,
|
||||
content: vec![LoggedContentPart::Text {
|
||||
text: "secret prompt".into(),
|
||||
}],
|
||||
},
|
||||
LoggedItem::Reasoning {
|
||||
text: "secret reasoning".into(),
|
||||
summary: Vec::new(),
|
||||
encrypted_content: None,
|
||||
signature: None,
|
||||
},
|
||||
LoggedItem::Message {
|
||||
role: LoggedRole::Assistant,
|
||||
content: vec![LoggedContentPart::Text {
|
||||
text: "visible".into(),
|
||||
}],
|
||||
},
|
||||
]
|
||||
.into_iter()
|
||||
.map(|item| LoggedHistoryEntry {
|
||||
item,
|
||||
metadata: LoggedSessionHistoryMetadata {
|
||||
entry_id: LoggedSessionHistoryEntryId::new(),
|
||||
origin: LoggedSessionHistoryOrigin::LegacyUnknown,
|
||||
derivation: None,
|
||||
},
|
||||
})
|
||||
.collect(),
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
}];
|
||||
|
||||
let first = project_session_snapshot(session_id, &log);
|
||||
let second = project_session_snapshot(session_id, &log);
|
||||
assert_eq!(first, second);
|
||||
assert_eq!(first.entries.len(), 1);
|
||||
assert_eq!(first.entries[0].timestamp, 1);
|
||||
assert_eq!(
|
||||
first.entries[0].provenance,
|
||||
SessionEntryProvenance::LegacyUnknown
|
||||
);
|
||||
let json = serde_json::to_string(&first).unwrap();
|
||||
assert!(!json.contains("secret prompt"));
|
||||
assert!(!json.contains("secret reasoning"));
|
||||
assert!(json.contains("visible"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn compacted_checkpoint_restores_uploaded_file_segments() {
|
||||
let session_id = crate::new_session_id();
|
||||
let user_entry_id = LoggedSessionHistoryEntryId::new();
|
||||
let file = protocol::UploadedFileRef {
|
||||
artifact_id: "019ca7c8-57b6-7f05-8edf-524147aba7b3".into(),
|
||||
file_name: "notes.md".into(),
|
||||
media_type: "text/markdown".into(),
|
||||
created_at_ms: 7,
|
||||
availability: protocol::UploadedFileAvailability::Available,
|
||||
byte_len: 12,
|
||||
sha256: "a".repeat(64),
|
||||
source_entry_id: Some(user_entry_id.0.clone()),
|
||||
};
|
||||
let segment = Segment::UploadedFile { file };
|
||||
let log = vec![
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 10,
|
||||
session_id,
|
||||
system_prompt: None,
|
||||
config: RequestConfig::default(),
|
||||
history: vec![LoggedHistoryEntry {
|
||||
item: LoggedItem::Message {
|
||||
role: LoggedRole::User,
|
||||
content: vec![LoggedContentPart::Text {
|
||||
text: "[Attached file: notes.md]".into(),
|
||||
}],
|
||||
},
|
||||
metadata: LoggedSessionHistoryMetadata {
|
||||
entry_id: user_entry_id,
|
||||
origin: LoggedSessionHistoryOrigin::HumanInput {
|
||||
account_id: "account-1".into(),
|
||||
},
|
||||
derivation: None,
|
||||
},
|
||||
}],
|
||||
forked_from: None,
|
||||
compacted_from: Some(crate::SegmentOrigin {
|
||||
segment_id: crate::new_segment_id(),
|
||||
at_turn_index: 1,
|
||||
}),
|
||||
},
|
||||
LogEntry::InputSegmentsCheckpoint {
|
||||
ts: 10,
|
||||
user_segments: vec![vec![segment.clone()]],
|
||||
},
|
||||
];
|
||||
|
||||
let snapshot = project_current_session_snapshot(&log);
|
||||
assert_eq!(
|
||||
snapshot.entries[0].data,
|
||||
SessionSnapshotEntryData::UserInput {
|
||||
segments: vec![segment]
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn annotated_user_input_attaches_segments_to_first_user_role_entry_for_any_origin() {
|
||||
let session_id = crate::new_session_id();
|
||||
let segments = vec![Segment::Text {
|
||||
content: "normal submit".into(),
|
||||
}];
|
||||
|
||||
for origin in [
|
||||
LoggedSessionHistoryOrigin::LegacyUnknown,
|
||||
LoggedSessionHistoryOrigin::FlowInstruction {
|
||||
selector: "builtin:coder-review".into(),
|
||||
definition_id: "flow-definition".into(),
|
||||
definition_revision: 7,
|
||||
instance_id: "flow-instance".into(),
|
||||
state_id: "implement".into(),
|
||||
},
|
||||
] {
|
||||
let user_entry_id = LoggedSessionHistoryEntryId::new();
|
||||
let source_entry_id = LoggedSessionHistoryEntryId::new();
|
||||
let log = vec![
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1,
|
||||
session_id,
|
||||
system_prompt: None,
|
||||
config: RequestConfig::default(),
|
||||
history: Vec::new(),
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
},
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 2,
|
||||
segments: segments.clone(),
|
||||
history: vec![
|
||||
LoggedHistoryEntry {
|
||||
item: LoggedItem::Message {
|
||||
role: LoggedRole::System,
|
||||
content: vec![LoggedContentPart::Text {
|
||||
text: "flow instruction".into(),
|
||||
}],
|
||||
},
|
||||
metadata: LoggedSessionHistoryMetadata {
|
||||
entry_id: LoggedSessionHistoryEntryId::new(),
|
||||
origin: LoggedSessionHistoryOrigin::FlowInstruction {
|
||||
selector: "builtin:coder-review".into(),
|
||||
definition_id: "flow-definition".into(),
|
||||
definition_revision: 7,
|
||||
instance_id: "flow-instance".into(),
|
||||
state_id: "implement".into(),
|
||||
},
|
||||
derivation: None,
|
||||
},
|
||||
},
|
||||
LoggedHistoryEntry {
|
||||
item: LoggedItem::Message {
|
||||
role: LoggedRole::User,
|
||||
content: vec![LoggedContentPart::Text {
|
||||
text: "normal submit".into(),
|
||||
}],
|
||||
},
|
||||
metadata: LoggedSessionHistoryMetadata {
|
||||
entry_id: user_entry_id.clone(),
|
||||
origin: origin.clone(),
|
||||
derivation: Some(LoggedHistoryDerivation {
|
||||
sources: vec![source_entry_id.clone()],
|
||||
}),
|
||||
},
|
||||
},
|
||||
],
|
||||
extensions: Vec::new(),
|
||||
},
|
||||
];
|
||||
|
||||
let snapshot = project_current_session_snapshot(&log);
|
||||
assert_eq!(snapshot.entries.len(), 1);
|
||||
assert_eq!(snapshot.entries[0].entry_id, user_entry_id.0);
|
||||
assert_eq!(snapshot.entries[0].provenance, provenance(&origin));
|
||||
assert_eq!(snapshot.entries[0].derived_from, vec![source_entry_id.0]);
|
||||
assert_eq!(
|
||||
snapshot.entries[0].data,
|
||||
SessionSnapshotEntryData::UserInput {
|
||||
segments: segments.clone(),
|
||||
}
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn annotated_projection_preserves_identity_and_provenance() {
|
||||
let session_id = crate::new_session_id();
|
||||
let metadata = LoggedSessionHistoryMetadata {
|
||||
entry_id: LoggedSessionHistoryEntryId::new(),
|
||||
origin: LoggedSessionHistoryOrigin::ModelOutput {
|
||||
worker: LoggedWorkerSubject {
|
||||
workspace_id: None,
|
||||
runtime_id: None,
|
||||
worker_id: "worker".into(),
|
||||
},
|
||||
},
|
||||
derivation: None,
|
||||
};
|
||||
let expected_id = metadata.entry_id.0.clone();
|
||||
let log = vec![LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1,
|
||||
session_id,
|
||||
system_prompt: None,
|
||||
config: RequestConfig::default(),
|
||||
history: vec![LoggedHistoryEntry {
|
||||
item: LoggedItem::Message {
|
||||
role: LoggedRole::Assistant,
|
||||
content: vec![LoggedContentPart::Text { text: "ok".into() }],
|
||||
},
|
||||
metadata,
|
||||
}],
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
}];
|
||||
|
||||
let snapshot = project_session_snapshot(session_id, &log);
|
||||
assert_eq!(snapshot.entries[0].entry_id, expected_id);
|
||||
assert_eq!(
|
||||
snapshot.entries[0].provenance,
|
||||
SessionEntryProvenance::ModelOutput
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -4,11 +4,9 @@
|
||||
//! The caller (typically Worker) holds the Engine directly and calls these
|
||||
//! functions after state-mutating operations.
|
||||
|
||||
use crate::logged_item::{LoggedItem, to_logged};
|
||||
use crate::segment_log::{self, LogEntry, SegmentOrigin};
|
||||
use crate::store::{Store, StoreError};
|
||||
use crate::system_item::SystemItem;
|
||||
use crate::{SegmentId, SessionId};
|
||||
use crate::{LoggedHistoryEntry, LoggedSystemHistoryEntry, SegmentId, SessionId};
|
||||
use agen::EngineResult;
|
||||
use agen::llm_client::RequestConfig;
|
||||
use agen::llm_client::types::Item;
|
||||
@@ -18,7 +16,34 @@ use protocol::Segment;
|
||||
pub struct SegmentStartState<'a> {
|
||||
pub system_prompt: Option<&'a str>,
|
||||
pub config: &'a RequestConfig,
|
||||
pub history: &'a [Item],
|
||||
pub history: Vec<LoggedHistoryEntry>,
|
||||
pub user_segments: Vec<Vec<Segment>>,
|
||||
}
|
||||
|
||||
fn seed_entries(
|
||||
ts: u64,
|
||||
session_id: SessionId,
|
||||
state: SegmentStartState<'_>,
|
||||
forked_from: Option<SegmentOrigin>,
|
||||
compacted_from: Option<SegmentOrigin>,
|
||||
) -> Vec<LogEntry> {
|
||||
let entry = LogEntry::AnnotatedSegmentStart {
|
||||
ts,
|
||||
session_id,
|
||||
system_prompt: state.system_prompt.map(String::from),
|
||||
config: state.config.clone(),
|
||||
history: state.history,
|
||||
forked_from,
|
||||
compacted_from,
|
||||
};
|
||||
let mut entries = vec![entry];
|
||||
if !state.user_segments.is_empty() {
|
||||
entries.push(LogEntry::InputSegmentsCheckpoint {
|
||||
ts,
|
||||
user_segments: state.user_segments,
|
||||
});
|
||||
}
|
||||
entries
|
||||
}
|
||||
|
||||
/// Create a new session + initial segment, writing the initial
|
||||
@@ -44,16 +69,8 @@ pub fn create_segment_with_ids(
|
||||
segment_id: SegmentId,
|
||||
state: SegmentStartState<'_>,
|
||||
) -> Result<(), StoreError> {
|
||||
let entry = LogEntry::SegmentStart {
|
||||
ts: segment_log::now_millis(),
|
||||
session_id,
|
||||
system_prompt: state.system_prompt.map(String::from),
|
||||
config: state.config.clone(),
|
||||
history: to_logged(state.history),
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
};
|
||||
store.append(session_id, segment_id, &entry)
|
||||
let entries = seed_entries(segment_log::now_millis(), session_id, state, None, None);
|
||||
store.create_segment(session_id, segment_id, &entries)
|
||||
}
|
||||
|
||||
/// Create a compacted segment from an existing one. Inherits the source's
|
||||
@@ -70,19 +87,17 @@ pub fn create_compacted_segment(
|
||||
source_turn_count: usize,
|
||||
) -> Result<SegmentId, StoreError> {
|
||||
let segment_id = crate::new_segment_id();
|
||||
let entry = LogEntry::SegmentStart {
|
||||
ts: segment_log::now_millis(),
|
||||
session_id: source_session_id,
|
||||
system_prompt: state.system_prompt.map(String::from),
|
||||
config: state.config.clone(),
|
||||
history: to_logged(state.history),
|
||||
forked_from: None,
|
||||
compacted_from: Some(SegmentOrigin {
|
||||
let entries = seed_entries(
|
||||
segment_log::now_millis(),
|
||||
source_session_id,
|
||||
state,
|
||||
None,
|
||||
Some(SegmentOrigin {
|
||||
segment_id: source_segment_id,
|
||||
at_turn_index: source_turn_count,
|
||||
}),
|
||||
};
|
||||
store.append(source_session_id, segment_id, &entry)?;
|
||||
);
|
||||
store.create_segment(source_session_id, segment_id, &entries)?;
|
||||
Ok(segment_id)
|
||||
}
|
||||
|
||||
@@ -154,21 +169,19 @@ pub fn ensure_head_or_fork(
|
||||
}
|
||||
let source_segment_id = *segment_id;
|
||||
let fork_id = crate::new_segment_id();
|
||||
let entry = LogEntry::SegmentStart {
|
||||
ts: segment_log::now_millis(),
|
||||
let entries = seed_entries(
|
||||
segment_log::now_millis(),
|
||||
session_id,
|
||||
system_prompt: state.system_prompt.map(String::from),
|
||||
config: state.config.clone(),
|
||||
history: to_logged(state.history),
|
||||
forked_from: Some(SegmentOrigin {
|
||||
state,
|
||||
Some(SegmentOrigin {
|
||||
segment_id: source_segment_id,
|
||||
at_turn_index,
|
||||
}),
|
||||
compacted_from: None,
|
||||
};
|
||||
store.create_segment(session_id, fork_id, &[entry])?;
|
||||
None,
|
||||
);
|
||||
store.create_segment(session_id, fork_id, &entries)?;
|
||||
*segment_id = fork_id;
|
||||
*entries_written = 1;
|
||||
*entries_written = entries.len();
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -183,8 +196,9 @@ pub fn save_user_input(
|
||||
session_id: SessionId,
|
||||
segment_id: SegmentId,
|
||||
segments: Vec<Segment>,
|
||||
history: Vec<LoggedHistoryEntry>,
|
||||
) -> Result<(), StoreError> {
|
||||
save_user_input_with_extensions(store, session_id, segment_id, segments, Vec::new())
|
||||
save_user_input_with_extensions(store, session_id, segment_id, segments, history, Vec::new())
|
||||
}
|
||||
|
||||
/// Atomically persist one typed user submission and Runtime-owned session
|
||||
@@ -194,15 +208,17 @@ pub fn save_user_input_with_extensions(
|
||||
session_id: SessionId,
|
||||
segment_id: SegmentId,
|
||||
segments: Vec<Segment>,
|
||||
history: Vec<LoggedHistoryEntry>,
|
||||
extensions: Vec<segment_log::SessionExtension>,
|
||||
) -> Result<(), StoreError> {
|
||||
append_entry(
|
||||
store,
|
||||
session_id,
|
||||
segment_id,
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: segment_log::now_millis(),
|
||||
segments,
|
||||
history,
|
||||
extensions,
|
||||
},
|
||||
)
|
||||
@@ -220,64 +236,57 @@ pub fn save_delta(
|
||||
store: &impl Store,
|
||||
session_id: SessionId,
|
||||
segment_id: SegmentId,
|
||||
new_items: &[Item],
|
||||
new_items: &[LoggedHistoryEntry],
|
||||
) -> Result<(), StoreError> {
|
||||
if new_items.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let ts = segment_log::now_millis();
|
||||
for item in new_items {
|
||||
for entry in new_items {
|
||||
let item = Item::from(entry.item.clone());
|
||||
if item.is_user_message() {
|
||||
// Already persisted by save_user_input at submit time.
|
||||
continue;
|
||||
}
|
||||
let entry = classify_history_item(item, ts);
|
||||
let entry = classify_logged_history_entry(entry.clone(), ts);
|
||||
append_entry(store, session_id, segment_id, entry)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Map one history item to its singular `LogEntry` form. Used by the
|
||||
/// fallback `save_delta` path and the controller's worker-callback
|
||||
/// classifier so write classification lives in one place.
|
||||
pub fn classify_history_item(item: &Item, ts: u64) -> LogEntry {
|
||||
/// Map one annotated history entry to its singular `LogEntry` form. Used by
|
||||
/// the fallback `save_delta` path and the controller's worker-callback
|
||||
/// classifier so write classification lives in one place without discarding
|
||||
/// identity or provenance.
|
||||
/// Map one already-annotated history entry to its singular canonical record
|
||||
/// without changing its identity or provenance.
|
||||
pub fn classify_logged_history_entry(entry: LoggedHistoryEntry, ts: u64) -> LogEntry {
|
||||
let item = Item::from(entry.item.clone());
|
||||
if item.is_tool_result() {
|
||||
LogEntry::ToolResult {
|
||||
ts,
|
||||
item: LoggedItem::from(item),
|
||||
}
|
||||
} else if item.is_assistant_message() || item.is_tool_call() || item.is_reasoning() {
|
||||
LogEntry::AssistantItem {
|
||||
ts,
|
||||
item: LoggedItem::from(item),
|
||||
}
|
||||
LogEntry::AnnotatedToolResult { ts, entry }
|
||||
} else {
|
||||
// Defensive: anything else (future Item kinds) routes through
|
||||
// AssistantItem rather than getting silently dropped.
|
||||
LogEntry::AssistantItem {
|
||||
ts,
|
||||
item: LoggedItem::from(item),
|
||||
}
|
||||
// Assistant messages, tool calls, reasoning, and future non-user
|
||||
// items all use the assistant-side canonical record.
|
||||
LogEntry::AnnotatedAssistantItem { ts, entry }
|
||||
}
|
||||
}
|
||||
|
||||
/// Append a single typed system item as `LogEntry::SystemItem`. Helper
|
||||
/// for the Worker-side interceptor commit path; mirrors the per-item
|
||||
/// commit shape used for assistant / tool result entries.
|
||||
/// Append one typed system item and its history metadata as a canonical
|
||||
/// `LogEntry::AnnotatedSystemItem`.
|
||||
pub fn append_system_item(
|
||||
store: &impl Store,
|
||||
session_id: SessionId,
|
||||
segment_id: SegmentId,
|
||||
item: SystemItem,
|
||||
entry: LoggedSystemHistoryEntry,
|
||||
) -> Result<(), StoreError> {
|
||||
append_entry(
|
||||
store,
|
||||
session_id,
|
||||
segment_id,
|
||||
LogEntry::SystemItem {
|
||||
LogEntry::AnnotatedSystemItem {
|
||||
ts: segment_log::now_millis(),
|
||||
item,
|
||||
entry,
|
||||
},
|
||||
)
|
||||
}
|
||||
@@ -426,20 +435,14 @@ pub fn save_config_changed(
|
||||
/// [`fork_at`] or [`ensure_head_or_fork`] instead.
|
||||
pub fn fork(
|
||||
store: &impl Store,
|
||||
source_session_id: SessionId,
|
||||
state: SegmentStartState<'_>,
|
||||
) -> Result<(SessionId, SegmentId), StoreError> {
|
||||
let session_id = crate::new_session_id();
|
||||
let fork_id = crate::new_segment_id();
|
||||
let entry = LogEntry::SegmentStart {
|
||||
ts: segment_log::now_millis(),
|
||||
session_id,
|
||||
system_prompt: state.system_prompt.map(String::from),
|
||||
config: state.config.clone(),
|
||||
history: to_logged(state.history),
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
};
|
||||
store.create_segment(session_id, fork_id, &[entry])?;
|
||||
let entries = seed_entries(segment_log::now_millis(), session_id, state, None, None);
|
||||
store.create_segment(session_id, fork_id, &entries)?;
|
||||
store.copy_committed_uploaded_files(source_session_id, session_id)?;
|
||||
Ok((session_id, fork_id))
|
||||
}
|
||||
|
||||
@@ -466,11 +469,18 @@ pub fn fork_at(
|
||||
) -> Result<SegmentId, StoreError> {
|
||||
let entries = store.read_all(source_session_id, source_id)?;
|
||||
let cut = if at_turn_index == 0 {
|
||||
// Branch directly after the SegmentStart (or whatever opens the
|
||||
// segment), before any turn completes.
|
||||
// Branch from the seeded state before any new turn completes. A typed
|
||||
// input checkpoint immediately following SegmentStart is part of that
|
||||
// seed and must stay atomic with its annotated history.
|
||||
entries
|
||||
.iter()
|
||||
.position(|e| !matches!(e, LogEntry::SegmentStart { .. }))
|
||||
.position(|entry| {
|
||||
!matches!(
|
||||
entry,
|
||||
LogEntry::AnnotatedSegmentStart { .. }
|
||||
| LogEntry::InputSegmentsCheckpoint { .. }
|
||||
)
|
||||
})
|
||||
.unwrap_or(entries.len())
|
||||
} else {
|
||||
entries
|
||||
@@ -482,19 +492,27 @@ pub fn fork_at(
|
||||
let state = segment_log::collect_state(&entries[..cut]);
|
||||
|
||||
let fork_id = crate::new_segment_id();
|
||||
let entry = LogEntry::SegmentStart {
|
||||
ts: segment_log::now_millis(),
|
||||
let ts = segment_log::now_millis();
|
||||
let entry = LogEntry::AnnotatedSegmentStart {
|
||||
ts,
|
||||
session_id: source_session_id,
|
||||
system_prompt: state.system_prompt,
|
||||
config: state.config,
|
||||
history: to_logged(&state.history),
|
||||
history: state.annotated_history,
|
||||
forked_from: Some(SegmentOrigin {
|
||||
segment_id: source_id,
|
||||
at_turn_index,
|
||||
}),
|
||||
compacted_from: None,
|
||||
};
|
||||
store.create_segment(source_session_id, fork_id, &[entry])?;
|
||||
let mut fork_entries = vec![entry];
|
||||
if !state.user_segments.is_empty() {
|
||||
fork_entries.push(LogEntry::InputSegmentsCheckpoint {
|
||||
ts,
|
||||
user_segments: state.user_segments,
|
||||
});
|
||||
}
|
||||
store.create_segment(source_session_id, fork_id, &fork_entries)?;
|
||||
Ok(fork_id)
|
||||
}
|
||||
|
||||
|
||||
@@ -16,7 +16,6 @@ use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::history::{LoggedHistoryEntry, LoggedSystemHistoryEntry};
|
||||
use crate::logged_item::LoggedItem;
|
||||
use crate::system_item::SystemItem;
|
||||
|
||||
/// A single segment log entry, serialized as one JSONL line.
|
||||
///
|
||||
@@ -50,28 +49,7 @@ impl SessionExtension {
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(tag = "kind", rename_all = "snake_case")]
|
||||
pub enum LogEntry {
|
||||
/// Segment start. Always the first entry in a segment log.
|
||||
/// For forked segments, `history` contains the seed state from the parent.
|
||||
SegmentStart {
|
||||
ts: u64,
|
||||
/// Session this segment belongs to. Compaction / fork inherits
|
||||
/// the source segment's session_id; only fresh "new conversation"
|
||||
/// segments mint a new session_id.
|
||||
session_id: crate::SessionId,
|
||||
system_prompt: Option<String>,
|
||||
config: RequestConfig,
|
||||
history: Vec<LoggedItem>,
|
||||
/// Origin: forked from a sibling segment at a specific turn boundary.
|
||||
/// The referenced segment is guaranteed to share `session_id`.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
forked_from: Option<SegmentOrigin>,
|
||||
/// Origin: compacted from a sibling segment at a specific turn boundary.
|
||||
/// The referenced segment is guaranteed to share `session_id`.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
compacted_from: Option<SegmentOrigin>,
|
||||
},
|
||||
|
||||
/// Schema-v2 segment seed. Retained entries keep their stable logical
|
||||
/// Canonical segment seed. Retained entries keep their stable logical
|
||||
/// identity and origin across fork/compaction/restore.
|
||||
AnnotatedSegmentStart {
|
||||
ts: u64,
|
||||
@@ -85,6 +63,14 @@ pub enum LogEntry {
|
||||
compacted_from: Option<SegmentOrigin>,
|
||||
},
|
||||
|
||||
/// Typed user-segment projection accompanying a compacted or forked
|
||||
/// SegmentStart history snapshot. This keeps attachment identity and
|
||||
/// metadata aligned with retained user entries without embedding bodies.
|
||||
InputSegmentsCheckpoint {
|
||||
ts: u64,
|
||||
user_segments: Vec<Vec<Segment>>,
|
||||
},
|
||||
|
||||
/// IDLE → active marker. Records the start of a new self-driving
|
||||
/// cycle (Invoke range). The range extends implicitly until the
|
||||
/// next `Invoke` entry; this entry carries the trigger only — the
|
||||
@@ -105,22 +91,7 @@ pub enum LogEntry {
|
||||
/// restore conservatively instead of re-running a dangling tool call.
|
||||
Invoke { ts: u64, trigger: InvokeKind },
|
||||
|
||||
/// User input accepted at submit time. Carries the original typed
|
||||
/// `Vec<Segment>` so clients can re-render typed atoms (paste chips,
|
||||
/// file refs) on segment restore.
|
||||
/// Replay flattens these into a `Item::user_message` for the worker
|
||||
/// history; the worker layer never sees segments directly.
|
||||
UserInput {
|
||||
ts: u64,
|
||||
segments: Vec<Segment>,
|
||||
/// Typed durable state committed atomically with this input record.
|
||||
/// Runtime-owned Flow invocation uses this to avoid a Backend-instance
|
||||
/// commit that can get ahead of Worker history.
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
extensions: Vec<SessionExtension>,
|
||||
},
|
||||
|
||||
/// Schema-v2 user submission with its exact model-visible entries. Typed
|
||||
/// Canonical user submission with its exact model-visible entries. Typed
|
||||
/// Flow instructions and caller-attributed input remain separate entries.
|
||||
AnnotatedUserInput {
|
||||
ts: u64,
|
||||
@@ -130,35 +101,19 @@ pub enum LogEntry {
|
||||
history: Vec<LoggedHistoryEntry>,
|
||||
},
|
||||
|
||||
/// Schema-v2 model output and metadata committed as one journal record.
|
||||
/// Canonical model output and metadata committed as one journal record.
|
||||
AnnotatedAssistantItem { ts: u64, entry: LoggedHistoryEntry },
|
||||
|
||||
/// One assistant-side item appended to history — assistant message,
|
||||
/// reasoning, or tool call. Singular: one entry per history item so
|
||||
/// the wire-side `Event::*` lane and on-disk LogEntry stay 1:1.
|
||||
AssistantItem { ts: u64, item: LoggedItem },
|
||||
|
||||
/// Schema-v2 tool output and metadata committed as one journal record.
|
||||
/// Canonical tool output and metadata committed as one journal record.
|
||||
AnnotatedToolResult { ts: u64, entry: LoggedHistoryEntry },
|
||||
|
||||
/// One tool-execution result appended to history.
|
||||
ToolResult { ts: u64, item: LoggedItem },
|
||||
|
||||
/// Schema-v2 typed system event and model-visible metadata committed
|
||||
/// Canonical typed system event and model-visible metadata committed
|
||||
/// together.
|
||||
AnnotatedSystemItem {
|
||||
ts: u64,
|
||||
entry: LoggedSystemHistoryEntry,
|
||||
},
|
||||
|
||||
/// One typed agent-injected system item: notification, child-Worker
|
||||
/// lifecycle event, `@<path>` / `/<slug>` resolution payload. Each
|
||||
/// `SystemItem` carries kind metadata that the LLM
|
||||
/// itself never sees (the LLM gets `Item::system_message` with the
|
||||
/// item's denormalised `body`), but live clients and replay paths
|
||||
/// dispatch on `kind` for typed rendering.
|
||||
SystemItem { ts: u64, item: SystemItem },
|
||||
|
||||
/// Turn boundary. Records the turn count after increment.
|
||||
TurnEnd { ts: u64, turn_count: usize },
|
||||
|
||||
@@ -260,6 +215,10 @@ pub struct RestoredState {
|
||||
pub system_prompt: Option<String>,
|
||||
pub config: RequestConfig,
|
||||
pub history: Vec<Item>,
|
||||
/// Canonical persisted history with stable identity and provenance. This is
|
||||
/// the authority for rewrites, forks, and annotated restore; `history` is
|
||||
/// retained as the model-facing item projection.
|
||||
pub annotated_history: Vec<LoggedHistoryEntry>,
|
||||
pub turn_count: usize,
|
||||
/// AgentTurns consumed by the active paused/yielded logical run.
|
||||
pub active_run_turn_count: Option<usize>,
|
||||
@@ -276,7 +235,7 @@ pub struct RestoredState {
|
||||
/// session-store は domain を不透明扱いし、各ドメインが自前で fold する。
|
||||
pub extensions: Vec<(String, serde_json::Value)>,
|
||||
/// User submissions in original typed form, in submit order.
|
||||
/// One entry per `LogEntry::UserInput`; the K-th entry corresponds to
|
||||
/// One entry per `LogEntry::AnnotatedUserInput`; the K-th entry corresponds to
|
||||
/// the K-th `Item::user_message` derived during replay (modulo
|
||||
/// pre-compaction history seeded via `SegmentStart.history`, whose
|
||||
/// original segments are not preserved). Used by clients to re-render
|
||||
@@ -291,6 +250,7 @@ pub fn collect_state(entries: &[LogEntry]) -> RestoredState {
|
||||
system_prompt: None,
|
||||
config: RequestConfig::default(),
|
||||
history: Vec::new(),
|
||||
annotated_history: Vec::new(),
|
||||
turn_count: 0,
|
||||
active_run_turn_count: None,
|
||||
last_run_interrupted: false,
|
||||
@@ -304,18 +264,6 @@ pub fn collect_state(entries: &[LogEntry]) -> RestoredState {
|
||||
state.entries_count += 1;
|
||||
|
||||
match entry {
|
||||
LogEntry::SegmentStart {
|
||||
session_id,
|
||||
system_prompt,
|
||||
config,
|
||||
history,
|
||||
..
|
||||
} => {
|
||||
state.session_id = Some(*session_id);
|
||||
state.system_prompt = system_prompt.clone();
|
||||
state.config = config.clone();
|
||||
state.history = history.iter().cloned().map(Item::from).collect();
|
||||
}
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
session_id,
|
||||
system_prompt,
|
||||
@@ -326,38 +274,29 @@ pub fn collect_state(entries: &[LogEntry]) -> RestoredState {
|
||||
state.session_id = Some(*session_id);
|
||||
state.system_prompt = system_prompt.clone();
|
||||
state.config = config.clone();
|
||||
state.annotated_history = history.clone();
|
||||
state.history = history
|
||||
.iter()
|
||||
.cloned()
|
||||
.map(|entry| Item::from(entry.item))
|
||||
.collect();
|
||||
}
|
||||
LogEntry::InputSegmentsCheckpoint { user_segments, .. } => {
|
||||
state.user_segments = user_segments.clone();
|
||||
}
|
||||
LogEntry::Invoke { .. } => {
|
||||
// A terminal run record below clears or refines this. If the
|
||||
// log ends first, restore must treat the turn as interrupted.
|
||||
state.last_run_interrupted = true;
|
||||
state.active_run_turn_count = Some(0);
|
||||
}
|
||||
LogEntry::UserInput {
|
||||
segments,
|
||||
extensions,
|
||||
..
|
||||
} => {
|
||||
let text = Segment::flatten_to_text(segments);
|
||||
state.history.push(Item::user_message(text));
|
||||
state.user_segments.push(segments.clone());
|
||||
state.extensions.extend(
|
||||
extensions
|
||||
.iter()
|
||||
.map(|extension| (extension.domain.clone(), extension.payload.clone())),
|
||||
);
|
||||
}
|
||||
LogEntry::AnnotatedUserInput {
|
||||
segments,
|
||||
extensions,
|
||||
history,
|
||||
..
|
||||
} => {
|
||||
state.annotated_history.extend(history.iter().cloned());
|
||||
state
|
||||
.history
|
||||
.extend(history.iter().cloned().map(|entry| Item::from(entry.item)));
|
||||
@@ -370,20 +309,16 @@ pub fn collect_state(entries: &[LogEntry]) -> RestoredState {
|
||||
}
|
||||
LogEntry::AnnotatedAssistantItem { entry, .. }
|
||||
| LogEntry::AnnotatedToolResult { entry, .. } => {
|
||||
state.annotated_history.push(entry.clone());
|
||||
state.history.push(Item::from(entry.item.clone()));
|
||||
}
|
||||
LogEntry::AnnotatedSystemItem { entry, .. } => {
|
||||
state.annotated_history.push(LoggedHistoryEntry {
|
||||
item: LoggedItem::from(entry.item.to_history_item()),
|
||||
metadata: entry.metadata.clone(),
|
||||
});
|
||||
state.history.push(entry.item.to_history_item());
|
||||
}
|
||||
LogEntry::AssistantItem { item, .. } => {
|
||||
state.history.push(Item::from(item.clone()));
|
||||
}
|
||||
LogEntry::ToolResult { item, .. } => {
|
||||
state.history.push(Item::from(item.clone()));
|
||||
}
|
||||
LogEntry::SystemItem { item, .. } => {
|
||||
state.history.push(item.to_history_item());
|
||||
}
|
||||
LogEntry::TurnEnd { turn_count, .. } => {
|
||||
if let Some(active_turn_count) = &mut state.active_run_turn_count {
|
||||
*active_turn_count += turn_count.saturating_sub(state.turn_count);
|
||||
@@ -465,6 +400,20 @@ pub fn now_millis() -> u64 {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::{
|
||||
LoggedSessionHistoryEntryId, LoggedSessionHistoryMetadata, LoggedSessionHistoryOrigin,
|
||||
};
|
||||
|
||||
fn annotated(item: Item) -> LoggedHistoryEntry {
|
||||
LoggedHistoryEntry {
|
||||
item: LoggedItem::from(item),
|
||||
metadata: LoggedSessionHistoryMetadata {
|
||||
entry_id: LoggedSessionHistoryEntryId::new(),
|
||||
origin: LoggedSessionHistoryOrigin::LegacyUnknown,
|
||||
derivation: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn replay_empty() {
|
||||
@@ -476,12 +425,12 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn replay_segment_start_sets_initial_state() {
|
||||
let state = collect_state(&[LogEntry::SegmentStart {
|
||||
let state = collect_state(&[LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: Some("You are helpful.".into()),
|
||||
config: RequestConfig::default().with_max_tokens(1024),
|
||||
history: vec![Item::user_message("seed").into()],
|
||||
history: vec![annotated(Item::user_message("seed"))],
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
}]);
|
||||
@@ -494,7 +443,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_full_turn() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -503,14 +452,15 @@ mod tests {
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
},
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 2000,
|
||||
extensions: vec![],
|
||||
segments: vec![Segment::text("Hello")],
|
||||
history: vec![annotated(Item::user_message("Hello"))],
|
||||
},
|
||||
LogEntry::AssistantItem {
|
||||
LogEntry::AnnotatedAssistantItem {
|
||||
ts: 3000,
|
||||
item: Item::assistant_message("Hi!").into(),
|
||||
entry: annotated(Item::assistant_message("Hi!")),
|
||||
},
|
||||
LogEntry::TurnEnd {
|
||||
ts: 3100,
|
||||
@@ -531,7 +481,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_incomplete_invoke_is_interrupted() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -544,14 +494,15 @@ mod tests {
|
||||
ts: 2000,
|
||||
trigger: InvokeKind::UserSend,
|
||||
},
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 2001,
|
||||
extensions: vec![],
|
||||
segments: vec![Segment::text("run a tool")],
|
||||
history: vec![annotated(Item::user_message("run a tool"))],
|
||||
},
|
||||
LogEntry::AssistantItem {
|
||||
LogEntry::AnnotatedAssistantItem {
|
||||
ts: 3000,
|
||||
item: Item::tool_call("call_1", "side_effect", "{}").into(),
|
||||
entry: annotated(Item::tool_call("call_1", "side_effect", "{}")),
|
||||
},
|
||||
]);
|
||||
|
||||
@@ -561,7 +512,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_with_tool_calls() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -570,22 +521,27 @@ mod tests {
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
},
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 2000,
|
||||
extensions: vec![],
|
||||
segments: vec![Segment::text("Check weather")],
|
||||
history: vec![annotated(Item::user_message("Check weather"))],
|
||||
},
|
||||
LogEntry::AssistantItem {
|
||||
LogEntry::AnnotatedAssistantItem {
|
||||
ts: 3000,
|
||||
item: Item::tool_call("call_1", "get_weather", r#"{"city":"Tokyo"}"#).into(),
|
||||
entry: annotated(Item::tool_call(
|
||||
"call_1",
|
||||
"get_weather",
|
||||
r#"{"city":"Tokyo"}"#,
|
||||
)),
|
||||
},
|
||||
LogEntry::ToolResult {
|
||||
LogEntry::AnnotatedToolResult {
|
||||
ts: 3500,
|
||||
item: Item::tool_result("call_1", "Sunny, 25C").into(),
|
||||
entry: annotated(Item::tool_result("call_1", "Sunny, 25C")),
|
||||
},
|
||||
LogEntry::AssistantItem {
|
||||
LogEntry::AnnotatedAssistantItem {
|
||||
ts: 4000,
|
||||
item: Item::assistant_message("It's sunny in Tokyo!").into(),
|
||||
entry: annotated(Item::assistant_message("It's sunny in Tokyo!")),
|
||||
},
|
||||
LogEntry::TurnEnd {
|
||||
ts: 4100,
|
||||
@@ -599,9 +555,9 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn replay_restores_durable_tool_image_detail() {
|
||||
let entry = LogEntry::ToolResult {
|
||||
let entry = LogEntry::AnnotatedToolResult {
|
||||
ts: 3500,
|
||||
item: Item::tool_result_item_with_attachments(
|
||||
entry: annotated(Item::tool_result_item_with_attachments(
|
||||
"call_image",
|
||||
"attached",
|
||||
None,
|
||||
@@ -609,8 +565,7 @@ mod tests {
|
||||
vec![agen::tool::Attachment::Image(
|
||||
agen::tool::ImageAttachment::new("image/png", b"durable-image".to_vec()),
|
||||
)],
|
||||
)
|
||||
.into(),
|
||||
)),
|
||||
};
|
||||
let persisted = serde_json::to_string(&entry).unwrap();
|
||||
let restored_entry: LogEntry = serde_json::from_str(&persisted).unwrap();
|
||||
@@ -630,7 +585,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_config_changed() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -650,7 +605,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_llm_usage_appends_to_usage_history() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -659,10 +614,11 @@ mod tests {
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
},
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 2000,
|
||||
extensions: vec![],
|
||||
segments: vec![Segment::text("hi")],
|
||||
history: vec![annotated(Item::user_message("hi"))],
|
||||
},
|
||||
LogEntry::LlmUsage {
|
||||
ts: 2100,
|
||||
@@ -672,9 +628,9 @@ mod tests {
|
||||
cache_write_tokens: 0,
|
||||
output_tokens: 10,
|
||||
},
|
||||
LogEntry::AssistantItem {
|
||||
LogEntry::AnnotatedAssistantItem {
|
||||
ts: 2200,
|
||||
item: Item::assistant_message("yo").into(),
|
||||
entry: annotated(Item::assistant_message("yo")),
|
||||
},
|
||||
LogEntry::LlmUsage {
|
||||
ts: 3100,
|
||||
@@ -698,7 +654,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_without_llm_usage_keeps_usage_history_empty() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -707,10 +663,11 @@ mod tests {
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
},
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 2000,
|
||||
extensions: vec![],
|
||||
segments: vec![Segment::text("hi")],
|
||||
history: vec![annotated(Item::user_message("hi"))],
|
||||
},
|
||||
]);
|
||||
assert!(state.usage_history.is_empty());
|
||||
@@ -771,7 +728,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_invoke_marker_only_mutates_interrupted_state() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 0,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -784,10 +741,11 @@ mod tests {
|
||||
ts: 100,
|
||||
trigger: InvokeKind::UserSend,
|
||||
},
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 101,
|
||||
extensions: vec![],
|
||||
segments: vec![Segment::text("hi")],
|
||||
history: vec![annotated(Item::user_message("hi"))],
|
||||
},
|
||||
LogEntry::TurnEnd {
|
||||
ts: 200,
|
||||
@@ -806,7 +764,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_paused_turn_abandoned_clears_interrupted_marker() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 0,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -830,7 +788,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_restores_active_run_budget_across_compaction_checkpoint() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 0,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -861,7 +819,7 @@ mod tests {
|
||||
}))
|
||||
.expect("legacy run-completed entry");
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 0,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -924,7 +882,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_extension_collects_domain_payload_pairs() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -983,9 +941,12 @@ mod tests {
|
||||
#[test]
|
||||
fn user_input_extensions_restore_with_the_same_committed_input() {
|
||||
let segments = vec![Segment::text("Flow instructions"), Segment::text("Ticket")];
|
||||
let entry = LogEntry::UserInput {
|
||||
let entry = LogEntry::AnnotatedUserInput {
|
||||
ts: 9999,
|
||||
segments: segments.clone(),
|
||||
history: vec![annotated(Item::user_message(Segment::flatten_to_text(
|
||||
&segments,
|
||||
)))],
|
||||
extensions: vec![SessionExtension::new(
|
||||
"flow.runtime.v1",
|
||||
serde_json::json!({ "state": "implement", "revision": 0 }),
|
||||
@@ -1000,7 +961,7 @@ mod tests {
|
||||
assert_eq!(state.extensions[0].1["state"], "implement");
|
||||
}
|
||||
|
||||
/// Mixed segments survive a JSON round-trip through `LogEntry::UserInput`,
|
||||
/// Mixed segments survive a JSON round-trip through `LogEntry::AnnotatedUserInput`,
|
||||
/// and `collect_state` derives `Item::user_message` from the flattened
|
||||
/// text while preserving the original segments separately. This covers
|
||||
/// the segments → flatten → Item replay path from the ticket.
|
||||
@@ -1020,16 +981,19 @@ mod tests {
|
||||
path: "src/main.rs".into(),
|
||||
},
|
||||
];
|
||||
let entry = LogEntry::UserInput {
|
||||
let entry = LogEntry::AnnotatedUserInput {
|
||||
ts: 4242,
|
||||
extensions: vec![],
|
||||
segments: segments.clone(),
|
||||
history: vec![annotated(Item::user_message(Segment::flatten_to_text(
|
||||
&segments,
|
||||
)))],
|
||||
};
|
||||
// JSON round-trip preserves the variant byte-for-byte.
|
||||
let json = serde_json::to_string(&entry).unwrap();
|
||||
let parsed: LogEntry = serde_json::from_str(&json).unwrap();
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
|
||||
@@ -13,7 +13,10 @@
|
||||
|
||||
use crate::event_trace::TraceEntry;
|
||||
use crate::segment_log::LogEntry;
|
||||
use crate::{SegmentId, SessionId};
|
||||
use crate::{
|
||||
PasteArtifactLimits, SegmentId, SessionId, UploadedFileLimits, UploadedFileUploadContext,
|
||||
};
|
||||
use protocol::{PasteArtifactRef, UploadedFileRef};
|
||||
|
||||
/// Errors from the persistence store.
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
@@ -29,6 +32,42 @@ pub enum StoreError {
|
||||
|
||||
#[error("log corrupted at line {line}: {message}")]
|
||||
Corrupt { line: usize, message: String },
|
||||
|
||||
#[error("paste artifact storage is unavailable")]
|
||||
PasteArtifactUnsupported,
|
||||
|
||||
#[error("paste artifact not found: {0}")]
|
||||
PasteArtifactNotFound(String),
|
||||
|
||||
#[error("paste artifact integrity check failed: {0}")]
|
||||
PasteArtifactIntegrity(String),
|
||||
|
||||
#[error("paste artifact size limit exceeded: {0}")]
|
||||
PasteArtifactLimit(String),
|
||||
|
||||
#[error("uploaded file is too large")]
|
||||
ArtifactTooLarge,
|
||||
|
||||
#[error("session artifact aggregate quota exceeded")]
|
||||
ArtifactQuotaExceeded,
|
||||
|
||||
#[error("uploaded file reference integrity check failed")]
|
||||
ArtifactIntegrityMismatch,
|
||||
|
||||
#[error("uploaded file name is invalid")]
|
||||
InvalidUploadedFileName,
|
||||
|
||||
#[error("uploaded file media type is invalid")]
|
||||
InvalidUploadedFileMediaType,
|
||||
|
||||
#[error("uploaded file is already committed to session history")]
|
||||
ArtifactAlreadyCommitted,
|
||||
|
||||
#[error("artifact id is invalid")]
|
||||
InvalidArtifactId,
|
||||
|
||||
#[error("artifact timestamp is invalid")]
|
||||
InvalidTimestamp,
|
||||
}
|
||||
|
||||
/// Sync persistence backend for segment logs.
|
||||
@@ -117,6 +156,97 @@ pub trait Store: Send + Sync {
|
||||
segment_id: SegmentId,
|
||||
) -> Result<usize, StoreError>;
|
||||
|
||||
/// Store a large paste before its reference is committed to history.
|
||||
fn write_paste_artifact(
|
||||
&self,
|
||||
_session_id: SessionId,
|
||||
_source_entry_id: &str,
|
||||
_content: &str,
|
||||
_limits: PasteArtifactLimits,
|
||||
) -> Result<PasteArtifactRef, StoreError> {
|
||||
Err(StoreError::PasteArtifactUnsupported)
|
||||
}
|
||||
|
||||
/// Read and verify one artifact owned by `session_id`.
|
||||
fn read_paste_artifact(
|
||||
&self,
|
||||
_session_id: SessionId,
|
||||
_artifact_id: &str,
|
||||
) -> Result<(PasteArtifactRef, String), StoreError> {
|
||||
Err(StoreError::PasteArtifactUnsupported)
|
||||
}
|
||||
|
||||
/// Persist a client-local file before a submission references it.
|
||||
fn write_uploaded_file(
|
||||
&self,
|
||||
_session_id: SessionId,
|
||||
_file_name: &str,
|
||||
_media_type: &str,
|
||||
_content: &[u8],
|
||||
_limits: UploadedFileLimits,
|
||||
) -> Result<UploadedFileRef, StoreError> {
|
||||
Err(StoreError::PasteArtifactUnsupported)
|
||||
}
|
||||
|
||||
fn write_uploaded_file_with_context(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
file_name: &str,
|
||||
media_type: &str,
|
||||
content: &[u8],
|
||||
_context: &UploadedFileUploadContext,
|
||||
limits: UploadedFileLimits,
|
||||
) -> Result<UploadedFileRef, StoreError> {
|
||||
self.write_uploaded_file(session_id, file_name, media_type, content, limits)
|
||||
}
|
||||
|
||||
/// Read and integrity-check an uploaded file owned by `session_id`.
|
||||
fn read_uploaded_file(
|
||||
&self,
|
||||
_session_id: SessionId,
|
||||
_reference: &UploadedFileRef,
|
||||
) -> Result<Vec<u8>, StoreError> {
|
||||
Err(StoreError::PasteArtifactUnsupported)
|
||||
}
|
||||
|
||||
fn read_uploaded_file_by_id(
|
||||
&self,
|
||||
_session_id: SessionId,
|
||||
_artifact_id: &str,
|
||||
) -> Result<(UploadedFileRef, Vec<u8>), StoreError> {
|
||||
Err(StoreError::PasteArtifactUnsupported)
|
||||
}
|
||||
|
||||
fn bind_uploaded_file(
|
||||
&self,
|
||||
_session_id: SessionId,
|
||||
_reference: &UploadedFileRef,
|
||||
_source_entry_id: &str,
|
||||
) -> Result<UploadedFileRef, StoreError> {
|
||||
Err(StoreError::PasteArtifactUnsupported)
|
||||
}
|
||||
|
||||
/// Delete an uncommitted uploaded file owned by `session_id`.
|
||||
fn delete_uploaded_file(
|
||||
&self,
|
||||
_session_id: SessionId,
|
||||
_artifact_id: &str,
|
||||
) -> Result<bool, StoreError> {
|
||||
Err(StoreError::PasteArtifactUnsupported)
|
||||
}
|
||||
|
||||
fn delete_uncommitted_uploaded_files(&self, _session_id: SessionId) -> Result<u64, StoreError> {
|
||||
Ok(0)
|
||||
}
|
||||
|
||||
fn copy_committed_uploaded_files(
|
||||
&self,
|
||||
_source_session_id: SessionId,
|
||||
_target_session_id: SessionId,
|
||||
) -> Result<u64, StoreError> {
|
||||
Ok(0)
|
||||
}
|
||||
|
||||
/// Append a trace entry to the debug event trace file.
|
||||
fn append_trace(
|
||||
&self,
|
||||
|
||||
@@ -8,7 +8,7 @@
|
||||
//! `kind` instead of parsing text prefixes like `[Notification] …` or
|
||||
//! `[File: …]`.
|
||||
//!
|
||||
//! Persisted as the payload of [`crate::LogEntry::SystemItem`] (one
|
||||
//! Persisted as the payload of [`crate::LogEntry::AnnotatedSystemItem`] (one
|
||||
//! entry per item), and broadcast live as the payload of
|
||||
//! `Event::SystemItem` on the wire.
|
||||
//!
|
||||
|
||||
@@ -0,0 +1,534 @@
|
||||
use std::{
|
||||
fs,
|
||||
path::Path,
|
||||
time::{SystemTime, UNIX_EPOCH},
|
||||
};
|
||||
|
||||
use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64};
|
||||
use fs4::fs_std::FileExt;
|
||||
use protocol::{UploadedFileAvailability, UploadedFileRef};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use sha2::{Digest, Sha256};
|
||||
use unicode_normalization::UnicodeNormalization;
|
||||
use unicode_properties::general_category::{GeneralCategory, UnicodeGeneralCategory};
|
||||
use unicode_security::{confusable_detection::skeleton, mixed_script::MixedScript};
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::StoreError;
|
||||
|
||||
type Result<T> = std::result::Result<T, StoreError>;
|
||||
|
||||
pub const DEFAULT_MAX_UPLOADED_FILE_BYTES: u64 = 10 * 1024 * 1024;
|
||||
pub const DEFAULT_MAX_SESSION_ARTIFACT_BYTES: u64 = 32 * 1024 * 1024;
|
||||
pub const DEFAULT_MAX_FILES_PER_SUBMISSION: usize = 8;
|
||||
pub const DEFAULT_MAX_SESSION_UPLOADED_FILES: u64 = 256;
|
||||
const MAX_FILE_NAME_CHARS: usize = 255;
|
||||
const MAX_MEDIA_TYPE_BYTES: usize = 127;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct UploadedFileLimits {
|
||||
pub max_file_bytes: u64,
|
||||
pub max_session_bytes: u64,
|
||||
}
|
||||
|
||||
impl Default for UploadedFileLimits {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
max_file_bytes: DEFAULT_MAX_UPLOADED_FILE_BYTES,
|
||||
max_session_bytes: DEFAULT_MAX_SESSION_ARTIFACT_BYTES,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct UploadedFileUploadContext {
|
||||
pub upload_id: String,
|
||||
pub principal_id: String,
|
||||
pub workspace_id: String,
|
||||
pub runtime_id: String,
|
||||
pub worker_id: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize)]
|
||||
struct StoredUploadedFile {
|
||||
file_name: String,
|
||||
media_type: String,
|
||||
created_at_ms: u64,
|
||||
byte_len: u64,
|
||||
sha256: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
source_entry_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
upload_context: Option<UploadedFileUploadContext>,
|
||||
content_base64: String,
|
||||
}
|
||||
|
||||
pub(crate) fn validate_file_name(file_name: &str) -> Result<()> {
|
||||
let normalized: String = file_name.nfkc().collect();
|
||||
let has_unsafe_component = file_name
|
||||
.split('.')
|
||||
.filter(|part| !part.is_empty())
|
||||
.any(|part| {
|
||||
let confusable_skeleton: String = skeleton(part).collect();
|
||||
let ascii_confusable = part.chars().any(|ch| !ch.is_ascii())
|
||||
&& confusable_skeleton.is_ascii()
|
||||
&& !confusable_skeleton.eq_ignore_ascii_case(part);
|
||||
!part.is_single_script() || ascii_confusable
|
||||
});
|
||||
|
||||
if file_name.is_empty()
|
||||
|| file_name.chars().count() > MAX_FILE_NAME_CHARS
|
||||
|| file_name == "."
|
||||
|| file_name == ".."
|
||||
|| normalized != file_name
|
||||
|| has_unsafe_component
|
||||
|| file_name.chars().any(|ch| {
|
||||
ch.is_control()
|
||||
|| ch.general_category() == GeneralCategory::Format
|
||||
|| matches!(ch, '/' | '\\')
|
||||
})
|
||||
{
|
||||
return Err(StoreError::InvalidUploadedFileName);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) fn validate_media_type(media_type: &str) -> Result<()> {
|
||||
let valid = !media_type.is_empty()
|
||||
&& media_type.len() <= MAX_MEDIA_TYPE_BYTES
|
||||
&& media_type.is_ascii()
|
||||
&& !media_type
|
||||
.bytes()
|
||||
.any(|byte| byte.is_ascii_control() || byte == b' ')
|
||||
&& media_type.split_once('/').is_some_and(|(kind, subtype)| {
|
||||
!kind.is_empty()
|
||||
&& !subtype.is_empty()
|
||||
&& kind.bytes().chain(subtype.bytes()).all(|byte| {
|
||||
byte.is_ascii_alphanumeric()
|
||||
|| matches!(
|
||||
byte,
|
||||
b'!' | b'#' | b'$' | b'&' | b'^' | b'_' | b'.' | b'+' | b'-'
|
||||
)
|
||||
})
|
||||
});
|
||||
let allowed = media_type.starts_with("text/")
|
||||
|| matches!(
|
||||
media_type,
|
||||
"application/json"
|
||||
| "application/pdf"
|
||||
| "image/png"
|
||||
| "image/jpeg"
|
||||
| "image/gif"
|
||||
| "image/webp"
|
||||
);
|
||||
if !valid || !allowed {
|
||||
return Err(StoreError::InvalidUploadedFileMediaType);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn normalized_file_name(file_name: &str) -> String {
|
||||
file_name.nfkc().flat_map(char::to_lowercase).collect()
|
||||
}
|
||||
|
||||
fn validate_content(media_type: &str, content: &[u8]) -> Result<()> {
|
||||
if content.is_empty() {
|
||||
return Err(StoreError::InvalidUploadedFileMediaType);
|
||||
}
|
||||
let matches_declared_type = if media_type.starts_with("text/") {
|
||||
std::str::from_utf8(content).is_ok()
|
||||
} else {
|
||||
match media_type {
|
||||
"application/json" => serde_json::from_slice::<serde_json::Value>(content).is_ok(),
|
||||
"application/pdf" => content.starts_with(b"%PDF-"),
|
||||
"image/png" => content.starts_with(b"\x89PNG\r\n\x1a\n"),
|
||||
"image/jpeg" => content.starts_with(&[0xff, 0xd8, 0xff]),
|
||||
"image/gif" => content.starts_with(b"GIF87a") || content.starts_with(b"GIF89a"),
|
||||
"image/webp" => {
|
||||
content.len() >= 12 && content.starts_with(b"RIFF") && &content[8..12] == b"WEBP"
|
||||
}
|
||||
_ => false,
|
||||
}
|
||||
};
|
||||
if !matches_declared_type {
|
||||
return Err(StoreError::ArtifactIntegrityMismatch);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn record_path(dir: &Path, artifact_id: &str) -> Result<std::path::PathBuf> {
|
||||
let id = Uuid::parse_str(artifact_id).map_err(|_| StoreError::InvalidArtifactId)?;
|
||||
Ok(dir.join(format!("{id}.file.json")))
|
||||
}
|
||||
|
||||
fn now_ms() -> Result<u64> {
|
||||
let value = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map_err(|_| StoreError::InvalidTimestamp)?
|
||||
.as_millis();
|
||||
u64::try_from(value).map_err(|_| StoreError::InvalidTimestamp)
|
||||
}
|
||||
|
||||
fn digest(bytes: &[u8]) -> String {
|
||||
Sha256::digest(bytes)
|
||||
.iter()
|
||||
.map(|byte| format!("{byte:02x}"))
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(crate) fn stored_uploaded_file_usage(dir: &Path) -> Result<(u64, u64)> {
|
||||
if !dir.exists() {
|
||||
return Ok((0, 0));
|
||||
}
|
||||
let mut bytes = 0_u64;
|
||||
let mut count = 0_u64;
|
||||
for entry in fs::read_dir(dir)? {
|
||||
let entry = entry?;
|
||||
let path = entry.path();
|
||||
if !entry.file_type()?.is_file()
|
||||
|| !path
|
||||
.file_name()
|
||||
.and_then(|name| name.to_str())
|
||||
.is_some_and(|name| name.ends_with(".file.json"))
|
||||
{
|
||||
continue;
|
||||
}
|
||||
let stored: StoredUploadedFile = serde_json::from_slice(&fs::read(&path)?)?;
|
||||
bytes = bytes
|
||||
.checked_add(stored.byte_len)
|
||||
.ok_or(StoreError::ArtifactQuotaExceeded)?;
|
||||
count = count
|
||||
.checked_add(1)
|
||||
.ok_or(StoreError::ArtifactQuotaExceeded)?;
|
||||
}
|
||||
Ok((bytes, count))
|
||||
}
|
||||
|
||||
pub(crate) fn write_uploaded_file(
|
||||
dir: &Path,
|
||||
file_name: &str,
|
||||
media_type: &str,
|
||||
content: &[u8],
|
||||
context: Option<&UploadedFileUploadContext>,
|
||||
limits: UploadedFileLimits,
|
||||
) -> Result<UploadedFileRef> {
|
||||
validate_file_name(file_name)?;
|
||||
validate_media_type(media_type)?;
|
||||
validate_content(media_type, content)?;
|
||||
let byte_len = u64::try_from(content.len()).map_err(|_| StoreError::ArtifactTooLarge)?;
|
||||
let sha256 = digest(content);
|
||||
if byte_len > limits.max_file_bytes {
|
||||
return Err(StoreError::ArtifactTooLarge);
|
||||
}
|
||||
|
||||
fs::create_dir_all(dir)?;
|
||||
let aggregate_lock = fs::OpenOptions::new()
|
||||
.create(true)
|
||||
.read(true)
|
||||
.write(true)
|
||||
.open(dir.join(".aggregate.lock"))?;
|
||||
FileExt::lock_exclusive(&aggregate_lock)?;
|
||||
let (paste_bytes, _) = crate::paste_artifact::stored_paste_usage(dir)?;
|
||||
let (file_bytes, file_count) = stored_uploaded_file_usage(dir)?;
|
||||
let normalized_name = normalized_file_name(file_name);
|
||||
for entry in fs::read_dir(dir)? {
|
||||
let path = entry?.path();
|
||||
if !path
|
||||
.file_name()
|
||||
.and_then(|name| name.to_str())
|
||||
.is_some_and(|name| name.ends_with(".file.json"))
|
||||
{
|
||||
continue;
|
||||
}
|
||||
let stored: StoredUploadedFile = serde_json::from_slice(&fs::read(&path)?)?;
|
||||
let same_context = context.is_some() && stored.upload_context.as_ref() == context;
|
||||
let same_uncommitted_name = stored.source_entry_id.is_none()
|
||||
&& normalized_file_name(&stored.file_name) == normalized_name;
|
||||
if same_context || same_uncommitted_name {
|
||||
if stored.file_name == file_name
|
||||
&& stored.media_type == media_type
|
||||
&& stored.byte_len == byte_len
|
||||
&& stored.sha256 == sha256
|
||||
&& stored.upload_context.as_ref() == context
|
||||
{
|
||||
let artifact_id = path
|
||||
.file_name()
|
||||
.and_then(|name| name.to_str())
|
||||
.and_then(|name| name.strip_suffix(".file.json"))
|
||||
.ok_or(StoreError::InvalidArtifactId)?
|
||||
.to_string();
|
||||
return Ok(UploadedFileRef {
|
||||
artifact_id,
|
||||
file_name: stored.file_name,
|
||||
media_type: stored.media_type,
|
||||
created_at_ms: stored.created_at_ms,
|
||||
availability: UploadedFileAvailability::Available,
|
||||
byte_len: stored.byte_len,
|
||||
sha256: stored.sha256,
|
||||
source_entry_id: None,
|
||||
});
|
||||
}
|
||||
return Err(StoreError::InvalidUploadedFileName);
|
||||
}
|
||||
}
|
||||
if file_count >= DEFAULT_MAX_SESSION_UPLOADED_FILES {
|
||||
return Err(StoreError::ArtifactQuotaExceeded);
|
||||
}
|
||||
if paste_bytes
|
||||
.checked_add(file_bytes)
|
||||
.and_then(|total| total.checked_add(byte_len))
|
||||
.is_none_or(|total| total > limits.max_session_bytes)
|
||||
{
|
||||
return Err(StoreError::ArtifactQuotaExceeded);
|
||||
}
|
||||
|
||||
let artifact_id = Uuid::now_v7().to_string();
|
||||
let created_at_ms = now_ms()?;
|
||||
let stored = StoredUploadedFile {
|
||||
file_name: file_name.to_owned(),
|
||||
media_type: media_type.to_owned(),
|
||||
created_at_ms,
|
||||
byte_len,
|
||||
sha256: sha256.clone(),
|
||||
source_entry_id: None,
|
||||
upload_context: context.cloned(),
|
||||
content_base64: BASE64.encode(content),
|
||||
};
|
||||
let path = record_path(dir, &artifact_id)?;
|
||||
let temp = dir.join(format!(".{artifact_id}.file.tmp"));
|
||||
fs::write(&temp, serde_json::to_vec(&stored)?)?;
|
||||
fs::rename(&temp, &path)?;
|
||||
|
||||
Ok(UploadedFileRef {
|
||||
artifact_id,
|
||||
file_name: file_name.to_owned(),
|
||||
media_type: media_type.to_owned(),
|
||||
created_at_ms,
|
||||
availability: UploadedFileAvailability::Available,
|
||||
byte_len,
|
||||
sha256,
|
||||
source_entry_id: None,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn read_uploaded_file_by_id(
|
||||
dir: &Path,
|
||||
artifact_id: &str,
|
||||
) -> Result<(UploadedFileRef, Vec<u8>)> {
|
||||
let stored: StoredUploadedFile =
|
||||
serde_json::from_slice(&fs::read(record_path(dir, artifact_id)?)?)?;
|
||||
let content = BASE64
|
||||
.decode(&stored.content_base64)
|
||||
.map_err(|_| StoreError::ArtifactIntegrityMismatch)?;
|
||||
if u64::try_from(content.len()).ok() != Some(stored.byte_len)
|
||||
|| digest(&content) != stored.sha256
|
||||
{
|
||||
return Err(StoreError::ArtifactIntegrityMismatch);
|
||||
}
|
||||
let reference = UploadedFileRef {
|
||||
artifact_id: artifact_id.to_owned(),
|
||||
file_name: stored.file_name,
|
||||
media_type: stored.media_type,
|
||||
created_at_ms: stored.created_at_ms,
|
||||
availability: UploadedFileAvailability::Available,
|
||||
byte_len: stored.byte_len,
|
||||
sha256: stored.sha256,
|
||||
source_entry_id: stored.source_entry_id,
|
||||
};
|
||||
Ok((reference, content))
|
||||
}
|
||||
|
||||
pub(crate) fn read_uploaded_file(dir: &Path, reference: &UploadedFileRef) -> Result<Vec<u8>> {
|
||||
let (stored_reference, content) = read_uploaded_file_by_id(dir, &reference.artifact_id)?;
|
||||
if stored_reference.file_name != reference.file_name
|
||||
|| stored_reference.media_type != reference.media_type
|
||||
|| stored_reference.created_at_ms != reference.created_at_ms
|
||||
|| stored_reference.byte_len != reference.byte_len
|
||||
|| stored_reference.sha256 != reference.sha256
|
||||
|| stored_reference.source_entry_id != reference.source_entry_id
|
||||
{
|
||||
return Err(StoreError::ArtifactIntegrityMismatch);
|
||||
}
|
||||
Ok(content)
|
||||
}
|
||||
|
||||
pub(crate) fn clear_uploaded_file_binding(
|
||||
dir: &Path,
|
||||
artifact_id: &str,
|
||||
expected_source_entry_id: &str,
|
||||
) -> Result<()> {
|
||||
fs::create_dir_all(dir)?;
|
||||
let aggregate_lock = fs::OpenOptions::new()
|
||||
.create(true)
|
||||
.read(true)
|
||||
.write(true)
|
||||
.open(dir.join(".aggregate.lock"))?;
|
||||
FileExt::lock_exclusive(&aggregate_lock)?;
|
||||
let path = record_path(dir, artifact_id)?;
|
||||
let mut stored: StoredUploadedFile = serde_json::from_slice(&fs::read(&path)?)?;
|
||||
if stored.source_entry_id.as_deref() != Some(expected_source_entry_id) {
|
||||
return Err(StoreError::ArtifactIntegrityMismatch);
|
||||
}
|
||||
stored.source_entry_id = None;
|
||||
let temp = dir.join(format!(".{artifact_id}.file.unbind.tmp"));
|
||||
fs::write(&temp, serde_json::to_vec(&stored)?)?;
|
||||
fs::rename(temp, path)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) fn bind_uploaded_file(
|
||||
dir: &Path,
|
||||
reference: &UploadedFileRef,
|
||||
source_entry_id: &str,
|
||||
) -> Result<UploadedFileRef> {
|
||||
if source_entry_id.is_empty() || reference.source_entry_id.is_some() {
|
||||
return Err(StoreError::ArtifactIntegrityMismatch);
|
||||
}
|
||||
let aggregate_lock = fs::OpenOptions::new()
|
||||
.create(true)
|
||||
.read(true)
|
||||
.write(true)
|
||||
.open(dir.join(".aggregate.lock"))?;
|
||||
FileExt::lock_exclusive(&aggregate_lock)?;
|
||||
let (stored_reference, _) = read_uploaded_file_by_id(dir, &reference.artifact_id)?;
|
||||
if stored_reference.file_name != reference.file_name
|
||||
|| stored_reference.media_type != reference.media_type
|
||||
|| stored_reference.created_at_ms != reference.created_at_ms
|
||||
|| stored_reference.byte_len != reference.byte_len
|
||||
|| stored_reference.sha256 != reference.sha256
|
||||
{
|
||||
return Err(StoreError::ArtifactIntegrityMismatch);
|
||||
}
|
||||
let path = record_path(dir, &reference.artifact_id)?;
|
||||
let mut stored: StoredUploadedFile = serde_json::from_slice(&fs::read(&path)?)?;
|
||||
if stored.source_entry_id.is_some() {
|
||||
return Err(StoreError::ArtifactAlreadyCommitted);
|
||||
}
|
||||
stored.source_entry_id = Some(source_entry_id.to_owned());
|
||||
let temp = dir.join(format!(".{}.file.bind.tmp", reference.artifact_id));
|
||||
fs::write(&temp, serde_json::to_vec(&stored)?)?;
|
||||
fs::rename(&temp, path)?;
|
||||
let mut bound = reference.clone();
|
||||
bound.source_entry_id = Some(source_entry_id.to_owned());
|
||||
Ok(bound)
|
||||
}
|
||||
|
||||
pub(crate) fn list_uploaded_file_refs(dir: &Path) -> Result<Vec<UploadedFileRef>> {
|
||||
if !dir.exists() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
let mut refs = Vec::new();
|
||||
for entry in fs::read_dir(dir)? {
|
||||
let path = entry?.path();
|
||||
let Some(artifact_id) = path
|
||||
.file_name()
|
||||
.and_then(|name| name.to_str())
|
||||
.and_then(|name| name.strip_suffix(".file.json"))
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
refs.push(read_uploaded_file_by_id(dir, artifact_id)?.0);
|
||||
}
|
||||
Ok(refs)
|
||||
}
|
||||
|
||||
pub(crate) fn copy_committed_uploaded_files(source_dir: &Path, target_dir: &Path) -> Result<u64> {
|
||||
if !source_dir.exists() {
|
||||
return Ok(0);
|
||||
}
|
||||
fs::create_dir_all(target_dir)?;
|
||||
let target_lock = fs::OpenOptions::new()
|
||||
.create(true)
|
||||
.read(true)
|
||||
.write(true)
|
||||
.open(target_dir.join(".aggregate.lock"))?;
|
||||
FileExt::lock_exclusive(&target_lock)?;
|
||||
let mut copied = 0_u64;
|
||||
for entry in fs::read_dir(source_dir)? {
|
||||
let entry = entry?;
|
||||
let path = entry.path();
|
||||
let Some(name) = path.file_name().and_then(|name| name.to_str()) else {
|
||||
continue;
|
||||
};
|
||||
if !name.ends_with(".file.json") {
|
||||
continue;
|
||||
}
|
||||
let bytes = fs::read(&path)?;
|
||||
let stored: StoredUploadedFile = serde_json::from_slice(&bytes)?;
|
||||
if stored.source_entry_id.is_none() {
|
||||
continue;
|
||||
}
|
||||
let target = target_dir.join(name);
|
||||
if target.exists() {
|
||||
let existing: StoredUploadedFile = serde_json::from_slice(&fs::read(&target)?)?;
|
||||
if existing.sha256 != stored.sha256
|
||||
|| existing.file_name != stored.file_name
|
||||
|| existing.source_entry_id != stored.source_entry_id
|
||||
{
|
||||
return Err(StoreError::ArtifactIntegrityMismatch);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
let temp = target_dir.join(format!(".{name}.copy.tmp"));
|
||||
fs::write(&temp, &bytes)?;
|
||||
fs::rename(temp, target)?;
|
||||
copied = copied
|
||||
.checked_add(1)
|
||||
.ok_or(StoreError::ArtifactQuotaExceeded)?;
|
||||
}
|
||||
Ok(copied)
|
||||
}
|
||||
|
||||
pub(crate) fn delete_uncommitted_uploaded_files(dir: &Path) -> Result<u64> {
|
||||
fs::create_dir_all(dir)?;
|
||||
let aggregate_lock = fs::OpenOptions::new()
|
||||
.create(true)
|
||||
.read(true)
|
||||
.write(true)
|
||||
.open(dir.join(".aggregate.lock"))?;
|
||||
FileExt::lock_exclusive(&aggregate_lock)?;
|
||||
let mut removed = 0_u64;
|
||||
for entry in fs::read_dir(dir)? {
|
||||
let entry = entry?;
|
||||
let path = entry.path();
|
||||
if !path
|
||||
.file_name()
|
||||
.and_then(|name| name.to_str())
|
||||
.is_some_and(|name| name.ends_with(".file.json"))
|
||||
{
|
||||
continue;
|
||||
}
|
||||
let stored: StoredUploadedFile = serde_json::from_slice(&fs::read(&path)?)?;
|
||||
if stored.source_entry_id.is_none() {
|
||||
fs::remove_file(path)?;
|
||||
removed = removed
|
||||
.checked_add(1)
|
||||
.ok_or(StoreError::ArtifactQuotaExceeded)?;
|
||||
}
|
||||
}
|
||||
Ok(removed)
|
||||
}
|
||||
|
||||
pub(crate) fn delete_uploaded_file(dir: &Path, artifact_id: &str) -> Result<bool> {
|
||||
fs::create_dir_all(dir)?;
|
||||
let aggregate_lock = fs::OpenOptions::new()
|
||||
.create(true)
|
||||
.read(true)
|
||||
.write(true)
|
||||
.open(dir.join(".aggregate.lock"))?;
|
||||
FileExt::lock_exclusive(&aggregate_lock)?;
|
||||
let path = record_path(dir, artifact_id)?;
|
||||
let stored = match fs::read(&path) {
|
||||
Ok(bytes) => serde_json::from_slice::<StoredUploadedFile>(&bytes)?,
|
||||
Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(false),
|
||||
Err(error) => return Err(error.into()),
|
||||
};
|
||||
if stored.source_entry_id.is_some() {
|
||||
return Err(StoreError::ArtifactAlreadyCommitted);
|
||||
}
|
||||
match fs::remove_file(path) {
|
||||
Ok(()) => Ok(true),
|
||||
Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(false),
|
||||
Err(error) => Err(error.into()),
|
||||
}
|
||||
}
|
||||
@@ -608,6 +608,24 @@ where
|
||||
) -> Result<usize, crate::StoreError> {
|
||||
self.session_store.read_entry_count(session_id, segment_id)
|
||||
}
|
||||
fn write_paste_artifact(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
source_entry_id: &str,
|
||||
content: &str,
|
||||
limits: crate::PasteArtifactLimits,
|
||||
) -> Result<protocol::PasteArtifactRef, crate::StoreError> {
|
||||
self.session_store
|
||||
.write_paste_artifact(session_id, source_entry_id, content, limits)
|
||||
}
|
||||
fn read_paste_artifact(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
artifact_id: &str,
|
||||
) -> Result<(protocol::PasteArtifactRef, String), crate::StoreError> {
|
||||
self.session_store
|
||||
.read_paste_artifact(session_id, artifact_id)
|
||||
}
|
||||
fn append_trace(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
|
||||
@@ -10,9 +10,11 @@
|
||||
//! every later operation must use that same ID.
|
||||
|
||||
use crate::event_trace::TraceEntry;
|
||||
use crate::paste_artifact::{read_from_dir, write_to_dir};
|
||||
use crate::segment_log::LogEntry;
|
||||
use crate::store::{Store, StoreError};
|
||||
use crate::{SegmentId, SessionId};
|
||||
use crate::{PasteArtifactLimits, SegmentId, SessionId};
|
||||
use protocol::PasteArtifactRef;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::fs::{self, File, OpenOptions};
|
||||
use std::io::{Read, Seek, SeekFrom, Write};
|
||||
@@ -20,10 +22,12 @@ use std::path::{Path, PathBuf};
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::SystemTime;
|
||||
|
||||
const SESSION_SCHEMA_VERSION: u32 = 2;
|
||||
const SESSION_SCHEMA_VERSION: u32 = 3;
|
||||
const PREVIOUS_SESSION_SCHEMA_VERSION: u32 = 2;
|
||||
const LEGACY_SESSION_SCHEMA_VERSION: u32 = 1;
|
||||
const SESSION_FILE: &str = "session.json";
|
||||
const SEGMENTS_DIR: &str = "segments";
|
||||
const PASTE_ARTIFACTS_DIR: &str = "artifacts/paste";
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct WorkerSessionStore {
|
||||
@@ -47,9 +51,15 @@ impl WorkerSessionStore {
|
||||
Ok(bytes) => {
|
||||
let mut manifest: SessionManifest = serde_json::from_slice(&bytes)?;
|
||||
match manifest.schema_version {
|
||||
SESSION_SCHEMA_VERSION => {}
|
||||
LEGACY_SESSION_SCHEMA_VERSION => {
|
||||
validate_legacy_segment_logs(&root)?;
|
||||
SESSION_SCHEMA_VERSION => {
|
||||
validate_canonical_segment_logs(&root)?;
|
||||
}
|
||||
PREVIOUS_SESSION_SCHEMA_VERSION | LEGACY_SESSION_SCHEMA_VERSION => {
|
||||
migrate_segment_logs_to_v3(
|
||||
&root,
|
||||
manifest.session_id,
|
||||
manifest.schema_version,
|
||||
)?;
|
||||
manifest.schema_version = SESSION_SCHEMA_VERSION;
|
||||
atomic_write_json(&root.join(SESSION_FILE), &manifest)?;
|
||||
}
|
||||
@@ -144,6 +154,41 @@ impl WorkerSessionStore {
|
||||
.join(format!("{segment_id}.trace.jsonl"))
|
||||
}
|
||||
|
||||
fn append_log_entry(&self, path: &Path, entry: &LogEntry) -> Result<(), StoreError> {
|
||||
let _guard = self
|
||||
.append_lock
|
||||
.lock()
|
||||
.map_err(|_| std::io::Error::other("Worker Session append lock was poisoned"))?;
|
||||
let mut file = OpenOptions::new()
|
||||
.create(true)
|
||||
.read(true)
|
||||
.write(true)
|
||||
.append(true)
|
||||
.open(path)?;
|
||||
let committed_len = truncate_uncommitted_tail(&mut file)?;
|
||||
file.seek(SeekFrom::Start(0))?;
|
||||
let mut existing = Vec::new();
|
||||
file.read_to_end(&mut existing)?;
|
||||
parse_jsonl::<LogEntry>(&existing)?;
|
||||
let line = serde_json::to_string(entry)?;
|
||||
let mut record = Vec::with_capacity(line.len() + 1);
|
||||
record.extend_from_slice(line.as_bytes());
|
||||
record.push(b'\n');
|
||||
if let Err(write_error) = file.write_all(&record) {
|
||||
return match file.set_len(committed_len) {
|
||||
Ok(()) => Err(write_error.into()),
|
||||
Err(rollback_error) => Err(std::io::Error::new(
|
||||
rollback_error.kind(),
|
||||
format!(
|
||||
"session append failed ({write_error}) and rollback failed: {rollback_error}"
|
||||
),
|
||||
)
|
||||
.into()),
|
||||
};
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn append_line(&self, path: &Path, line: &str) -> Result<(), StoreError> {
|
||||
let _guard = self
|
||||
.append_lock
|
||||
@@ -183,7 +228,7 @@ impl Store for WorkerSessionStore {
|
||||
entry: &LogEntry,
|
||||
) -> Result<(), StoreError> {
|
||||
self.ensure_session(session_id, true)?;
|
||||
self.append_line(&self.log_path(segment_id), &serde_json::to_string(entry)?)
|
||||
self.append_log_entry(&self.log_path(segment_id), entry)
|
||||
}
|
||||
|
||||
fn read_all(
|
||||
@@ -275,6 +320,35 @@ impl Store for WorkerSessionStore {
|
||||
.count())
|
||||
}
|
||||
|
||||
fn write_paste_artifact(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
source_entry_id: &str,
|
||||
content: &str,
|
||||
limits: PasteArtifactLimits,
|
||||
) -> Result<PasteArtifactRef, StoreError> {
|
||||
self.ensure_session(session_id, true)?;
|
||||
let _guard = self
|
||||
.append_lock
|
||||
.lock()
|
||||
.map_err(|_| std::io::Error::other("Worker Session append lock was poisoned"))?;
|
||||
write_to_dir(
|
||||
&self.root.join(PASTE_ARTIFACTS_DIR),
|
||||
source_entry_id,
|
||||
content,
|
||||
limits,
|
||||
)
|
||||
}
|
||||
|
||||
fn read_paste_artifact(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
artifact_id: &str,
|
||||
) -> Result<(PasteArtifactRef, String), StoreError> {
|
||||
self.ensure_session(session_id, false)?;
|
||||
read_from_dir(&self.root.join(PASTE_ARTIFACTS_DIR), artifact_id)
|
||||
}
|
||||
|
||||
fn append_trace(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
@@ -286,37 +360,138 @@ impl Store for WorkerSessionStore {
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_legacy_segment_logs(root: &Path) -> Result<(), StoreError> {
|
||||
fn segment_log_paths(root: &Path) -> Result<Vec<(SegmentId, PathBuf)>, StoreError> {
|
||||
let segments = root.join(SEGMENTS_DIR);
|
||||
if !segments.exists() {
|
||||
return Ok(());
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
let mut paths = Vec::new();
|
||||
for entry in fs::read_dir(&segments)? {
|
||||
let entry = entry?;
|
||||
let path = entry.path();
|
||||
let metadata = fs::symlink_metadata(&path)?;
|
||||
let Some(name) = path.file_name().and_then(|name| name.to_str()) else {
|
||||
continue;
|
||||
return Err(StoreError::Corrupt {
|
||||
line: 0,
|
||||
message: format!("non-UTF-8 Worker Session segment path: {}", path.display()),
|
||||
});
|
||||
};
|
||||
if !name.ends_with(".jsonl") || name.ends_with(".trace.jsonl") {
|
||||
if name.ends_with(".trace.jsonl") || name.starts_with('.') {
|
||||
continue;
|
||||
}
|
||||
let contents = fs::read_to_string(&path)?;
|
||||
for (line_index, line) in contents.lines().enumerate() {
|
||||
if line.trim().is_empty() {
|
||||
continue;
|
||||
}
|
||||
serde_json::from_str::<LogEntry>(line).map_err(|error| StoreError::Corrupt {
|
||||
line: line_index + 1,
|
||||
if !name.ends_with(".jsonl") {
|
||||
continue;
|
||||
}
|
||||
if !metadata.file_type().is_file() {
|
||||
return Err(StoreError::Corrupt {
|
||||
line: 0,
|
||||
message: format!(
|
||||
"cannot migrate legacy Worker Session log {}: {error}",
|
||||
"Worker Session segment is not a regular file: {}",
|
||||
path.display()
|
||||
),
|
||||
});
|
||||
}
|
||||
let segment_id =
|
||||
name.trim_end_matches(".jsonl")
|
||||
.parse()
|
||||
.map_err(|_| StoreError::Corrupt {
|
||||
line: 0,
|
||||
message: format!("invalid Worker Session segment name: {name}"),
|
||||
})?;
|
||||
paths.push((segment_id, path));
|
||||
}
|
||||
paths.sort_by_key(|(segment_id, _)| *segment_id);
|
||||
Ok(paths)
|
||||
}
|
||||
|
||||
fn migrate_segment_logs_to_v3(
|
||||
root: &Path,
|
||||
session_id: SessionId,
|
||||
source_schema_version: u32,
|
||||
) -> Result<(), StoreError> {
|
||||
struct MigrationPlan {
|
||||
path: PathBuf,
|
||||
source: Vec<u8>,
|
||||
output: Vec<u8>,
|
||||
}
|
||||
|
||||
// Phase 1 is strictly read-only. Every segment must parse and canonicalize
|
||||
// successfully before the first authoritative byte is replaced.
|
||||
let mut plans = Vec::new();
|
||||
for (segment_id, path) in segment_log_paths(root)? {
|
||||
let source = fs::read(&path)?;
|
||||
let canonical = parse_legacy_jsonl(source_schema_version, session_id, segment_id, &source)
|
||||
.map_err(|error| StoreError::Corrupt {
|
||||
line: 0,
|
||||
message: format!(
|
||||
"cannot migrate Worker Session log {}: {error}",
|
||||
path.display()
|
||||
),
|
||||
})?;
|
||||
let mut output = Vec::new();
|
||||
for entry in canonical {
|
||||
serde_json::to_writer(&mut output, &entry)?;
|
||||
output.push(b'\n');
|
||||
}
|
||||
plans.push(MigrationPlan {
|
||||
path,
|
||||
source,
|
||||
output,
|
||||
});
|
||||
}
|
||||
|
||||
// Fence the complete preflight snapshot before starting phase 2. Session
|
||||
// open is the exclusive restore boundary; this additionally fails closed
|
||||
// if an unexpected writer raced the preflight.
|
||||
for plan in &plans {
|
||||
if fs::read(&plan.path)? != plan.source {
|
||||
return Err(StoreError::Corrupt {
|
||||
line: 0,
|
||||
message: format!(
|
||||
"Worker Session segment changed during migration: {}",
|
||||
plan.path.display()
|
||||
),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
for plan in plans {
|
||||
atomic_write_bytes(&plan.path, &plan.output)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_canonical_segment_logs(root: &Path) -> Result<(), StoreError> {
|
||||
for (_, path) in segment_log_paths(root)? {
|
||||
let _: Vec<LogEntry> = parse_jsonl(&fs::read(&path)?)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn parse_legacy_jsonl(
|
||||
schema_version: u32,
|
||||
session_id: SessionId,
|
||||
segment_id: SegmentId,
|
||||
bytes: &[u8],
|
||||
) -> Result<Vec<LogEntry>, serde_json::Error> {
|
||||
let text = std::str::from_utf8(bytes).map_err(|error| {
|
||||
serde_json::Error::io(std::io::Error::new(std::io::ErrorKind::InvalidData, error))
|
||||
})?;
|
||||
text.lines()
|
||||
.enumerate()
|
||||
.filter(|(_, line)| !line.trim().is_empty())
|
||||
.map(|(line_index, line)| {
|
||||
crate::legacy_session_log::decode_entry(
|
||||
schema_version,
|
||||
line,
|
||||
session_id,
|
||||
segment_id,
|
||||
line_index,
|
||||
)
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn atomic_write_json<T: Serialize>(path: &Path, value: &T) -> Result<(), StoreError> {
|
||||
let mut bytes = serde_json::to_vec_pretty(value)?;
|
||||
bytes.push(b'\n');
|
||||
@@ -418,7 +593,21 @@ fn truncate_uncommitted_tail(file: &mut File) -> std::io::Result<u64> {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::{Store, new_segment_id, new_session_id};
|
||||
use crate::{
|
||||
LoggedHistoryEntry, LoggedItem, LoggedSessionHistoryEntryId, LoggedSessionHistoryMetadata,
|
||||
LoggedSessionHistoryOrigin, Store, new_segment_id, new_session_id,
|
||||
};
|
||||
|
||||
fn annotated(item: agen::Item) -> LoggedHistoryEntry {
|
||||
LoggedHistoryEntry {
|
||||
item: LoggedItem::from(item),
|
||||
metadata: LoggedSessionHistoryMetadata {
|
||||
entry_id: LoggedSessionHistoryEntryId::new(),
|
||||
origin: LoggedSessionHistoryOrigin::LegacyUnknown,
|
||||
derivation: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn canonical_layout_and_single_session_invariant() {
|
||||
@@ -445,7 +634,46 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn schema_v1_logs_are_validated_and_promoted_to_v2() {
|
||||
fn worker_session_store_keeps_paste_artifacts_inside_retention_root() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let store = WorkerSessionStore::new(root.path().join("session")).unwrap();
|
||||
let session_id = new_session_id();
|
||||
store
|
||||
.create_segment(session_id, new_segment_id(), &[])
|
||||
.unwrap();
|
||||
let content = "large paste body\n終端\n";
|
||||
let reference = store
|
||||
.write_paste_artifact(
|
||||
session_id,
|
||||
"entry-1",
|
||||
content,
|
||||
PasteArtifactLimits::default(),
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert!(
|
||||
root.path()
|
||||
.join(format!(
|
||||
"session/{PASTE_ARTIFACTS_DIR}/{}.json",
|
||||
reference.artifact_id
|
||||
))
|
||||
.is_file()
|
||||
);
|
||||
assert_eq!(
|
||||
store
|
||||
.read_paste_artifact(session_id, &reference.artifact_id)
|
||||
.unwrap()
|
||||
.1,
|
||||
content
|
||||
);
|
||||
assert!(matches!(
|
||||
store.read_paste_artifact(new_session_id(), &reference.artifact_id),
|
||||
Err(StoreError::Corrupt { .. })
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn schema_v1_logs_are_rewritten_and_promoted_to_v3() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let session_id = new_session_id();
|
||||
let segment_id = new_segment_id();
|
||||
@@ -467,7 +695,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn schema_v1_migration_rejects_corrupt_log_before_manifest_update() {
|
||||
fn schema_v1_migration_rejects_corrupt_log_before_v3_manifest_update() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let session_id = new_session_id();
|
||||
let manifest = SessionManifest {
|
||||
@@ -492,6 +720,280 @@ mod tests {
|
||||
assert_eq!(persisted.schema_version, LEGACY_SESSION_SCHEMA_VERSION);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn schema_v2_migration_rewrites_legacy_records_with_stable_unknown_provenance() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let session_id = new_session_id();
|
||||
let segment_id = new_segment_id();
|
||||
fs::create_dir_all(root.path().join(SEGMENTS_DIR)).unwrap();
|
||||
atomic_write_json(
|
||||
&root.path().join(SESSION_FILE),
|
||||
&SessionManifest {
|
||||
schema_version: PREVIOUS_SESSION_SCHEMA_VERSION,
|
||||
session_id,
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
let source = vec![
|
||||
serde_json::json!({
|
||||
"kind": "segment_start",
|
||||
"ts": 1,
|
||||
"session_id": session_id,
|
||||
"system_prompt": null,
|
||||
"config": agen::llm_client::RequestConfig::default(),
|
||||
"history": [LoggedItem::from(agen::Item::assistant_message("prior"))],
|
||||
"forked_from": null,
|
||||
"compacted_from": null
|
||||
}),
|
||||
serde_json::json!({
|
||||
"kind": "user_input",
|
||||
"ts": 2,
|
||||
"segments": [{ "kind": "text", "content": "hello" }],
|
||||
"extensions": []
|
||||
}),
|
||||
serde_json::json!({
|
||||
"kind": "assistant_item",
|
||||
"ts": 3,
|
||||
"item": LoggedItem::from(agen::Item::assistant_message("reply"))
|
||||
}),
|
||||
];
|
||||
let path = root
|
||||
.path()
|
||||
.join(SEGMENTS_DIR)
|
||||
.join(format!("{segment_id}.jsonl"));
|
||||
let mut bytes = Vec::new();
|
||||
for entry in source {
|
||||
serde_json::to_writer(&mut bytes, &entry).unwrap();
|
||||
bytes.push(b'\n');
|
||||
}
|
||||
fs::write(&path, bytes).unwrap();
|
||||
|
||||
let store = WorkerSessionStore::new(root.path()).unwrap();
|
||||
let first = store.read_all(session_id, segment_id).unwrap();
|
||||
assert!(matches!(first[0], LogEntry::AnnotatedSegmentStart { .. }));
|
||||
assert!(matches!(first[1], LogEntry::AnnotatedUserInput { .. }));
|
||||
assert!(matches!(first[2], LogEntry::AnnotatedAssistantItem { .. }));
|
||||
let first_bytes = fs::read(&path).unwrap();
|
||||
drop(store);
|
||||
|
||||
let reopened = WorkerSessionStore::new(root.path()).unwrap();
|
||||
assert_eq!(fs::read(&path).unwrap(), first_bytes);
|
||||
let snapshot = crate::public_snapshot::project_current_session_snapshot(
|
||||
&reopened.read_all(session_id, segment_id).unwrap(),
|
||||
);
|
||||
assert_eq!(snapshot.entries.len(), 3);
|
||||
assert_eq!(
|
||||
snapshot
|
||||
.entries
|
||||
.iter()
|
||||
.map(|entry| entry.timestamp)
|
||||
.collect::<Vec<_>>(),
|
||||
vec![1, 2, 3]
|
||||
);
|
||||
assert!(snapshot.entries.iter().all(|entry| {
|
||||
entry.provenance == protocol::SessionEntryProvenance::LegacyUnknown
|
||||
&& entry.entry_id.len() <= 64
|
||||
}));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn schema_v2_preflight_keeps_earlier_segments_unchanged_when_later_is_corrupt() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let session_id = new_session_id();
|
||||
let valid_segment = uuid::Uuid::from_u128(1);
|
||||
let corrupt_segment = uuid::Uuid::from_u128(2);
|
||||
fs::create_dir_all(root.path().join(SEGMENTS_DIR)).unwrap();
|
||||
atomic_write_json(
|
||||
&root.path().join(SESSION_FILE),
|
||||
&SessionManifest {
|
||||
schema_version: PREVIOUS_SESSION_SCHEMA_VERSION,
|
||||
session_id,
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
let manifest_before = fs::read(root.path().join(SESSION_FILE)).unwrap();
|
||||
|
||||
let valid_path = root
|
||||
.path()
|
||||
.join(SEGMENTS_DIR)
|
||||
.join(format!("{valid_segment}.jsonl"));
|
||||
let valid_entry = serde_json::json!({
|
||||
"kind": "segment_start",
|
||||
"ts": 1,
|
||||
"session_id": session_id,
|
||||
"system_prompt": null,
|
||||
"config": agen::llm_client::RequestConfig::default(),
|
||||
"history": [LoggedItem::from(agen::Item::assistant_message("prior"))],
|
||||
"forked_from": null,
|
||||
"compacted_from": null
|
||||
});
|
||||
let mut valid_bytes = serde_json::to_vec(&valid_entry).unwrap();
|
||||
valid_bytes.push(b'\n');
|
||||
fs::write(&valid_path, &valid_bytes).unwrap();
|
||||
let corrupt_path = root
|
||||
.path()
|
||||
.join(SEGMENTS_DIR)
|
||||
.join(format!("{corrupt_segment}.jsonl"));
|
||||
fs::write(&corrupt_path, b"{not-json}\n").unwrap();
|
||||
let corrupt_before = fs::read(&corrupt_path).unwrap();
|
||||
|
||||
let error = match WorkerSessionStore::new(root.path()) {
|
||||
Ok(_) => panic!("later corrupt segment must fail migration preflight"),
|
||||
Err(error) => error,
|
||||
};
|
||||
assert!(matches!(error, StoreError::Corrupt { .. }));
|
||||
assert_eq!(fs::read(&valid_path).unwrap(), valid_bytes);
|
||||
assert_eq!(fs::read(&corrupt_path).unwrap(), corrupt_before);
|
||||
assert_eq!(
|
||||
fs::read(root.path().join(SESSION_FILE)).unwrap(),
|
||||
manifest_before
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn current_jsonl_requires_annotations_across_append_rewrite_and_reopen() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let session_id = new_session_id();
|
||||
let segment_id = new_segment_id();
|
||||
let store = WorkerSessionStore::new(root.path()).unwrap();
|
||||
store
|
||||
.create_segment(
|
||||
session_id,
|
||||
segment_id,
|
||||
&[LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1,
|
||||
session_id,
|
||||
system_prompt: None,
|
||||
config: agen::llm_client::RequestConfig::default(),
|
||||
history: vec![annotated(agen::Item::user_message("seed"))],
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
}],
|
||||
)
|
||||
.unwrap();
|
||||
store
|
||||
.append(
|
||||
session_id,
|
||||
segment_id,
|
||||
&LogEntry::AnnotatedAssistantItem {
|
||||
ts: 2,
|
||||
entry: annotated(agen::Item::assistant_message("reply")),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let before_rewrite = store.read_all(session_id, segment_id).unwrap();
|
||||
store
|
||||
.create_segment(session_id, segment_id, &before_rewrite)
|
||||
.unwrap();
|
||||
drop(store);
|
||||
|
||||
let reopened = WorkerSessionStore::new(root.path()).unwrap();
|
||||
let restored = reopened.read_all(session_id, segment_id).unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(&restored).unwrap(),
|
||||
serde_json::to_value(&before_rewrite).unwrap()
|
||||
);
|
||||
for entry in &restored {
|
||||
match entry {
|
||||
LogEntry::AnnotatedSegmentStart { history, .. } => assert!(history.iter().all(
|
||||
|entry| !entry.metadata.entry_id.0.is_empty()
|
||||
&& matches!(
|
||||
entry.metadata.origin,
|
||||
LoggedSessionHistoryOrigin::LegacyUnknown
|
||||
)
|
||||
)),
|
||||
LogEntry::AnnotatedAssistantItem { entry, .. } => {
|
||||
assert!(!entry.metadata.entry_id.0.is_empty());
|
||||
assert!(matches!(
|
||||
entry.metadata.origin,
|
||||
LoggedSessionHistoryOrigin::LegacyUnknown
|
||||
));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
let log = fs::read_to_string(reopened.log_path(segment_id)).unwrap();
|
||||
for line in log.lines() {
|
||||
let value: serde_json::Value = serde_json::from_str(line).unwrap();
|
||||
let kind = value["kind"].as_str().unwrap();
|
||||
assert!(
|
||||
!matches!(
|
||||
kind,
|
||||
"segment_start"
|
||||
| "user_input"
|
||||
| "assistant_item"
|
||||
| "tool_result"
|
||||
| "system_item"
|
||||
),
|
||||
"current-schema JSONL contains legacy history record: {kind}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn schema_v3_rejects_legacy_records_and_new_writes_are_canonical() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let session_id = new_session_id();
|
||||
let segment_id = new_segment_id();
|
||||
let store = WorkerSessionStore::new(root.path()).unwrap();
|
||||
store
|
||||
.create_segment(
|
||||
session_id,
|
||||
segment_id,
|
||||
&[LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1,
|
||||
session_id,
|
||||
system_prompt: None,
|
||||
config: agen::llm_client::RequestConfig::default(),
|
||||
history: vec![annotated(agen::Item::assistant_message("seed"))],
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
}],
|
||||
)
|
||||
.unwrap();
|
||||
store
|
||||
.append(
|
||||
session_id,
|
||||
segment_id,
|
||||
&LogEntry::AnnotatedUserInput {
|
||||
ts: 2,
|
||||
segments: vec![protocol::Segment::Text {
|
||||
content: "new".into(),
|
||||
}],
|
||||
history: vec![annotated(agen::Item::user_message("new"))],
|
||||
extensions: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
let entries = store.read_all(session_id, segment_id).unwrap();
|
||||
assert!(matches!(entries[0], LogEntry::AnnotatedSegmentStart { .. }));
|
||||
assert!(matches!(entries[1], LogEntry::AnnotatedUserInput { .. }));
|
||||
drop(store);
|
||||
|
||||
let path = root
|
||||
.path()
|
||||
.join(SEGMENTS_DIR)
|
||||
.join(format!("{segment_id}.jsonl"));
|
||||
let mut file = OpenOptions::new().append(true).open(path).unwrap();
|
||||
serde_json::to_writer(
|
||||
&mut file,
|
||||
&serde_json::json!({
|
||||
"kind": "system_item",
|
||||
"ts": 3,
|
||||
"item": { "kind": "legacy_ignored", "slug": "legacy" }
|
||||
}),
|
||||
)
|
||||
.unwrap();
|
||||
file.write_all(b"\n").unwrap();
|
||||
let error = match WorkerSessionStore::new(root.path()) {
|
||||
Ok(_) => panic!("schema v3 must reject a legacy history record"),
|
||||
Err(error) => error,
|
||||
};
|
||||
assert!(matches!(error, StoreError::Corrupt { .. }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reopen_preserves_session_and_segment_ids() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
|
||||
@@ -1,12 +1,25 @@
|
||||
use agen::EngineResult;
|
||||
use agen::llm_client::types::{Item, RequestConfig};
|
||||
use session_store::{
|
||||
FsStore, LogEntry, Store, TraceEntry, collect_state, new_segment_id, new_session_id,
|
||||
FsStore, LogEntry, LoggedHistoryEntry, LoggedItem, LoggedSessionHistoryEntryId,
|
||||
LoggedSessionHistoryMetadata, LoggedSessionHistoryOrigin, Store, TraceEntry, collect_state,
|
||||
new_segment_id, new_session_id,
|
||||
};
|
||||
use std::io::Write;
|
||||
|
||||
fn annotated(item: Item) -> LoggedHistoryEntry {
|
||||
LoggedHistoryEntry {
|
||||
item: LoggedItem::from(item),
|
||||
metadata: LoggedSessionHistoryMetadata {
|
||||
entry_id: LoggedSessionHistoryEntryId::new(),
|
||||
origin: LoggedSessionHistoryOrigin::LegacyUnknown,
|
||||
derivation: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn nil_session_start(ts: u64, session_id: uuid::Uuid) -> LogEntry {
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts,
|
||||
session_id,
|
||||
system_prompt: None,
|
||||
@@ -25,7 +38,7 @@ fn round_trip_write_and_read() {
|
||||
let segid = new_segment_id();
|
||||
|
||||
let entries = vec![
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: sid,
|
||||
system_prompt: Some("You are helpful.".into()),
|
||||
@@ -34,14 +47,15 @@ fn round_trip_write_and_read() {
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
},
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 2000,
|
||||
extensions: vec![],
|
||||
segments: vec![protocol::Segment::text("Hello")],
|
||||
history: vec![annotated(Item::user_message("Hello"))],
|
||||
},
|
||||
LogEntry::AssistantItem {
|
||||
LogEntry::AnnotatedAssistantItem {
|
||||
ts: 3000,
|
||||
item: Item::assistant_message("Hi there!").into(),
|
||||
entry: annotated(Item::assistant_message("Hi there!")),
|
||||
},
|
||||
LogEntry::TurnEnd {
|
||||
ts: 3100,
|
||||
@@ -79,14 +93,14 @@ fn create_segment_writes_all_entries() {
|
||||
let sid = new_session_id();
|
||||
let segid = new_segment_id();
|
||||
|
||||
let entries = [LogEntry::SegmentStart {
|
||||
let entries = [LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: sid,
|
||||
system_prompt: None,
|
||||
config: RequestConfig::default(),
|
||||
history: vec![
|
||||
Item::user_message("seed").into(),
|
||||
Item::assistant_message("ok").into(),
|
||||
annotated(Item::user_message("seed")),
|
||||
annotated(Item::assistant_message("ok")),
|
||||
],
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
@@ -205,7 +219,7 @@ fn read_entry_count_matches_append_tally() {
|
||||
let segid = new_segment_id();
|
||||
|
||||
let entries = [
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: sid,
|
||||
system_prompt: None,
|
||||
@@ -214,10 +228,11 @@ fn read_entry_count_matches_append_tally() {
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
},
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 2000,
|
||||
extensions: vec![],
|
||||
segments: vec![protocol::Segment::text("Hello")],
|
||||
history: vec![annotated(Item::user_message("Hello"))],
|
||||
},
|
||||
];
|
||||
|
||||
@@ -254,10 +269,11 @@ fn unterminated_utf8_tail_is_ignored_and_replaced_on_append() {
|
||||
assert_eq!(store.read_all(sid, segid).unwrap().len(), 1);
|
||||
assert_eq!(store.read_entry_count(sid, segid).unwrap(), 1);
|
||||
|
||||
let next = LogEntry::UserInput {
|
||||
let next = LogEntry::AnnotatedUserInput {
|
||||
ts: 2,
|
||||
extensions: vec![],
|
||||
segments: vec![protocol::Segment::text("recovered")],
|
||||
history: vec![annotated(Item::user_message("recovered"))],
|
||||
};
|
||||
store.append(sid, segid, &next).unwrap();
|
||||
|
||||
|
||||
@@ -3,19 +3,35 @@ mod common;
|
||||
use std::ops::{Deref, DerefMut};
|
||||
use std::sync::Arc;
|
||||
|
||||
use agen::interceptor::{Interceptor, TurnEndAction};
|
||||
use agen::interceptor::{AssistantTurnEndContext, Interceptor, InterceptorResult, TurnEndAction};
|
||||
use agen::llm_client::event::{Event, ResponseStatus, StatusEvent};
|
||||
use agen::llm_client::types::{Item, RequestConfig};
|
||||
use agen::tool::{Tool, ToolDefinition, ToolError, ToolMeta, ToolOutput};
|
||||
use agen::{Engine, History};
|
||||
use async_trait::async_trait;
|
||||
use common::MockLlmClient;
|
||||
use protocol::{Segment, SessionSnapshotEntryData, UploadedFileAvailability, UploadedFileRef};
|
||||
use session_store::{FsStore, LogEntry, SegmentStartState, Store, collect_state};
|
||||
|
||||
// =============================================================================
|
||||
// Helpers
|
||||
// =============================================================================
|
||||
|
||||
fn annotated(items: &[Item]) -> Vec<session_store::LoggedHistoryEntry> {
|
||||
items
|
||||
.iter()
|
||||
.cloned()
|
||||
.map(|item| session_store::LoggedHistoryEntry {
|
||||
item: session_store::LoggedItem::from(item),
|
||||
metadata: session_store::LoggedSessionHistoryMetadata {
|
||||
entry_id: session_store::LoggedSessionHistoryEntryId::new(),
|
||||
origin: session_store::LoggedSessionHistoryOrigin::LegacyUnknown,
|
||||
derivation: None,
|
||||
},
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn simple_text_events() -> Vec<Event> {
|
||||
vec![
|
||||
Event::text_block_start(0),
|
||||
@@ -84,8 +100,11 @@ struct PausePolicy;
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for PausePolicy {
|
||||
async fn on_turn_end(&self, _history: &[Item]) -> TurnEndAction {
|
||||
TurnEndAction::Pause
|
||||
async fn on_assistant_turn_end(
|
||||
&self,
|
||||
_context: AssistantTurnEndContext<'_>,
|
||||
) -> InterceptorResult<TurnEndAction> {
|
||||
Ok(TurnEndAction::Pause)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -144,6 +163,7 @@ async fn run_and_persist(
|
||||
session_id,
|
||||
segment_id,
|
||||
vec![protocol::Segment::text(input)],
|
||||
annotated(&[Item::user_message(input)]),
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
@@ -154,8 +174,8 @@ async fn run_and_persist(
|
||||
worker.engine = locked.unlock();
|
||||
|
||||
let projected = worker.history();
|
||||
let new_items = &projected[history_before..];
|
||||
session_store::save_delta(store, session_id, segment_id, new_items).unwrap();
|
||||
let new_items = annotated(&projected[history_before..]);
|
||||
session_store::save_delta(store, session_id, segment_id, &new_items).unwrap();
|
||||
session_store::save_turn_end(store, session_id, segment_id, worker.turn_count()).unwrap();
|
||||
|
||||
match &result {
|
||||
@@ -178,7 +198,7 @@ async fn run_and_persist(
|
||||
)
|
||||
.unwrap();
|
||||
}
|
||||
agen::EngineRunExit::Interrupted(agen::StopReason::LimitReached) => {
|
||||
agen::EngineRunExit::Interrupted(agen::RunInterruptionReason::LimitReached) => {
|
||||
session_store::save_run_completed(
|
||||
store,
|
||||
session_id,
|
||||
@@ -219,7 +239,8 @@ async fn session_run_logs_entries() {
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: &worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -237,7 +258,10 @@ async fn session_run_logs_entries() {
|
||||
);
|
||||
|
||||
// First entry is SegmentStart
|
||||
assert!(matches!(&entries[0], LogEntry::SegmentStart { .. }));
|
||||
assert!(matches!(
|
||||
&entries[0],
|
||||
LogEntry::AnnotatedSegmentStart { .. }
|
||||
));
|
||||
|
||||
// Has a RunCompleted with Finished
|
||||
let has_finished = entries.iter().any(|e| {
|
||||
@@ -264,7 +288,8 @@ async fn session_restore_round_trip() {
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: &worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -303,7 +328,8 @@ async fn session_run_with_tool_call() {
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: &worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -314,12 +340,12 @@ async fn session_run_with_tool_call() {
|
||||
|
||||
let has_tool_results = entries
|
||||
.iter()
|
||||
.any(|e| matches!(e, LogEntry::ToolResult { .. }));
|
||||
.any(|e| matches!(e, LogEntry::AnnotatedToolResult { .. }));
|
||||
assert!(has_tool_results, "should have ToolResult entry");
|
||||
|
||||
let has_assistant = entries
|
||||
.iter()
|
||||
.any(|e| matches!(e, LogEntry::AssistantItem { .. }));
|
||||
.any(|e| matches!(e, LogEntry::AnnotatedAssistantItem { .. }));
|
||||
assert!(has_assistant, "should have AssistantItem entry");
|
||||
}
|
||||
|
||||
@@ -327,7 +353,8 @@ async fn session_run_with_tool_call() {
|
||||
async fn session_resume_after_pause() {
|
||||
let (_dir, store) = make_store();
|
||||
|
||||
// First run: tool call with pause policy → Paused
|
||||
// First terminal assistant response requests a tool; the assistant-turn
|
||||
// interceptor pauses before the Engine enters the tool phase.
|
||||
let client = MockLlmClient::with_responses(tool_call_events());
|
||||
let mut worker = TestWorker::new(Engine::new(client));
|
||||
worker.register_tool(weather_tool_definition());
|
||||
@@ -338,7 +365,8 @@ async fn session_resume_after_pause() {
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: &worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -362,7 +390,7 @@ async fn session_resume_after_pause() {
|
||||
// Restore state and verify
|
||||
let state = session_store::restore(&store, sid, segid).unwrap();
|
||||
assert!(state.last_run_interrupted);
|
||||
assert_eq!(state.active_run_turn_count, Some(2));
|
||||
assert_eq!(state.active_run_turn_count, Some(1));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -377,7 +405,8 @@ async fn session_fork_creates_new_session() {
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: &worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -385,25 +414,38 @@ async fn session_fork_creates_new_session() {
|
||||
let (worker, _) = run_and_persist(worker, &store, sid, segid, "Hello").await;
|
||||
|
||||
let original_history_len = worker.history().len();
|
||||
let source_user_segments = session_store::restore(&store, sid, segid)
|
||||
.unwrap()
|
||||
.user_segments;
|
||||
let (fork_sid, fork_segid) = session_store::fork(
|
||||
&store,
|
||||
sid,
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: &worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
user_segments: source_user_segments.clone(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
assert_ne!(fork_sid, sid, "`fork` mints a fresh Session");
|
||||
|
||||
// Fork should have a SegmentStart with the current history
|
||||
// Fork should have an annotated seed and typed input checkpoint.
|
||||
let fork_entries = store.read_all(fork_sid, fork_segid).unwrap();
|
||||
assert_eq!(fork_entries.len(), 1);
|
||||
assert!(matches!(&fork_entries[0], LogEntry::SegmentStart { .. }));
|
||||
assert_eq!(fork_entries.len(), 2);
|
||||
assert!(matches!(
|
||||
&fork_entries[0],
|
||||
LogEntry::AnnotatedSegmentStart { .. }
|
||||
));
|
||||
assert!(matches!(
|
||||
&fork_entries[1],
|
||||
LogEntry::InputSegmentsCheckpoint { .. }
|
||||
));
|
||||
|
||||
let fork_state = collect_state(&fork_entries);
|
||||
assert_eq!(fork_state.session_id, Some(fork_sid));
|
||||
assert_eq!(fork_state.history.len(), original_history_len);
|
||||
assert_eq!(fork_state.user_segments, source_user_segments);
|
||||
assert_eq!(fork_state.system_prompt.as_deref(), Some("System prompt"));
|
||||
}
|
||||
|
||||
@@ -418,7 +460,8 @@ async fn session_fork_at_truncates_within_session() {
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: &worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -432,7 +475,11 @@ async fn session_fork_at_truncates_within_session() {
|
||||
let fork_segid = session_store::fork_at(&store, sid, segid, worker.turn_count()).unwrap();
|
||||
|
||||
let fork_entries = store.read_all(sid, fork_segid).unwrap();
|
||||
assert_eq!(fork_entries.len(), 1); // Just the new SegmentStart
|
||||
assert_eq!(fork_entries.len(), 2);
|
||||
assert!(matches!(
|
||||
&fork_entries[1],
|
||||
LogEntry::InputSegmentsCheckpoint { .. }
|
||||
));
|
||||
|
||||
let fork_state = collect_state(&fork_entries);
|
||||
assert_eq!(fork_state.session_id, Some(sid), "fork_at inherits Session");
|
||||
@@ -444,7 +491,25 @@ async fn session_fork_at_truncates_within_session() {
|
||||
.position(|e| matches!(e, LogEntry::TurnEnd { turn_count, .. } if *turn_count == worker.turn_count()))
|
||||
.expect("source segment has the matching TurnEnd");
|
||||
let source_state_at_fork = collect_state(&all_entries[..=turn_end_pos]);
|
||||
assert_eq!(fork_state.user_segments, source_state_at_fork.user_segments);
|
||||
assert_eq!(fork_state.history.len(), source_state_at_fork.history.len());
|
||||
assert_eq!(
|
||||
fork_state.annotated_history, source_state_at_fork.annotated_history,
|
||||
"fork_at must preserve every retained history entry identity and provenance",
|
||||
);
|
||||
assert!(fork_state.annotated_history.iter().all(|entry| {
|
||||
!entry.metadata.entry_id.0.is_empty()
|
||||
&& matches!(
|
||||
entry.metadata.origin,
|
||||
session_store::LoggedSessionHistoryOrigin::LegacyUnknown
|
||||
| session_store::LoggedSessionHistoryOrigin::HumanInput { .. }
|
||||
| session_store::LoggedSessionHistoryOrigin::WorkerInput { .. }
|
||||
| session_store::LoggedSessionHistoryOrigin::BackendInstruction { .. }
|
||||
| session_store::LoggedSessionHistoryOrigin::ModelOutput { .. }
|
||||
| session_store::LoggedSessionHistoryOrigin::ToolOutput { .. }
|
||||
| session_store::LoggedSessionHistoryOrigin::DerivedSummary
|
||||
)
|
||||
}));
|
||||
|
||||
// list_segments should show both source and fork in the same Session.
|
||||
let segs = store.list_segments(sid).unwrap();
|
||||
@@ -452,6 +517,84 @@ async fn session_fork_at_truncates_within_session() {
|
||||
assert!(segs.contains(&fork_segid));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rewound_fork_preserves_uploaded_file_segments_in_snapshot() {
|
||||
let (_dir, store) = make_store();
|
||||
let config = RequestConfig::default();
|
||||
let (sid, segid) = session_store::create_segment(
|
||||
&store,
|
||||
SegmentStartState {
|
||||
system_prompt: Some("System prompt"),
|
||||
config: &config,
|
||||
history: Vec::new(),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
let uploaded = UploadedFileRef {
|
||||
artifact_id: "uploaded-file-1".into(),
|
||||
file_name: "notes.txt".into(),
|
||||
media_type: "text/plain".into(),
|
||||
created_at_ms: 123,
|
||||
availability: UploadedFileAvailability::Available,
|
||||
byte_len: 5,
|
||||
sha256: "a".repeat(64),
|
||||
source_entry_id: Some("entry-1".into()),
|
||||
};
|
||||
let segments = vec![Segment::UploadedFile {
|
||||
file: uploaded.clone(),
|
||||
}];
|
||||
session_store::save_user_input(
|
||||
&store,
|
||||
sid,
|
||||
segid,
|
||||
segments.clone(),
|
||||
annotated(&[Item::user_message(Segment::flatten_to_text(&segments))]),
|
||||
)
|
||||
.unwrap();
|
||||
session_store::save_turn_end(&store, sid, segid, 1).unwrap();
|
||||
|
||||
let fork_segid = session_store::fork_at(&store, sid, segid, 1).unwrap();
|
||||
let fork_entries = store.read_all(sid, fork_segid).unwrap();
|
||||
let snapshot = session_store::public_snapshot::project_session_snapshot(sid, &fork_entries);
|
||||
|
||||
assert!(fork_entries.iter().any(|entry| matches!(
|
||||
entry,
|
||||
LogEntry::InputSegmentsCheckpoint { user_segments, .. }
|
||||
if user_segments == &vec![segments.clone()]
|
||||
)));
|
||||
assert!(snapshot.entries.iter().any(|entry| matches!(
|
||||
&entry.data,
|
||||
SessionSnapshotEntryData::UserInput { segments: restored }
|
||||
if restored == &segments
|
||||
)));
|
||||
|
||||
let fork_state = collect_state(&fork_entries);
|
||||
let (copied_session_id, copied_segment_id) = session_store::fork(
|
||||
&store,
|
||||
sid,
|
||||
SegmentStartState {
|
||||
system_prompt: fork_state.system_prompt.as_deref(),
|
||||
config: &fork_state.config,
|
||||
history: fork_state.annotated_history.clone(),
|
||||
user_segments: fork_state.user_segments.clone(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
let copied_entries = store
|
||||
.read_all(copied_session_id, copied_segment_id)
|
||||
.unwrap();
|
||||
let copied_snapshot = session_store::public_snapshot::project_session_snapshot(
|
||||
copied_session_id,
|
||||
&copied_entries,
|
||||
);
|
||||
assert!(copied_snapshot.entries.iter().any(|entry| matches!(
|
||||
&entry.data,
|
||||
SessionSnapshotEntryData::UserInput { segments: restored }
|
||||
if restored == &segments
|
||||
)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn session_config_changed_logged() {
|
||||
let (_dir, store) = make_store();
|
||||
@@ -463,7 +606,8 @@ async fn session_config_changed_logged() {
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: &worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -496,7 +640,8 @@ async fn session_auto_forks_on_conflict() {
|
||||
SegmentStartState {
|
||||
system_prompt: worker_a.get_system_prompt(),
|
||||
config: worker_a.request_config(),
|
||||
history: &worker_a.history(),
|
||||
history: annotated(&worker_a.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -505,12 +650,14 @@ async fn session_auto_forks_on_conflict() {
|
||||
let mut entries_written: usize = 1;
|
||||
|
||||
// Simulate another Worker writing to the same segment behind our back.
|
||||
let extra_entry = LogEntry::UserInput {
|
||||
ts: 9999,
|
||||
extensions: vec![],
|
||||
segments: vec![protocol::Segment::text("Interloper")],
|
||||
};
|
||||
store.append(sid, original_segid, &extra_entry).unwrap();
|
||||
session_store::save_user_input(
|
||||
&store,
|
||||
sid,
|
||||
original_segid,
|
||||
vec![protocol::Segment::text("Interloper")],
|
||||
annotated(&[Item::user_message("Interloper")]),
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
// Now the on-disk count exceeds our tally — ensure_head_or_fork should auto-fork.
|
||||
session_store::ensure_head_or_fork(
|
||||
@@ -522,7 +669,8 @@ async fn session_auto_forks_on_conflict() {
|
||||
SegmentStartState {
|
||||
system_prompt: worker_a.get_system_prompt(),
|
||||
config: worker_a.request_config(),
|
||||
history: &worker_a.history(),
|
||||
history: annotated(&worker_a.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -543,7 +691,7 @@ async fn session_auto_forks_on_conflict() {
|
||||
// The new segment records its lineage forward via forked_from; the
|
||||
// source segment is left immutable (no terminal marker written back).
|
||||
match &fork_entries[0] {
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
forked_from: Some(origin),
|
||||
..
|
||||
} => {
|
||||
@@ -563,7 +711,7 @@ async fn session_auto_forks_on_conflict() {
|
||||
);
|
||||
let has_interloper = original_entries
|
||||
.iter()
|
||||
.any(|e| matches!(e, LogEntry::UserInput { .. }));
|
||||
.any(|e| matches!(e, LogEntry::AnnotatedUserInput { .. }));
|
||||
assert!(has_interloper);
|
||||
}
|
||||
|
||||
@@ -581,7 +729,8 @@ async fn nested_past_fork_leaves_ancestors_immutable() {
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: &worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -612,13 +761,20 @@ async fn nested_past_fork_leaves_ancestors_immutable() {
|
||||
let fork1_entries = store.read_all(sid, fork1).unwrap();
|
||||
assert_eq!(
|
||||
fork1_entries.len(),
|
||||
1,
|
||||
"fork1 is just its SegmentStart seed"
|
||||
2,
|
||||
"fork1 stores its SegmentStart and typed input checkpoint"
|
||||
);
|
||||
|
||||
// fork2's lineage points at fork1, not the root.
|
||||
match &store.read_all(sid, fork2).unwrap()[0] {
|
||||
LogEntry::SegmentStart {
|
||||
// fork2's lineage points at fork1, not the root, and the typed seed remains
|
||||
// intact across the nested turn-zero fork.
|
||||
let fork2_entries = store.read_all(sid, fork2).unwrap();
|
||||
assert_eq!(fork2_entries.len(), 2);
|
||||
assert_eq!(
|
||||
collect_state(&fork2_entries).user_segments,
|
||||
collect_state(&fork1_entries).user_segments
|
||||
);
|
||||
match &fork2_entries[0] {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
forked_from: Some(origin),
|
||||
..
|
||||
} => assert_eq!(origin.segment_id, fork1),
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
[package]
|
||||
name = "standalone"
|
||||
description = "In-process standalone Worker host"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
|
||||
[dependencies]
|
||||
agen.workspace = true
|
||||
client.workspace = true
|
||||
fs4.workspace = true
|
||||
manifest.workspace = true
|
||||
protocol.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
session-store.workspace = true
|
||||
thiserror.workspace = true
|
||||
tokio = { workspace = true, features = ["rt", "sync", "time"] }
|
||||
uuid = { workspace = true, features = ["v7"] }
|
||||
worker.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
async-trait.workspace = true
|
||||
futures.workspace = true
|
||||
tempfile.workspace = true
|
||||
tokio = { workspace = true, features = ["macros", "rt-multi-thread", "time"] }
|
||||
@@ -0,0 +1,544 @@
|
||||
use std::path::PathBuf;
|
||||
use std::time::Duration;
|
||||
|
||||
use agen::llm_client::client::LlmClient;
|
||||
use client::Client;
|
||||
use client::transport::in_process::{Peer as InProcessPeer, Socket as InProcessSocket};
|
||||
use manifest::ScopeRule;
|
||||
use protocol::stream::{decode_method, encode_event};
|
||||
use protocol::{Event, Method, WorkerId};
|
||||
use session_store::{
|
||||
CombinedStore, FsStore, FsWorkerStore, WorkerActiveSegmentRef, WorkerMetadataStore,
|
||||
};
|
||||
use thiserror::Error;
|
||||
use worker::bootstrap::{
|
||||
WorkerBootstrap, WorkerBootstrapError, WorkerBootstrapLayout, bash_output_dir_for_worker_id,
|
||||
};
|
||||
use worker::controller::WorkerControllerTransport;
|
||||
use worker::ipc::protocol_session::{
|
||||
WorkerProtocolSessionStreams, dispatch_worker_protocol_method, live_log_entry_event,
|
||||
subscribe_worker_protocol_session,
|
||||
};
|
||||
use worker::runtime::worker_allocation::ScopeLockError;
|
||||
use worker::{BootstrappedWorker, WorkerError, WorkerFilesystemAuthority, WorkerWorkspaceContext};
|
||||
|
||||
use crate::launch::ResolvedStandaloneLaunch;
|
||||
use crate::store::{
|
||||
StaleLeasePolicy, StandaloneShutdownReason, StandaloneStoreError, StandaloneWorkerLease,
|
||||
StandaloneWorkerRecord, StandaloneWorkerStore,
|
||||
};
|
||||
|
||||
const DEFAULT_SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(10);
|
||||
type StandaloneBackingStore = CombinedStore<FsStore, FsWorkerStore>;
|
||||
|
||||
/// One client-owned top-level Worker and its standalone Worker authority.
|
||||
///
|
||||
/// The host deliberately exposes the existing typed Worker protocol rather than owning an
|
||||
/// HTTP/WebSocket server or creating Runtime/Workspace/Ticket/Workdir domain records.
|
||||
pub struct StandaloneHost {
|
||||
handle: worker::WorkerHandle,
|
||||
shutdown: Option<worker::controller::ShutdownReceiver>,
|
||||
shutdown_timeout: Duration,
|
||||
store: StandaloneWorkerStore,
|
||||
worker_store: FsWorkerStore,
|
||||
record: StandaloneWorkerRecord,
|
||||
lease: Option<StandaloneWorkerLease>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Error)]
|
||||
pub enum StandaloneStartupError {
|
||||
#[error("the standalone state store could not be opened or validated")]
|
||||
StateStore,
|
||||
#[error("the standalone Worker is already active")]
|
||||
WorkerActive,
|
||||
#[error("the standalone Worker lease cannot be observed safely; recovery is rejected")]
|
||||
LeaseLivenessUnknown,
|
||||
#[error("the standalone Worker working directory is unavailable or changed")]
|
||||
WorkingDirectoryUnavailable,
|
||||
#[error(
|
||||
"requested scope `{}` conflicts with worker allocation `{competitor}` rule `{}`",
|
||||
requested_rule.target.display(),
|
||||
competitor_rule.target.display()
|
||||
)]
|
||||
ScopeConflict {
|
||||
competitor: String,
|
||||
requested_rule: ScopeRule,
|
||||
competitor_rule: ScopeRule,
|
||||
},
|
||||
#[error("the resolved Worker configuration or persisted history is invalid")]
|
||||
WorkerConfiguration,
|
||||
#[error("the configured model provider is unavailable")]
|
||||
ModelProvider,
|
||||
#[error("the fixed standalone feature composition could not be installed")]
|
||||
FeatureComposition,
|
||||
#[error("the in-process Worker controller could not start")]
|
||||
Controller,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
|
||||
pub enum StandaloneShutdownError {
|
||||
#[error("the standalone Worker did not stop before the shutdown deadline")]
|
||||
DeadlineExceeded,
|
||||
#[error("the standalone Worker shutdown confirmation was lost")]
|
||||
ConfirmationLost,
|
||||
#[error("the standalone Worker final state could not be committed")]
|
||||
StateStore,
|
||||
}
|
||||
|
||||
impl StandaloneHost {
|
||||
pub async fn start(launch: ResolvedStandaloneLaunch) -> Result<Self, StandaloneStartupError> {
|
||||
Self::start_with_optional_model_client(launch, None).await
|
||||
}
|
||||
|
||||
pub async fn start_with_model_client<C>(
|
||||
launch: ResolvedStandaloneLaunch,
|
||||
model_client: C,
|
||||
) -> Result<Self, StandaloneStartupError>
|
||||
where
|
||||
C: LlmClient + 'static,
|
||||
{
|
||||
Self::start_with_optional_model_client(launch, Some(Box::new(model_client))).await
|
||||
}
|
||||
|
||||
async fn start_with_optional_model_client(
|
||||
launch: ResolvedStandaloneLaunch,
|
||||
model_client: Option<Box<dyn LlmClient>>,
|
||||
) -> Result<Self, StandaloneStartupError> {
|
||||
let store =
|
||||
StandaloneWorkerStore::open(&launch.state_dir).map_err(classify_store_startup_error)?;
|
||||
let allocation = store
|
||||
.allocate(&launch.cwd, StaleLeasePolicy::Reject)
|
||||
.map_err(classify_store_startup_error)?;
|
||||
let worker_id = allocation.worker_id();
|
||||
|
||||
// WorkerId is the stable identity. The current Worker store remains
|
||||
// name-keyed, so keep its derived storage key separate from the
|
||||
// user-facing profile name.
|
||||
let manifest = launch.profile.manifest.clone();
|
||||
let storage_key = format!("standalone-{worker_id}");
|
||||
let mut bootstrap_manifest = manifest.clone();
|
||||
bootstrap_manifest.worker.name = storage_key.clone();
|
||||
let (backing_store, worker_store) = match backing_store(&store, worker_id) {
|
||||
Ok(stores) => stores,
|
||||
Err(error) => {
|
||||
let _ = store.abandon_allocation(allocation);
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
let filesystem_authority =
|
||||
WorkerFilesystemAuthority::local(launch.cwd.clone(), launch.cwd.clone());
|
||||
let workspace_context = WorkerWorkspaceContext::local_filesystem(None);
|
||||
let runtime_base = store.runtime_dir(worker_id);
|
||||
let bash_output_dir = bash_output_dir_for_worker_id(worker_id);
|
||||
|
||||
let mut bootstrap = WorkerBootstrap::new(
|
||||
bootstrap_manifest,
|
||||
backing_store,
|
||||
launch.prompt_catalog,
|
||||
workspace_context,
|
||||
filesystem_authority,
|
||||
WorkerBootstrapLayout::Direct {
|
||||
runtime_base,
|
||||
bash_output_dir,
|
||||
},
|
||||
WorkerControllerTransport::InProcess,
|
||||
);
|
||||
if let Some(model_client) = model_client {
|
||||
bootstrap = bootstrap.with_model_client(model_client);
|
||||
}
|
||||
let started = match bootstrap.start().await {
|
||||
Ok(started) => started,
|
||||
Err(error) => {
|
||||
let _ = store.abandon_allocation(allocation);
|
||||
return Err(classify_startup_error(error));
|
||||
}
|
||||
};
|
||||
let active = match active_pointer(&worker_store, &storage_key) {
|
||||
Ok(active) => active,
|
||||
Err(error) => {
|
||||
stop_started_worker(started).await;
|
||||
let _ = store.abandon_allocation(allocation);
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
let record = match store.commit_created(
|
||||
&allocation,
|
||||
manifest,
|
||||
storage_key,
|
||||
active.session_id,
|
||||
active.segment_id,
|
||||
) {
|
||||
Ok(record) => record,
|
||||
Err(_) => {
|
||||
stop_started_worker(started).await;
|
||||
let _ = store.abandon_allocation(allocation);
|
||||
return Err(StandaloneStartupError::StateStore);
|
||||
}
|
||||
};
|
||||
Ok(Self::from_started(
|
||||
started,
|
||||
store,
|
||||
worker_store,
|
||||
record,
|
||||
allocation.into_lease(),
|
||||
))
|
||||
}
|
||||
|
||||
pub async fn restore(
|
||||
state_dir: PathBuf,
|
||||
worker_id: WorkerId,
|
||||
) -> Result<Self, StandaloneStartupError> {
|
||||
Self::restore_with_optional_model_client(state_dir, worker_id, None).await
|
||||
}
|
||||
|
||||
pub async fn restore_with_model_client<C>(
|
||||
state_dir: PathBuf,
|
||||
worker_id: WorkerId,
|
||||
model_client: C,
|
||||
) -> Result<Self, StandaloneStartupError>
|
||||
where
|
||||
C: LlmClient + 'static,
|
||||
{
|
||||
Self::restore_with_optional_model_client(state_dir, worker_id, Some(Box::new(model_client)))
|
||||
.await
|
||||
}
|
||||
|
||||
async fn restore_with_optional_model_client(
|
||||
state_dir: PathBuf,
|
||||
worker_id: WorkerId,
|
||||
model_client: Option<Box<dyn LlmClient>>,
|
||||
) -> Result<Self, StandaloneStartupError> {
|
||||
let store = StandaloneWorkerStore::open(state_dir).map_err(classify_store_startup_error)?;
|
||||
let record = store
|
||||
.load(worker_id)
|
||||
.map_err(classify_store_startup_error)?;
|
||||
record.cwd.verify().map_err(classify_store_startup_error)?;
|
||||
let lease = store
|
||||
.acquire_lease(worker_id, StaleLeasePolicy::Recover)
|
||||
.map_err(classify_store_startup_error)?;
|
||||
let (backing_store, worker_store) = backing_store(&store, worker_id)?;
|
||||
let storage_key = record.storage_key.clone();
|
||||
let mut manifest = record.manifest.clone();
|
||||
manifest.worker.name = storage_key.clone();
|
||||
let filesystem_authority = WorkerFilesystemAuthority::local(
|
||||
record.cwd.canonical_path.clone(),
|
||||
record.cwd.canonical_path.clone(),
|
||||
);
|
||||
let workspace_context = WorkerWorkspaceContext::local_filesystem(None);
|
||||
let runtime_base = store.runtime_dir(worker_id);
|
||||
let bash_output_dir = bash_output_dir_for_worker_id(worker_id);
|
||||
|
||||
let mut bootstrap = WorkerBootstrap::new(
|
||||
manifest,
|
||||
backing_store,
|
||||
worker::PromptCatalogSource::builtins_only(),
|
||||
workspace_context,
|
||||
filesystem_authority,
|
||||
WorkerBootstrapLayout::Direct {
|
||||
runtime_base,
|
||||
bash_output_dir,
|
||||
},
|
||||
WorkerControllerTransport::InProcess,
|
||||
);
|
||||
if let Some(model_client) = model_client {
|
||||
bootstrap = bootstrap.with_model_client(model_client);
|
||||
}
|
||||
let prepared = bootstrap
|
||||
.prepare_restored(&storage_key)
|
||||
.await
|
||||
.map_err(classify_startup_error)?;
|
||||
let started = prepared.start().await.map_err(classify_startup_error)?;
|
||||
let active = match active_pointer(&worker_store, &storage_key) {
|
||||
Ok(active) => active,
|
||||
Err(error) => {
|
||||
stop_started_worker(started).await;
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
let record =
|
||||
match store.update_active_pointer(&record, active.session_id, active.segment_id) {
|
||||
Ok(record) => record,
|
||||
Err(_) => {
|
||||
stop_started_worker(started).await;
|
||||
lease.retain();
|
||||
return Err(StandaloneStartupError::StateStore);
|
||||
}
|
||||
};
|
||||
Ok(Self::from_started(
|
||||
started,
|
||||
store,
|
||||
worker_store,
|
||||
record,
|
||||
lease,
|
||||
))
|
||||
}
|
||||
|
||||
fn from_started(
|
||||
started: BootstrappedWorker,
|
||||
store: StandaloneWorkerStore,
|
||||
worker_store: FsWorkerStore,
|
||||
record: StandaloneWorkerRecord,
|
||||
lease: StandaloneWorkerLease,
|
||||
) -> Self {
|
||||
Self {
|
||||
handle: started.handle,
|
||||
shutdown: Some(started.shutdown),
|
||||
shutdown_timeout: DEFAULT_SHUTDOWN_TIMEOUT,
|
||||
store,
|
||||
worker_store,
|
||||
record,
|
||||
lease: Some(lease),
|
||||
}
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn worker_id(&self) -> WorkerId {
|
||||
self.record.worker_id
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn record(&self) -> &StandaloneWorkerRecord {
|
||||
&self.record
|
||||
}
|
||||
|
||||
/// Open one complete client-side Worker protocol session.
|
||||
///
|
||||
/// Working events, committed session entries, alert snapshots, and the
|
||||
/// initial history snapshot are merged behind the client boundary.
|
||||
pub fn connect(&self) -> Client<InProcessSocket> {
|
||||
let streams = subscribe_worker_protocol_session(&self.handle);
|
||||
let (socket, peer) = InProcessSocket::pair();
|
||||
tokio::spawn(run_protocol_session(self.handle.clone(), streams, peer));
|
||||
Client::new(socket)
|
||||
}
|
||||
|
||||
pub fn with_shutdown_timeout(mut self, shutdown_timeout: Duration) -> Self {
|
||||
self.shutdown_timeout = shutdown_timeout;
|
||||
self
|
||||
}
|
||||
|
||||
pub async fn shutdown(mut self) -> Result<(), StandaloneShutdownError> {
|
||||
let _ = self.handle.send(Method::Shutdown).await;
|
||||
let Some(shutdown) = self.shutdown.take() else {
|
||||
self.retain_lease();
|
||||
return Err(StandaloneShutdownError::ConfirmationLost);
|
||||
};
|
||||
match tokio::time::timeout(self.shutdown_timeout, shutdown).await {
|
||||
Ok(Ok(())) => {}
|
||||
Ok(Err(_)) => {
|
||||
self.retain_lease();
|
||||
return Err(StandaloneShutdownError::ConfirmationLost);
|
||||
}
|
||||
Err(_) => {
|
||||
self.retain_lease();
|
||||
return Err(StandaloneShutdownError::DeadlineExceeded);
|
||||
}
|
||||
}
|
||||
let active = match active_pointer(&self.worker_store, &self.record.storage_key) {
|
||||
Ok(active) => active,
|
||||
Err(_) => {
|
||||
self.retain_lease();
|
||||
return Err(StandaloneShutdownError::StateStore);
|
||||
}
|
||||
};
|
||||
if self
|
||||
.store
|
||||
.mark_stopped(
|
||||
&self.record,
|
||||
active.session_id,
|
||||
active.segment_id,
|
||||
StandaloneShutdownReason::UserExit,
|
||||
)
|
||||
.is_err()
|
||||
{
|
||||
self.retain_lease();
|
||||
return Err(StandaloneShutdownError::StateStore);
|
||||
}
|
||||
if let Some(lease) = self.lease.take() {
|
||||
lease
|
||||
.release()
|
||||
.map_err(|_| StandaloneShutdownError::StateStore)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn retain_lease(&mut self) {
|
||||
if let Some(lease) = self.lease.take() {
|
||||
lease.retain();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn run_protocol_session(
|
||||
handle: worker::WorkerHandle,
|
||||
streams: WorkerProtocolSessionStreams,
|
||||
mut peer: InProcessPeer,
|
||||
) {
|
||||
let WorkerProtocolSessionStreams {
|
||||
snapshot_event,
|
||||
mut log_entries,
|
||||
alert_snapshot,
|
||||
mut events,
|
||||
} = streams;
|
||||
|
||||
if !send_protocol_snapshot(&peer, alert_snapshot, snapshot_event).await {
|
||||
return;
|
||||
}
|
||||
|
||||
loop {
|
||||
tokio::select! {
|
||||
message = peer.next() => {
|
||||
let Some(message) = message else {
|
||||
return;
|
||||
};
|
||||
let Ok(method) = decode_method(&message) else {
|
||||
return;
|
||||
};
|
||||
if let Some(event) = dispatch_worker_protocol_method(&handle, method).await
|
||||
&& !send_protocol_event(&peer, event).await
|
||||
{
|
||||
return;
|
||||
}
|
||||
}
|
||||
event = events.recv() => {
|
||||
match event {
|
||||
Ok(event) => {
|
||||
if !send_protocol_event(&peer, event).await {
|
||||
return;
|
||||
}
|
||||
}
|
||||
Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {
|
||||
let replacement = subscribe_worker_protocol_session(&handle);
|
||||
let WorkerProtocolSessionStreams {
|
||||
snapshot_event,
|
||||
log_entries: replacement_log_entries,
|
||||
alert_snapshot,
|
||||
events: replacement_events,
|
||||
} = replacement;
|
||||
log_entries = replacement_log_entries;
|
||||
events = replacement_events;
|
||||
if !send_protocol_snapshot(&peer, alert_snapshot, snapshot_event).await {
|
||||
return;
|
||||
}
|
||||
}
|
||||
Err(tokio::sync::broadcast::error::RecvError::Closed) => return,
|
||||
}
|
||||
}
|
||||
entry = log_entries.recv() => {
|
||||
match entry {
|
||||
Ok(entry) => {
|
||||
if let Some(event) = live_log_entry_event(entry)
|
||||
&& !send_protocol_event(&peer, event).await
|
||||
{
|
||||
return;
|
||||
}
|
||||
}
|
||||
Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {
|
||||
let replacement = subscribe_worker_protocol_session(&handle);
|
||||
let WorkerProtocolSessionStreams {
|
||||
snapshot_event,
|
||||
log_entries: replacement_log_entries,
|
||||
alert_snapshot,
|
||||
events: replacement_events,
|
||||
} = replacement;
|
||||
log_entries = replacement_log_entries;
|
||||
events = replacement_events;
|
||||
if !send_protocol_snapshot(&peer, alert_snapshot, snapshot_event).await {
|
||||
return;
|
||||
}
|
||||
}
|
||||
Err(tokio::sync::broadcast::error::RecvError::Closed) => return,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn send_protocol_snapshot(
|
||||
peer: &InProcessPeer,
|
||||
alert_snapshot: Vec<protocol::Alert>,
|
||||
snapshot_event: Event,
|
||||
) -> bool {
|
||||
for alert in alert_snapshot {
|
||||
if !send_protocol_event(peer, Event::Alert(alert)).await {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
send_protocol_event(peer, snapshot_event).await
|
||||
}
|
||||
|
||||
async fn send_protocol_event(peer: &InProcessPeer, event: Event) -> bool {
|
||||
let Ok(message) = encode_event(&event) else {
|
||||
return false;
|
||||
};
|
||||
peer.send(message).await.is_ok()
|
||||
}
|
||||
|
||||
fn backing_store(
|
||||
store: &StandaloneWorkerStore,
|
||||
worker_id: WorkerId,
|
||||
) -> Result<(StandaloneBackingStore, FsWorkerStore), StandaloneStartupError> {
|
||||
let session_store = FsStore::new(store.sessions_dir(worker_id))
|
||||
.map_err(|_| StandaloneStartupError::StateStore)?;
|
||||
let worker_store = FsWorkerStore::new(store.worker_metadata_dir(worker_id))
|
||||
.map_err(|_| StandaloneStartupError::StateStore)?;
|
||||
Ok((
|
||||
CombinedStore::new(session_store, worker_store.clone()),
|
||||
worker_store,
|
||||
))
|
||||
}
|
||||
|
||||
fn active_pointer(
|
||||
worker_store: &FsWorkerStore,
|
||||
storage_key: &str,
|
||||
) -> Result<WorkerActiveSegmentRef, StandaloneStartupError> {
|
||||
worker_store
|
||||
.read_by_name(storage_key)
|
||||
.map_err(|_| StandaloneStartupError::StateStore)?
|
||||
.and_then(|metadata| metadata.active)
|
||||
.ok_or(StandaloneStartupError::StateStore)
|
||||
}
|
||||
|
||||
async fn stop_started_worker(started: BootstrappedWorker) {
|
||||
let _ = started.handle.send(Method::Shutdown).await;
|
||||
let _ = tokio::time::timeout(Duration::from_secs(2), started.shutdown).await;
|
||||
}
|
||||
|
||||
fn classify_store_startup_error(error: StandaloneStoreError) -> StandaloneStartupError {
|
||||
match error {
|
||||
StandaloneStoreError::WorkerLeased(_) => StandaloneStartupError::WorkerActive,
|
||||
StandaloneStoreError::LeaseLivenessUnknown(_) => {
|
||||
StandaloneStartupError::LeaseLivenessUnknown
|
||||
}
|
||||
StandaloneStoreError::CwdUnavailable(_)
|
||||
| StandaloneStoreError::CwdNotDirectory
|
||||
| StandaloneStoreError::CwdIdentityMismatch => {
|
||||
StandaloneStartupError::WorkingDirectoryUnavailable
|
||||
}
|
||||
_ => StandaloneStartupError::StateStore,
|
||||
}
|
||||
}
|
||||
|
||||
fn classify_startup_error(error: WorkerBootstrapError) -> StandaloneStartupError {
|
||||
match error {
|
||||
WorkerBootstrapError::Worker(WorkerError::ScopeLock(ScopeLockError::WriteConflict {
|
||||
competitor,
|
||||
rule,
|
||||
competitor_rule,
|
||||
})) => StandaloneStartupError::ScopeConflict {
|
||||
competitor,
|
||||
requested_rule: rule,
|
||||
competitor_rule,
|
||||
},
|
||||
WorkerBootstrapError::Worker(WorkerError::Provider(_)) => {
|
||||
StandaloneStartupError::ModelProvider
|
||||
}
|
||||
WorkerBootstrapError::Worker(_) => StandaloneStartupError::WorkerConfiguration,
|
||||
WorkerBootstrapError::Controller { source, .. }
|
||||
if source.kind() == std::io::ErrorKind::Other =>
|
||||
{
|
||||
StandaloneStartupError::FeatureComposition
|
||||
}
|
||||
WorkerBootstrapError::Controller { .. } => StandaloneStartupError::Controller,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use manifest::{
|
||||
ProfileExecutionTarget, ProfileResolveOptions, ProfileResolver, ProfileSelector,
|
||||
ResolvedProfile,
|
||||
};
|
||||
use thiserror::Error;
|
||||
use worker::PromptCatalogSource;
|
||||
|
||||
/// Process launch input resolved before any Worker/session side effect occurs.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct StandaloneLaunchConfig {
|
||||
pub cwd: PathBuf,
|
||||
pub state_dir: PathBuf,
|
||||
pub profile: ProfileSelector,
|
||||
pub worker_name: String,
|
||||
}
|
||||
|
||||
pub struct ResolvedStandaloneLaunch {
|
||||
pub cwd: PathBuf,
|
||||
pub state_dir: PathBuf,
|
||||
pub profile: ResolvedProfile,
|
||||
pub prompt_catalog: PromptCatalogSource,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
|
||||
pub enum StandaloneLaunchError {
|
||||
#[error("the standalone working directory is unavailable")]
|
||||
WorkingDirectoryUnavailable,
|
||||
#[error("path-based profiles are not standalone launch authority")]
|
||||
PathProfileUnsupported,
|
||||
#[error("the standalone profile could not be resolved")]
|
||||
ProfileResolutionFailed,
|
||||
}
|
||||
|
||||
impl StandaloneLaunchConfig {
|
||||
pub fn new(
|
||||
cwd: impl Into<PathBuf>,
|
||||
state_dir: impl Into<PathBuf>,
|
||||
profile: ProfileSelector,
|
||||
worker_name: impl Into<String>,
|
||||
) -> Self {
|
||||
Self {
|
||||
cwd: cwd.into(),
|
||||
state_dir: state_dir.into(),
|
||||
profile,
|
||||
worker_name: worker_name.into(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Resolve only built-in/XDG profile authority and bind standalone scope
|
||||
/// to the canonical process cwd. Repository-local profile discovery is
|
||||
/// deliberately not part of this path.
|
||||
pub fn resolve(self) -> Result<ResolvedStandaloneLaunch, StandaloneLaunchError> {
|
||||
if matches!(self.profile, ProfileSelector::Path { .. }) {
|
||||
return Err(StandaloneLaunchError::PathProfileUnsupported);
|
||||
}
|
||||
let cwd = canonical_directory(&self.cwd)?;
|
||||
let profile = ProfileResolver::new()
|
||||
.with_workspace_base(&cwd)
|
||||
.resolve_for_target(
|
||||
&self.profile,
|
||||
ProfileResolveOptions {
|
||||
worker_name: Some(self.worker_name),
|
||||
},
|
||||
ProfileExecutionTarget::Standalone,
|
||||
)
|
||||
.map_err(|_| StandaloneLaunchError::ProfileResolutionFailed)?;
|
||||
|
||||
Ok(ResolvedStandaloneLaunch {
|
||||
cwd,
|
||||
state_dir: self.state_dir,
|
||||
profile,
|
||||
prompt_catalog: PromptCatalogSource::builtins_only(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn canonical_directory(path: &Path) -> Result<PathBuf, StandaloneLaunchError> {
|
||||
let path = std::fs::canonicalize(path)
|
||||
.map_err(|_| StandaloneLaunchError::WorkingDirectoryUnavailable)?;
|
||||
if !path.is_dir() {
|
||||
return Err(StandaloneLaunchError::WorkingDirectoryUnavailable);
|
||||
}
|
||||
Ok(path)
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
//! In-process standalone host for one top-level Yoi Worker.
|
||||
//!
|
||||
//! The crate composes existing `worker`, `manifest`, `session-store`, and
|
||||
//! `workdir` contracts. It intentionally owns no TUI, Runtime, Workspace
|
||||
//! Server, HTTP, WebSocket, subprocess Worker, or alternative execution path.
|
||||
|
||||
pub mod host;
|
||||
pub mod launch;
|
||||
pub mod store;
|
||||
|
||||
pub use host::{StandaloneHost, StandaloneShutdownError, StandaloneStartupError};
|
||||
pub use launch::{ResolvedStandaloneLaunch, StandaloneLaunchConfig, StandaloneLaunchError};
|
||||
pub use protocol::WorkerId;
|
||||
pub use store::{
|
||||
StaleLeasePolicy, StandaloneCwdIdentity, StandaloneListScope, StandaloneShutdownReason,
|
||||
StandaloneStoreError, StandaloneWorkerRecord, StandaloneWorkerStatus, StandaloneWorkerStore,
|
||||
};
|
||||
@@ -0,0 +1,741 @@
|
||||
use std::fs::{self, File, OpenOptions};
|
||||
use std::io::{self, Write};
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use fs4::fs_std::FileExt;
|
||||
use manifest::WorkerManifest;
|
||||
use protocol::WorkerId;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use session_store::{SegmentId, SessionId};
|
||||
use thiserror::Error;
|
||||
use uuid::Uuid;
|
||||
|
||||
const RECORD_FILE: &str = "record.json";
|
||||
const COMMIT_MARKER: &str = "commit.pending";
|
||||
const LEASE_FILE: &str = "lease.json";
|
||||
const LEASE_LOCK_FILE: &str = "lease.lock";
|
||||
const SESSIONS_DIR: &str = "sessions";
|
||||
const WORKER_DIR: &str = "worker";
|
||||
const SCHEMA_VERSION: u32 = 1;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct StandaloneCwdIdentity {
|
||||
pub canonical_path: PathBuf,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub device: Option<u64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub inode: Option<u64>,
|
||||
}
|
||||
|
||||
impl StandaloneCwdIdentity {
|
||||
pub fn capture(path: impl AsRef<Path>) -> Result<Self, StandaloneStoreError> {
|
||||
let canonical_path =
|
||||
fs::canonicalize(path).map_err(StandaloneStoreError::CwdUnavailable)?;
|
||||
let metadata =
|
||||
fs::metadata(&canonical_path).map_err(StandaloneStoreError::CwdUnavailable)?;
|
||||
if !metadata.is_dir() {
|
||||
return Err(StandaloneStoreError::CwdNotDirectory);
|
||||
}
|
||||
#[cfg(unix)]
|
||||
let (device, inode) = {
|
||||
use std::os::unix::fs::MetadataExt;
|
||||
(Some(metadata.dev()), Some(metadata.ino()))
|
||||
};
|
||||
#[cfg(not(unix))]
|
||||
let (device, inode) = (None, None);
|
||||
Ok(Self {
|
||||
canonical_path,
|
||||
device,
|
||||
inode,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn verify(&self) -> Result<PathBuf, StandaloneStoreError> {
|
||||
let current = Self::capture(&self.canonical_path)?;
|
||||
if current != *self {
|
||||
return Err(StandaloneStoreError::CwdIdentityMismatch);
|
||||
}
|
||||
Ok(current.canonical_path)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum StandaloneWorkerStatus {
|
||||
Active,
|
||||
Stopped,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum StandaloneShutdownReason {
|
||||
UserExit,
|
||||
StartupFailed,
|
||||
ControllerError,
|
||||
ProcessInterrupted,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct StandaloneWorkerRecord {
|
||||
pub schema_version: u32,
|
||||
pub revision: u64,
|
||||
pub worker_id: WorkerId,
|
||||
/// User-facing Worker name resolved from the profile.
|
||||
pub worker_name: String,
|
||||
/// Internal key used by the current name-keyed Worker store.
|
||||
pub storage_key: String,
|
||||
pub cwd: StandaloneCwdIdentity,
|
||||
pub manifest: WorkerManifest,
|
||||
pub active_session_id: SessionId,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub active_segment_id: Option<SegmentId>,
|
||||
pub status: StandaloneWorkerStatus,
|
||||
pub created_at_unix_ms: u64,
|
||||
pub updated_at_unix_ms: u64,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub shutdown_reason: Option<StandaloneShutdownReason>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum StandaloneListScope {
|
||||
CurrentCwd,
|
||||
All,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum StaleLeasePolicy {
|
||||
Reject,
|
||||
Recover,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct StandaloneWorkerStore {
|
||||
root: PathBuf,
|
||||
}
|
||||
|
||||
impl StandaloneWorkerStore {
|
||||
pub fn open(root: impl Into<PathBuf>) -> Result<Self, StandaloneStoreError> {
|
||||
let root = root.into();
|
||||
fs::create_dir_all(&root).map_err(StandaloneStoreError::Io)?;
|
||||
if !fs::metadata(&root)
|
||||
.map_err(StandaloneStoreError::Io)?
|
||||
.is_dir()
|
||||
{
|
||||
return Err(StandaloneStoreError::NotDirectory);
|
||||
}
|
||||
Ok(Self { root })
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn root(&self) -> &Path {
|
||||
&self.root
|
||||
}
|
||||
|
||||
pub fn allocate(
|
||||
&self,
|
||||
cwd: impl AsRef<Path>,
|
||||
policy: StaleLeasePolicy,
|
||||
) -> Result<StandaloneWorkerAllocation, StandaloneStoreError> {
|
||||
let worker_id = WorkerId::now_v7();
|
||||
let cwd = StandaloneCwdIdentity::capture(cwd)?;
|
||||
let dir = self.worker_dir(worker_id);
|
||||
fs::create_dir(&dir).map_err(StandaloneStoreError::Io)?;
|
||||
fs::create_dir(dir.join(SESSIONS_DIR)).map_err(StandaloneStoreError::Io)?;
|
||||
fs::create_dir(dir.join(WORKER_DIR)).map_err(StandaloneStoreError::Io)?;
|
||||
let lease = self.acquire_lease(worker_id, policy)?;
|
||||
Ok(StandaloneWorkerAllocation {
|
||||
worker_id,
|
||||
cwd,
|
||||
lease,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn commit_created(
|
||||
&self,
|
||||
allocation: &StandaloneWorkerAllocation,
|
||||
manifest: WorkerManifest,
|
||||
storage_key: String,
|
||||
active_session_id: SessionId,
|
||||
active_segment_id: Option<SegmentId>,
|
||||
) -> Result<StandaloneWorkerRecord, StandaloneStoreError> {
|
||||
let now = now_unix_ms()?;
|
||||
let record = StandaloneWorkerRecord {
|
||||
schema_version: SCHEMA_VERSION,
|
||||
revision: 1,
|
||||
worker_id: allocation.worker_id,
|
||||
worker_name: manifest.worker.name.clone(),
|
||||
storage_key,
|
||||
cwd: allocation.cwd.clone(),
|
||||
manifest,
|
||||
active_session_id,
|
||||
active_segment_id,
|
||||
status: StandaloneWorkerStatus::Active,
|
||||
created_at_unix_ms: now,
|
||||
updated_at_unix_ms: now,
|
||||
shutdown_reason: None,
|
||||
};
|
||||
self.commit_record(None, &record)?;
|
||||
Ok(record)
|
||||
}
|
||||
|
||||
pub fn load(&self, id: WorkerId) -> Result<StandaloneWorkerRecord, StandaloneStoreError> {
|
||||
let dir = self.worker_dir(id);
|
||||
if dir.join(COMMIT_MARKER).exists() {
|
||||
return Err(StandaloneStoreError::IncompleteCommit(id));
|
||||
}
|
||||
let bytes = fs::read(dir.join(RECORD_FILE)).map_err(|error| {
|
||||
if error.kind() == io::ErrorKind::NotFound {
|
||||
StandaloneStoreError::WorkerNotFound(id)
|
||||
} else {
|
||||
StandaloneStoreError::Io(error)
|
||||
}
|
||||
})?;
|
||||
let record: StandaloneWorkerRecord = serde_json::from_slice(&bytes)
|
||||
.map_err(|source| StandaloneStoreError::CorruptRecord { id, source })?;
|
||||
if record.schema_version > SCHEMA_VERSION {
|
||||
return Err(StandaloneStoreError::NewerSchema {
|
||||
id,
|
||||
found: record.schema_version,
|
||||
supported: SCHEMA_VERSION,
|
||||
});
|
||||
}
|
||||
if record.schema_version != SCHEMA_VERSION || record.worker_id != id {
|
||||
return Err(StandaloneStoreError::InvalidRecord(id));
|
||||
}
|
||||
Ok(record)
|
||||
}
|
||||
|
||||
pub fn list(
|
||||
&self,
|
||||
cwd: impl AsRef<Path>,
|
||||
scope: StandaloneListScope,
|
||||
limit: usize,
|
||||
) -> Result<Vec<StandaloneWorkerRecord>, StandaloneStoreError> {
|
||||
let current_cwd = (scope == StandaloneListScope::CurrentCwd)
|
||||
.then(|| StandaloneCwdIdentity::capture(cwd))
|
||||
.transpose()?;
|
||||
let mut records = Vec::new();
|
||||
for entry in fs::read_dir(&self.root).map_err(StandaloneStoreError::Io)? {
|
||||
let entry = entry.map_err(StandaloneStoreError::Io)?;
|
||||
if !entry
|
||||
.file_type()
|
||||
.map_err(StandaloneStoreError::Io)?
|
||||
.is_dir()
|
||||
{
|
||||
continue;
|
||||
}
|
||||
let Ok(id) = entry.file_name().to_string_lossy().parse() else {
|
||||
continue;
|
||||
};
|
||||
let record = self.load(id)?;
|
||||
if current_cwd.as_ref().is_none_or(|cwd| &record.cwd == cwd) {
|
||||
records.push(record);
|
||||
}
|
||||
}
|
||||
records.sort_by(|left, right| {
|
||||
right
|
||||
.updated_at_unix_ms
|
||||
.cmp(&left.updated_at_unix_ms)
|
||||
.then_with(|| right.worker_id.to_string().cmp(&left.worker_id.to_string()))
|
||||
});
|
||||
records.truncate(limit);
|
||||
Ok(records)
|
||||
}
|
||||
|
||||
pub fn acquire_lease(
|
||||
&self,
|
||||
id: WorkerId,
|
||||
policy: StaleLeasePolicy,
|
||||
) -> Result<StandaloneWorkerLease, StandaloneStoreError> {
|
||||
let dir = self.worker_dir(id);
|
||||
let path = dir.join(LEASE_FILE);
|
||||
let _guard = LeaseMutationGuard::acquire(&dir)?;
|
||||
let lease = LeaseRecord::current()?;
|
||||
loop {
|
||||
match OpenOptions::new().write(true).create_new(true).open(&path) {
|
||||
Ok(mut file) => {
|
||||
serde_json::to_writer(&mut file, &lease).map_err(StandaloneStoreError::Json)?;
|
||||
file.write_all(b"\n").map_err(StandaloneStoreError::Io)?;
|
||||
file.sync_all().map_err(StandaloneStoreError::Io)?;
|
||||
sync_directory(&dir)?;
|
||||
return Ok(StandaloneWorkerLease {
|
||||
path,
|
||||
lease_id: lease.lease_id,
|
||||
released: false,
|
||||
});
|
||||
}
|
||||
Err(error) if error.kind() == io::ErrorKind::AlreadyExists => {
|
||||
let existing = read_lease(&path, id)?;
|
||||
match existing.liveness() {
|
||||
LeaseLiveness::Live => {
|
||||
return Err(StandaloneStoreError::WorkerLeased(id));
|
||||
}
|
||||
LeaseLiveness::Unknown => {
|
||||
return Err(StandaloneStoreError::LeaseLivenessUnknown(id));
|
||||
}
|
||||
LeaseLiveness::Stale => {}
|
||||
}
|
||||
if policy == StaleLeasePolicy::Reject {
|
||||
return Err(StandaloneStoreError::StaleLease(id));
|
||||
}
|
||||
fs::remove_file(&path).map_err(StandaloneStoreError::Io)?;
|
||||
sync_directory(&dir)?;
|
||||
}
|
||||
Err(error) => return Err(StandaloneStoreError::Io(error)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn update_active_pointer(
|
||||
&self,
|
||||
record: &StandaloneWorkerRecord,
|
||||
active_session_id: SessionId,
|
||||
active_segment_id: Option<SegmentId>,
|
||||
) -> Result<StandaloneWorkerRecord, StandaloneStoreError> {
|
||||
let mut next = record.clone();
|
||||
next.revision = next.revision.saturating_add(1);
|
||||
next.updated_at_unix_ms = now_unix_ms()?;
|
||||
next.active_session_id = active_session_id;
|
||||
next.active_segment_id = active_segment_id;
|
||||
next.status = StandaloneWorkerStatus::Active;
|
||||
next.shutdown_reason = None;
|
||||
self.commit_record(Some(record.revision), &next)?;
|
||||
Ok(next)
|
||||
}
|
||||
|
||||
pub fn mark_stopped(
|
||||
&self,
|
||||
record: &StandaloneWorkerRecord,
|
||||
active_session_id: SessionId,
|
||||
active_segment_id: Option<SegmentId>,
|
||||
reason: StandaloneShutdownReason,
|
||||
) -> Result<StandaloneWorkerRecord, StandaloneStoreError> {
|
||||
let mut next = record.clone();
|
||||
next.revision = next.revision.saturating_add(1);
|
||||
next.updated_at_unix_ms = now_unix_ms()?;
|
||||
next.active_session_id = active_session_id;
|
||||
next.active_segment_id = active_segment_id;
|
||||
next.status = StandaloneWorkerStatus::Stopped;
|
||||
next.shutdown_reason = Some(reason);
|
||||
self.commit_record(Some(record.revision), &next)?;
|
||||
Ok(next)
|
||||
}
|
||||
|
||||
pub fn delete(&self, id: WorkerId) -> Result<(), StandaloneStoreError> {
|
||||
let record = self.load(id)?;
|
||||
if record.status != StandaloneWorkerStatus::Stopped {
|
||||
return Err(StandaloneStoreError::DeleteActive(id));
|
||||
}
|
||||
let worker_dir = self.worker_dir(id);
|
||||
let _guard = LeaseMutationGuard::acquire(&worker_dir)?;
|
||||
let lease_path = worker_dir.join(LEASE_FILE);
|
||||
if lease_path.exists() {
|
||||
let lease = read_lease(&lease_path, id)?;
|
||||
return Err(match lease.liveness() {
|
||||
LeaseLiveness::Live => StandaloneStoreError::WorkerLeased(id),
|
||||
LeaseLiveness::Stale => StandaloneStoreError::StaleLease(id),
|
||||
LeaseLiveness::Unknown => StandaloneStoreError::LeaseLivenessUnknown(id),
|
||||
});
|
||||
}
|
||||
fs::remove_dir_all(self.worker_dir(id)).map_err(StandaloneStoreError::Io)?;
|
||||
sync_directory(&self.root)
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn sessions_dir(&self, id: WorkerId) -> PathBuf {
|
||||
self.worker_dir(id).join(SESSIONS_DIR)
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn worker_metadata_dir(&self, id: WorkerId) -> PathBuf {
|
||||
self.worker_dir(id).join(WORKER_DIR)
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub(crate) fn runtime_dir(&self, id: WorkerId) -> PathBuf {
|
||||
self.worker_dir(id).join("runtime")
|
||||
}
|
||||
|
||||
pub(crate) fn abandon_allocation(
|
||||
&self,
|
||||
allocation: StandaloneWorkerAllocation,
|
||||
) -> Result<(), StandaloneStoreError> {
|
||||
let worker_id = allocation.worker_id;
|
||||
allocation.lease.release()?;
|
||||
fs::remove_dir_all(self.worker_dir(worker_id)).map_err(StandaloneStoreError::Io)?;
|
||||
sync_directory(&self.root)
|
||||
}
|
||||
|
||||
fn commit_record(
|
||||
&self,
|
||||
expected_revision: Option<u64>,
|
||||
next: &StandaloneWorkerRecord,
|
||||
) -> Result<(), StandaloneStoreError> {
|
||||
let dir = self.worker_dir(next.worker_id);
|
||||
let marker = dir.join(COMMIT_MARKER);
|
||||
let mut marker_file = OpenOptions::new()
|
||||
.write(true)
|
||||
.create_new(true)
|
||||
.open(&marker)
|
||||
.map_err(|error| {
|
||||
if error.kind() == io::ErrorKind::AlreadyExists {
|
||||
StandaloneStoreError::IncompleteCommit(next.worker_id)
|
||||
} else {
|
||||
StandaloneStoreError::Io(error)
|
||||
}
|
||||
})?;
|
||||
writeln!(marker_file, "{}", next.revision).map_err(StandaloneStoreError::Io)?;
|
||||
marker_file.sync_all().map_err(StandaloneStoreError::Io)?;
|
||||
sync_directory(&dir)?;
|
||||
|
||||
if let Some(expected) = expected_revision {
|
||||
let current = self.load_record_while_committing(next.worker_id)?;
|
||||
if current.revision != expected {
|
||||
let _ = fs::remove_file(&marker);
|
||||
return Err(StandaloneStoreError::RevisionConflict {
|
||||
id: next.worker_id,
|
||||
expected,
|
||||
found: current.revision,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
let temporary = dir.join(format!("record.{}.tmp", Uuid::now_v7()));
|
||||
let result = (|| {
|
||||
let mut file = OpenOptions::new()
|
||||
.write(true)
|
||||
.create_new(true)
|
||||
.open(&temporary)
|
||||
.map_err(StandaloneStoreError::Io)?;
|
||||
serde_json::to_writer_pretty(&mut file, next).map_err(StandaloneStoreError::Json)?;
|
||||
file.write_all(b"\n").map_err(StandaloneStoreError::Io)?;
|
||||
file.sync_all().map_err(StandaloneStoreError::Io)?;
|
||||
fs::rename(&temporary, dir.join(RECORD_FILE)).map_err(StandaloneStoreError::Io)?;
|
||||
sync_directory(&dir)?;
|
||||
fs::remove_file(&marker).map_err(StandaloneStoreError::Io)?;
|
||||
sync_directory(&dir)
|
||||
})();
|
||||
if result.is_err() {
|
||||
let _ = fs::remove_file(&temporary);
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
fn load_record_while_committing(
|
||||
&self,
|
||||
id: WorkerId,
|
||||
) -> Result<StandaloneWorkerRecord, StandaloneStoreError> {
|
||||
let bytes =
|
||||
fs::read(self.worker_dir(id).join(RECORD_FILE)).map_err(StandaloneStoreError::Io)?;
|
||||
serde_json::from_slice(&bytes)
|
||||
.map_err(|source| StandaloneStoreError::CorruptRecord { id, source })
|
||||
}
|
||||
|
||||
fn worker_dir(&self, id: WorkerId) -> PathBuf {
|
||||
self.root.join(id.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct StandaloneWorkerAllocation {
|
||||
worker_id: WorkerId,
|
||||
cwd: StandaloneCwdIdentity,
|
||||
lease: StandaloneWorkerLease,
|
||||
}
|
||||
|
||||
impl StandaloneWorkerAllocation {
|
||||
#[must_use]
|
||||
pub fn worker_id(&self) -> WorkerId {
|
||||
self.worker_id
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn cwd(&self) -> &StandaloneCwdIdentity {
|
||||
&self.cwd
|
||||
}
|
||||
|
||||
pub fn into_lease(self) -> StandaloneWorkerLease {
|
||||
self.lease
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct StandaloneWorkerLease {
|
||||
path: PathBuf,
|
||||
lease_id: Uuid,
|
||||
released: bool,
|
||||
}
|
||||
|
||||
impl StandaloneWorkerLease {
|
||||
pub fn release(mut self) -> Result<(), StandaloneStoreError> {
|
||||
self.release_inner()
|
||||
}
|
||||
|
||||
pub(crate) fn retain(mut self) {
|
||||
self.released = true;
|
||||
}
|
||||
|
||||
fn release_inner(&mut self) -> Result<(), StandaloneStoreError> {
|
||||
if self.released {
|
||||
return Ok(());
|
||||
}
|
||||
if self.path.exists() {
|
||||
let parent = self.path.parent().expect("lease parent");
|
||||
let _guard = LeaseMutationGuard::acquire(parent)?;
|
||||
let bytes = fs::read(&self.path).map_err(StandaloneStoreError::Io)?;
|
||||
let current: LeaseRecord =
|
||||
serde_json::from_slice(&bytes).map_err(StandaloneStoreError::Json)?;
|
||||
if current.lease_id != self.lease_id {
|
||||
return Err(StandaloneStoreError::LeaseOwnershipLost);
|
||||
}
|
||||
fs::remove_file(&self.path).map_err(StandaloneStoreError::Io)?;
|
||||
sync_directory(self.path.parent().expect("lease parent"))?;
|
||||
}
|
||||
self.released = true;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for StandaloneWorkerLease {
|
||||
fn drop(&mut self) {
|
||||
let _ = self.release_inner();
|
||||
}
|
||||
}
|
||||
|
||||
struct LeaseMutationGuard {
|
||||
file: File,
|
||||
}
|
||||
|
||||
impl LeaseMutationGuard {
|
||||
fn acquire(dir: &Path) -> Result<Self, StandaloneStoreError> {
|
||||
let file = OpenOptions::new()
|
||||
.read(true)
|
||||
.write(true)
|
||||
.create(true)
|
||||
.truncate(false)
|
||||
.open(dir.join(LEASE_LOCK_FILE))
|
||||
.map_err(StandaloneStoreError::Io)?;
|
||||
file.lock_exclusive().map_err(StandaloneStoreError::Io)?;
|
||||
Ok(Self { file })
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for LeaseMutationGuard {
|
||||
fn drop(&mut self) {
|
||||
let _ = FileExt::unlock(&self.file);
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize)]
|
||||
struct LeaseRecord {
|
||||
lease_id: Uuid,
|
||||
pid: u32,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
process_start_marker: Option<u64>,
|
||||
acquired_at_unix_ms: u64,
|
||||
}
|
||||
|
||||
impl LeaseRecord {
|
||||
fn current() -> Result<Self, StandaloneStoreError> {
|
||||
Ok(Self {
|
||||
lease_id: Uuid::now_v7(),
|
||||
pid: std::process::id(),
|
||||
process_start_marker: match observe_process(std::process::id()) {
|
||||
ProcessObservation::Running { start_marker } => Some(start_marker),
|
||||
ProcessObservation::Missing | ProcessObservation::Unobservable => None,
|
||||
},
|
||||
acquired_at_unix_ms: now_unix_ms()?,
|
||||
})
|
||||
}
|
||||
|
||||
fn liveness(&self) -> LeaseLiveness {
|
||||
classify_lease_liveness(self.process_start_marker, observe_process(self.pid))
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum LeaseLiveness {
|
||||
Live,
|
||||
Stale,
|
||||
Unknown,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum ProcessObservation {
|
||||
Running { start_marker: u64 },
|
||||
Missing,
|
||||
Unobservable,
|
||||
}
|
||||
|
||||
fn classify_lease_liveness(
|
||||
recorded_start_marker: Option<u64>,
|
||||
observation: ProcessObservation,
|
||||
) -> LeaseLiveness {
|
||||
match (recorded_start_marker, observation) {
|
||||
(Some(recorded), ProcessObservation::Running { start_marker })
|
||||
if recorded == start_marker =>
|
||||
{
|
||||
LeaseLiveness::Live
|
||||
}
|
||||
(Some(_), ProcessObservation::Running { .. }) | (_, ProcessObservation::Missing) => {
|
||||
LeaseLiveness::Stale
|
||||
}
|
||||
(None, ProcessObservation::Running { .. }) | (_, ProcessObservation::Unobservable) => {
|
||||
LeaseLiveness::Unknown
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn read_lease(path: &Path, id: WorkerId) -> Result<LeaseRecord, StandaloneStoreError> {
|
||||
let bytes = fs::read(path).map_err(StandaloneStoreError::Io)?;
|
||||
serde_json::from_slice(&bytes)
|
||||
.map_err(|source| StandaloneStoreError::CorruptLease { id, source })
|
||||
}
|
||||
|
||||
#[cfg(target_os = "linux")]
|
||||
fn observe_process(pid: u32) -> ProcessObservation {
|
||||
let stat = match fs::read_to_string(format!("/proc/{pid}/stat")) {
|
||||
Ok(stat) => stat,
|
||||
Err(error) if error.kind() == io::ErrorKind::NotFound => {
|
||||
return if pid != std::process::id() && linux_proc_is_observable() {
|
||||
ProcessObservation::Missing
|
||||
} else {
|
||||
ProcessObservation::Unobservable
|
||||
};
|
||||
}
|
||||
Err(_) => return ProcessObservation::Unobservable,
|
||||
};
|
||||
parse_linux_process_start_marker(&stat)
|
||||
.map(|start_marker| ProcessObservation::Running { start_marker })
|
||||
.unwrap_or(ProcessObservation::Unobservable)
|
||||
}
|
||||
|
||||
#[cfg(target_os = "linux")]
|
||||
fn linux_proc_is_observable() -> bool {
|
||||
fs::read_to_string("/proc/self/stat")
|
||||
.ok()
|
||||
.and_then(|stat| parse_linux_process_start_marker(&stat))
|
||||
.is_some()
|
||||
}
|
||||
|
||||
#[cfg(target_os = "linux")]
|
||||
fn parse_linux_process_start_marker(stat: &str) -> Option<u64> {
|
||||
let (_, tail) = stat.rsplit_once(") ")?;
|
||||
tail.split_whitespace().nth(19)?.parse().ok()
|
||||
}
|
||||
|
||||
#[cfg(not(target_os = "linux"))]
|
||||
fn observe_process(pid: u32) -> ProcessObservation {
|
||||
if pid == std::process::id() {
|
||||
ProcessObservation::Running { start_marker: 0 }
|
||||
} else {
|
||||
ProcessObservation::Unobservable
|
||||
}
|
||||
}
|
||||
|
||||
fn now_unix_ms() -> Result<u64, StandaloneStoreError> {
|
||||
let duration = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map_err(|_| StandaloneStoreError::Clock)?;
|
||||
u64::try_from(duration.as_millis()).map_err(|_| StandaloneStoreError::Clock)
|
||||
}
|
||||
|
||||
fn sync_directory(path: &Path) -> Result<(), StandaloneStoreError> {
|
||||
File::open(path)
|
||||
.and_then(|file| file.sync_all())
|
||||
.map_err(StandaloneStoreError::Io)
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum StandaloneStoreError {
|
||||
#[error("standalone state path is not a directory")]
|
||||
NotDirectory,
|
||||
#[error("standalone cwd is unavailable")]
|
||||
CwdUnavailable(#[source] io::Error),
|
||||
#[error("standalone cwd is not a directory")]
|
||||
CwdNotDirectory,
|
||||
#[error("standalone cwd identity no longer matches the persisted Worker")]
|
||||
CwdIdentityMismatch,
|
||||
#[error("standalone Worker {0} was not found")]
|
||||
WorkerNotFound(WorkerId),
|
||||
#[error("standalone Worker {0} has an incomplete metadata commit")]
|
||||
IncompleteCommit(WorkerId),
|
||||
#[error("standalone Worker {0} has invalid metadata")]
|
||||
InvalidRecord(WorkerId),
|
||||
#[error("standalone Worker {id} metadata is corrupt")]
|
||||
CorruptRecord {
|
||||
id: WorkerId,
|
||||
#[source]
|
||||
source: serde_json::Error,
|
||||
},
|
||||
#[error("standalone Worker {id} lease is corrupt")]
|
||||
CorruptLease {
|
||||
id: WorkerId,
|
||||
#[source]
|
||||
source: serde_json::Error,
|
||||
},
|
||||
#[error("standalone Worker {id} uses schema {found}, newer than supported schema {supported}")]
|
||||
NewerSchema {
|
||||
id: WorkerId,
|
||||
found: u32,
|
||||
supported: u32,
|
||||
},
|
||||
#[error("standalone Worker {0} is already active")]
|
||||
WorkerLeased(WorkerId),
|
||||
#[error("standalone Worker {0} lease liveness cannot be proven; recovery is rejected")]
|
||||
LeaseLivenessUnknown(WorkerId),
|
||||
#[error("standalone Worker {0} has a stale lease; explicit recovery is required")]
|
||||
StaleLease(WorkerId),
|
||||
#[error("standalone Worker lease ownership changed")]
|
||||
LeaseOwnershipLost,
|
||||
#[error("standalone Worker {0} must be stopped before deletion")]
|
||||
DeleteActive(WorkerId),
|
||||
#[error(
|
||||
"standalone Worker {id} metadata revision changed (expected {expected}, found {found})"
|
||||
)]
|
||||
RevisionConflict {
|
||||
id: WorkerId,
|
||||
expected: u64,
|
||||
found: u64,
|
||||
},
|
||||
#[error("system clock is before the Unix epoch or out of range")]
|
||||
Clock,
|
||||
#[error("standalone metadata serialization failed")]
|
||||
Json(#[source] serde_json::Error),
|
||||
#[error("standalone state I/O failed")]
|
||||
Io(#[source] io::Error),
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{LeaseLiveness, ProcessObservation, classify_lease_liveness};
|
||||
|
||||
#[test]
|
||||
fn lease_liveness_requires_positive_live_or_stale_evidence() {
|
||||
assert_eq!(
|
||||
classify_lease_liveness(Some(41), ProcessObservation::Running { start_marker: 41 }),
|
||||
LeaseLiveness::Live
|
||||
);
|
||||
assert_eq!(
|
||||
classify_lease_liveness(Some(41), ProcessObservation::Running { start_marker: 42 }),
|
||||
LeaseLiveness::Stale
|
||||
);
|
||||
assert_eq!(
|
||||
classify_lease_liveness(Some(41), ProcessObservation::Missing),
|
||||
LeaseLiveness::Stale
|
||||
);
|
||||
assert_eq!(
|
||||
classify_lease_liveness(None, ProcessObservation::Running { start_marker: 41 }),
|
||||
LeaseLiveness::Unknown
|
||||
);
|
||||
assert_eq!(
|
||||
classify_lease_liveness(Some(41), ProcessObservation::Unobservable),
|
||||
LeaseLiveness::Unknown
|
||||
);
|
||||
assert_eq!(
|
||||
classify_lease_liveness(None, ProcessObservation::Unobservable),
|
||||
LeaseLiveness::Unknown
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,612 @@
|
||||
use std::collections::VecDeque;
|
||||
use std::pin::Pin;
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::Duration;
|
||||
|
||||
use agen::llm_client::client::LlmClient;
|
||||
use agen::llm_client::error::ClientError;
|
||||
use agen::llm_client::event::{Event as LlmEvent, StopReason};
|
||||
use agen::llm_client::types::Request;
|
||||
use async_trait::async_trait;
|
||||
use client::Client;
|
||||
use client::transport::in_process::Socket as InProcessSocket;
|
||||
use futures::{Stream, stream};
|
||||
use protocol::{Event, Method};
|
||||
use standalone::{
|
||||
StaleLeasePolicy, StandaloneHost, StandaloneLaunchConfig, StandaloneListScope,
|
||||
StandaloneStartupError, StandaloneStoreError, StandaloneWorkerStatus, StandaloneWorkerStore,
|
||||
};
|
||||
use uuid::Uuid;
|
||||
|
||||
#[derive(Clone)]
|
||||
struct ScriptedClient {
|
||||
responses: Arc<Mutex<VecDeque<Vec<LlmEvent>>>>,
|
||||
requests: Arc<Mutex<Vec<Request>>>,
|
||||
}
|
||||
|
||||
impl ScriptedClient {
|
||||
fn new(responses: Vec<Vec<LlmEvent>>) -> Self {
|
||||
Self {
|
||||
responses: Arc::new(Mutex::new(responses.into())),
|
||||
requests: Arc::new(Mutex::new(Vec::new())),
|
||||
}
|
||||
}
|
||||
|
||||
fn requests(&self) -> Vec<Request> {
|
||||
self.requests.lock().expect("requests lock").clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LlmClient for ScriptedClient {
|
||||
async fn stream(
|
||||
&self,
|
||||
request: Request,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<LlmEvent, ClientError>> + Send>>, ClientError>
|
||||
{
|
||||
self.requests.lock().expect("requests lock").push(request);
|
||||
let response = self
|
||||
.responses
|
||||
.lock()
|
||||
.expect("responses lock")
|
||||
.pop_front()
|
||||
.expect("scripted response");
|
||||
Ok(Box::pin(stream::iter(response.into_iter().map(Ok))))
|
||||
}
|
||||
|
||||
fn clone_boxed(&self) -> Box<dyn LlmClient> {
|
||||
Box::new(self.clone())
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn in_process_host_runs_text_and_read_tool_then_shuts_down() {
|
||||
let temp = tempfile::tempdir().expect("tempdir");
|
||||
std::fs::write(temp.path().join("probe.txt"), "standalone tool evidence\n")
|
||||
.expect("write probe");
|
||||
let worker_name = format!("standalone-{}", Uuid::now_v7());
|
||||
let launch = StandaloneLaunchConfig::new(
|
||||
temp.path(),
|
||||
temp.path().join("state"),
|
||||
manifest::ProfileSelector::Default,
|
||||
&worker_name,
|
||||
)
|
||||
.resolve()
|
||||
.expect("resolve standalone profile");
|
||||
|
||||
let client = ScriptedClient::new(vec![
|
||||
vec![
|
||||
LlmEvent::tool_use_start(0, "read-1", "Read"),
|
||||
LlmEvent::tool_input_delta(0, r#"{"file_path":"probe.txt"}"#),
|
||||
LlmEvent::tool_use_stop(0),
|
||||
],
|
||||
vec![
|
||||
LlmEvent::text_block_start(0),
|
||||
LlmEvent::text_delta(0, "standalone response"),
|
||||
LlmEvent::text_block_stop(0, Some(StopReason::EndTurn)),
|
||||
],
|
||||
]);
|
||||
let inspection = client.clone();
|
||||
let host = StandaloneHost::start_with_model_client(launch, client)
|
||||
.await
|
||||
.expect("start in-process host");
|
||||
assert_eq!(host.record().worker_name, worker_name);
|
||||
assert_eq!(host.record().manifest.worker.name, worker_name);
|
||||
assert_eq!(
|
||||
host.record().storage_key,
|
||||
format!("standalone-{}", host.worker_id())
|
||||
);
|
||||
let mut protocol_client = host.connect();
|
||||
|
||||
protocol_client
|
||||
.send(&Method::run_text("read the probe"))
|
||||
.await
|
||||
.expect("submit input");
|
||||
|
||||
tokio::time::timeout(Duration::from_secs(30), async {
|
||||
let mut saw_user_message = false;
|
||||
let mut saw_text = false;
|
||||
let mut saw_tool_result = false;
|
||||
loop {
|
||||
match protocol_client
|
||||
.next_event()
|
||||
.await
|
||||
.expect("protocol event")
|
||||
.expect("worker event")
|
||||
{
|
||||
Event::UserMessage { segments }
|
||||
if format!("{segments:?}").contains("read the probe") =>
|
||||
{
|
||||
saw_user_message = true;
|
||||
}
|
||||
Event::TextDelta { text } if text.contains("standalone response") => {
|
||||
saw_text = true;
|
||||
}
|
||||
Event::ToolResult { .. } => {
|
||||
saw_tool_result = true;
|
||||
}
|
||||
Event::RunEnd { .. } => {
|
||||
assert!(
|
||||
saw_user_message,
|
||||
"stream must expose the committed user message"
|
||||
);
|
||||
assert!(saw_text, "stream must expose the model text delta");
|
||||
assert!(saw_tool_result, "stream must expose the tool result");
|
||||
break;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("run completed");
|
||||
|
||||
let requests = inspection.requests();
|
||||
assert_eq!(requests.len(), 2);
|
||||
let tool_names = requests[0]
|
||||
.tools
|
||||
.iter()
|
||||
.map(|tool| tool.name.as_str())
|
||||
.collect::<Vec<_>>();
|
||||
assert!(tool_names.contains(&"Read"));
|
||||
assert!(tool_names.contains(&"TaskCreate"));
|
||||
assert!(tool_names.contains(&"SubWorkerSpawn"));
|
||||
assert!(format!("{:?}", requests[1].items).contains("standalone tool evidence"));
|
||||
assert!(
|
||||
!temp
|
||||
.path()
|
||||
.join("state/runtime")
|
||||
.join(&worker_name)
|
||||
.join("worker.sock")
|
||||
.exists()
|
||||
);
|
||||
|
||||
host.shutdown().await.expect("graceful shutdown");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn startup_preserves_occupied_scope_conflict_details() {
|
||||
let temp = tempfile::tempdir().expect("tempdir");
|
||||
let cwd = temp.path().join("project");
|
||||
std::fs::create_dir(&cwd).expect("create project");
|
||||
|
||||
let first_launch = StandaloneLaunchConfig::new(
|
||||
&cwd,
|
||||
temp.path().join("first-state"),
|
||||
manifest::ProfileSelector::Default,
|
||||
"first-worker",
|
||||
)
|
||||
.resolve()
|
||||
.expect("resolve first launch");
|
||||
let first_host =
|
||||
StandaloneHost::start_with_model_client(first_launch, ScriptedClient::new(Vec::new()))
|
||||
.await
|
||||
.expect("start first host");
|
||||
let competitor = first_host.record().storage_key.clone();
|
||||
|
||||
let second_launch = StandaloneLaunchConfig::new(
|
||||
&cwd,
|
||||
temp.path().join("second-state"),
|
||||
manifest::ProfileSelector::Default,
|
||||
"second-worker",
|
||||
)
|
||||
.resolve()
|
||||
.expect("resolve second launch");
|
||||
let error =
|
||||
StandaloneHost::start_with_model_client(second_launch, ScriptedClient::new(Vec::new()))
|
||||
.await
|
||||
.err()
|
||||
.expect("occupied scope rejected");
|
||||
|
||||
first_host.shutdown().await.expect("shutdown first host");
|
||||
|
||||
let canonical_cwd = cwd.canonicalize().expect("canonical cwd");
|
||||
match &error {
|
||||
StandaloneStartupError::ScopeConflict {
|
||||
competitor: actual_competitor,
|
||||
requested_rule,
|
||||
competitor_rule,
|
||||
} => {
|
||||
assert_eq!(actual_competitor, &competitor);
|
||||
assert_eq!(requested_rule.target, canonical_cwd);
|
||||
assert_eq!(competitor_rule.target, canonical_cwd);
|
||||
}
|
||||
other => panic!("expected scope conflict, got {other:?}"),
|
||||
}
|
||||
assert_eq!(
|
||||
error.to_string(),
|
||||
format!(
|
||||
"requested scope `{}` conflicts with worker allocation `{competitor}` rule `{}`",
|
||||
canonical_cwd.display(),
|
||||
canonical_cwd.display()
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn state_store_failure_is_redacted_and_starts_no_controller() {
|
||||
let temp = tempfile::tempdir().expect("tempdir");
|
||||
let state_path = temp.path().join("state-file-with-secret-name");
|
||||
std::fs::write(&state_path, "not a directory").expect("write blocking file");
|
||||
let launch = StandaloneLaunchConfig::new(
|
||||
temp.path(),
|
||||
&state_path,
|
||||
manifest::ProfileSelector::Default,
|
||||
format!("standalone-failure-{}", Uuid::now_v7()),
|
||||
)
|
||||
.resolve()
|
||||
.expect("resolve launch");
|
||||
let client = ScriptedClient::new(Vec::new());
|
||||
|
||||
let error = StandaloneHost::start_with_model_client(launch, client)
|
||||
.await
|
||||
.err()
|
||||
.expect("state store startup rejected");
|
||||
assert_eq!(error, standalone::StandaloneStartupError::StateStore);
|
||||
assert_eq!(
|
||||
error.to_string(),
|
||||
"the standalone state store could not be opened or validated"
|
||||
);
|
||||
assert!(!error.to_string().contains("secret-name"));
|
||||
assert!(
|
||||
!temp
|
||||
.path()
|
||||
.join("state-file-with-secret-name/runtime")
|
||||
.exists()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn standalone_crate_has_no_tui_runtime_or_workspace_server_dependency() {
|
||||
let manifest = include_str!("../Cargo.toml");
|
||||
let dependencies = manifest
|
||||
.split("[dependencies]")
|
||||
.nth(1)
|
||||
.expect("dependencies section")
|
||||
.split("[dev-dependencies]")
|
||||
.next()
|
||||
.expect("dependency body");
|
||||
for forbidden in ["tui", "worker-runtime", "yoi-workspace-server"] {
|
||||
assert!(
|
||||
!dependencies.lines().any(|line| {
|
||||
line.split_once('=')
|
||||
.is_some_and(|(name, _)| name.trim() == forbidden)
|
||||
}),
|
||||
"standalone must not depend on {forbidden}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn launch_rejects_path_profile_before_worker_startup() {
|
||||
let temp = tempfile::tempdir().expect("tempdir");
|
||||
let error = StandaloneLaunchConfig::new(
|
||||
temp.path(),
|
||||
temp.path().join("state"),
|
||||
manifest::ProfileSelector::Path {
|
||||
path: temp.path().join("profile.dcdl"),
|
||||
},
|
||||
"standalone-path-profile",
|
||||
)
|
||||
.resolve()
|
||||
.err()
|
||||
.expect("path profile rejected");
|
||||
assert_eq!(
|
||||
error,
|
||||
standalone::StandaloneLaunchError::PathProfileUnsupported
|
||||
);
|
||||
}
|
||||
|
||||
type TestResult = Result<(), Box<dyn std::error::Error>>;
|
||||
|
||||
#[tokio::test]
|
||||
async fn standalone_restore_preserves_history_tasks_notifications_and_cwd_scope() -> TestResult {
|
||||
let temp = tempfile::tempdir()?;
|
||||
let cwd = temp.path().join("project");
|
||||
let state_dir = temp.path().join("client").join("standalone-workers");
|
||||
std::fs::create_dir_all(&cwd)?;
|
||||
let launch = StandaloneLaunchConfig::new(
|
||||
&cwd,
|
||||
&state_dir,
|
||||
manifest::ProfileSelector::Default,
|
||||
"display-name-is-not-session-identity",
|
||||
)
|
||||
.resolve()?;
|
||||
let first_client = ScriptedClient::new(vec![
|
||||
vec![
|
||||
LlmEvent::tool_use_start(0, "task-1", "TaskCreate"),
|
||||
LlmEvent::tool_input_delta(
|
||||
0,
|
||||
r#"{"subject":"persisted task","description":"survives restore"}"#,
|
||||
),
|
||||
LlmEvent::tool_use_stop(0),
|
||||
],
|
||||
vec![
|
||||
LlmEvent::text_block_start(0),
|
||||
LlmEvent::text_delta(0, "first answer"),
|
||||
LlmEvent::text_block_stop(0, Some(StopReason::EndTurn)),
|
||||
],
|
||||
vec![
|
||||
LlmEvent::text_block_start(0),
|
||||
LlmEvent::text_delta(0, "notification acknowledged"),
|
||||
LlmEvent::text_block_stop(0, Some(StopReason::EndTurn)),
|
||||
],
|
||||
]);
|
||||
let host = StandaloneHost::start_with_model_client(launch, first_client).await?;
|
||||
let worker_id = host.worker_id();
|
||||
let mut protocol_client = host.connect();
|
||||
protocol_client
|
||||
.send(&Method::run_text("first request"))
|
||||
.await?;
|
||||
wait_for_run_end(&mut protocol_client).await?;
|
||||
protocol_client
|
||||
.send(&Method::Notify {
|
||||
message: "persisted notification".to_string(),
|
||||
auto_run: true,
|
||||
})
|
||||
.await?;
|
||||
wait_for_run_end(&mut protocol_client).await?;
|
||||
host.shutdown().await?;
|
||||
|
||||
let store = StandaloneWorkerStore::open(&state_dir)?;
|
||||
let current = store.list(&cwd, StandaloneListScope::CurrentCwd, 100)?;
|
||||
assert_eq!(current.len(), 1);
|
||||
assert_eq!(current[0].worker_id, worker_id);
|
||||
assert_eq!(current[0].status, StandaloneWorkerStatus::Stopped);
|
||||
let other_cwd = temp.path().join("other");
|
||||
std::fs::create_dir(&other_cwd)?;
|
||||
assert!(
|
||||
store
|
||||
.list(&other_cwd, StandaloneListScope::CurrentCwd, 100)?
|
||||
.is_empty()
|
||||
);
|
||||
assert_eq!(
|
||||
store.list(&other_cwd, StandaloneListScope::All, 100)?.len(),
|
||||
1
|
||||
);
|
||||
|
||||
let second_client = ScriptedClient::new(vec![vec![
|
||||
LlmEvent::text_block_start(0),
|
||||
LlmEvent::text_delta(0, "second answer"),
|
||||
LlmEvent::text_block_stop(0, Some(StopReason::EndTurn)),
|
||||
]]);
|
||||
let second_inspection = second_client.clone();
|
||||
let host =
|
||||
StandaloneHost::restore_with_model_client(state_dir.clone(), worker_id, second_client)
|
||||
.await?;
|
||||
assert_eq!(
|
||||
host.record().worker_name,
|
||||
"display-name-is-not-session-identity"
|
||||
);
|
||||
assert_eq!(host.record().storage_key, format!("standalone-{worker_id}"));
|
||||
let mut protocol_client = host.connect();
|
||||
let snapshot = format!(
|
||||
"{:?}",
|
||||
protocol_client
|
||||
.next_event()
|
||||
.await
|
||||
.expect("restored protocol stream")
|
||||
.expect("restored snapshot")
|
||||
);
|
||||
assert!(snapshot.contains("first request"), "{snapshot}");
|
||||
assert!(snapshot.contains("first answer"), "{snapshot}");
|
||||
assert!(snapshot.contains("persisted task"), "{snapshot}");
|
||||
assert!(snapshot.contains("persisted notification"), "{snapshot}");
|
||||
|
||||
protocol_client
|
||||
.send(&Method::run_text("continue after restore"))
|
||||
.await?;
|
||||
wait_for_run_end(&mut protocol_client).await?;
|
||||
let request = second_inspection
|
||||
.requests()
|
||||
.into_iter()
|
||||
.next()
|
||||
.expect("restored run request");
|
||||
let projected = format!("{:?}", request.items);
|
||||
assert!(projected.contains("first answer"), "{projected}");
|
||||
assert!(projected.contains("persisted notification"), "{projected}");
|
||||
assert!(projected.contains("persisted task"), "{projected}");
|
||||
host.shutdown().await?;
|
||||
|
||||
store.delete(worker_id)?;
|
||||
assert!(cwd.exists(), "deleting session state must not mutate cwd");
|
||||
assert!(matches!(
|
||||
store.load(worker_id),
|
||||
Err(StandaloneStoreError::WorkerNotFound(_))
|
||||
));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn standalone_restore_rejects_concurrent_lease_and_missing_cwd() -> TestResult {
|
||||
let temp = tempfile::tempdir()?;
|
||||
let cwd = temp.path().join("project");
|
||||
let moved = temp.path().join("moved-project");
|
||||
let state_dir = temp.path().join("state");
|
||||
std::fs::create_dir(&cwd)?;
|
||||
let launch = StandaloneLaunchConfig::new(
|
||||
&cwd,
|
||||
&state_dir,
|
||||
manifest::ProfileSelector::Default,
|
||||
"standalone-lease-test",
|
||||
)
|
||||
.resolve()?;
|
||||
let host =
|
||||
StandaloneHost::start_with_model_client(launch, ScriptedClient::new(Vec::new())).await?;
|
||||
let worker_id = host.worker_id();
|
||||
let store = StandaloneWorkerStore::open(&state_dir)?;
|
||||
assert!(matches!(
|
||||
store.acquire_lease(worker_id, StaleLeasePolicy::Recover),
|
||||
Err(StandaloneStoreError::WorkerLeased(id)) if id == worker_id
|
||||
));
|
||||
let restore = StandaloneHost::restore_with_model_client(
|
||||
state_dir.clone(),
|
||||
worker_id,
|
||||
ScriptedClient::new(Vec::new()),
|
||||
)
|
||||
.await;
|
||||
assert!(matches!(restore, Err(StandaloneStartupError::WorkerActive)));
|
||||
host.shutdown().await?;
|
||||
|
||||
std::fs::rename(&cwd, &moved)?;
|
||||
let restore = StandaloneHost::restore_with_model_client(
|
||||
state_dir,
|
||||
worker_id,
|
||||
ScriptedClient::new(Vec::new()),
|
||||
)
|
||||
.await;
|
||||
assert!(matches!(
|
||||
restore,
|
||||
Err(StandaloneStartupError::WorkingDirectoryUnavailable)
|
||||
));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn standalone_restore_recovers_only_a_proven_stale_lease() -> TestResult {
|
||||
let temp = tempfile::tempdir()?;
|
||||
let state_dir = temp.path().join("state");
|
||||
let mut launch = StandaloneLaunchConfig::new(
|
||||
temp.path(),
|
||||
&state_dir,
|
||||
manifest::ProfileSelector::Default,
|
||||
"standalone-stale-lease-test",
|
||||
)
|
||||
.resolve()?;
|
||||
launch.profile.manifest.profile = Some(manifest::ProfileManifestSnapshot {
|
||||
source: manifest::ProfileSource::Registry {
|
||||
source: manifest::ProfileRegistrySource::User,
|
||||
name: "user-standalone".to_string(),
|
||||
path: None,
|
||||
provenance: Some("user-config-revision-7".to_string()),
|
||||
},
|
||||
profile: Some(manifest::ProfileMetadata {
|
||||
name: Some("User standalone".to_string()),
|
||||
description: None,
|
||||
format: None,
|
||||
}),
|
||||
});
|
||||
let host =
|
||||
StandaloneHost::start_with_model_client(launch, ScriptedClient::new(Vec::new())).await?;
|
||||
let worker_id = host.worker_id();
|
||||
host.shutdown().await?;
|
||||
let store = StandaloneWorkerStore::open(&state_dir)?;
|
||||
assert!(matches!(
|
||||
store.load(worker_id)?.manifest.profile,
|
||||
Some(manifest::ProfileManifestSnapshot {
|
||||
source: manifest::ProfileSource::Registry {
|
||||
source: manifest::ProfileRegistrySource::User,
|
||||
..
|
||||
},
|
||||
..
|
||||
})
|
||||
));
|
||||
let worker_dir = state_dir.join(worker_id.to_string());
|
||||
std::fs::write(
|
||||
worker_dir.join("lease.json"),
|
||||
serde_json::to_vec(&serde_json::json!({
|
||||
"lease_id": uuid::Uuid::now_v7(),
|
||||
"pid": u32::MAX,
|
||||
"process_start_marker": 1,
|
||||
"acquired_at_unix_ms": 1
|
||||
}))?,
|
||||
)?;
|
||||
|
||||
let host = StandaloneHost::restore_with_model_client(
|
||||
state_dir,
|
||||
worker_id,
|
||||
ScriptedClient::new(Vec::new()),
|
||||
)
|
||||
.await?;
|
||||
host.shutdown().await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn standalone_restore_rejects_lease_with_missing_start_marker() -> TestResult {
|
||||
let temp = tempfile::tempdir()?;
|
||||
let state_dir = temp.path().join("state");
|
||||
let launch = StandaloneLaunchConfig::new(
|
||||
temp.path(),
|
||||
&state_dir,
|
||||
manifest::ProfileSelector::Default,
|
||||
"standalone-unknown-lease-test",
|
||||
)
|
||||
.resolve()?;
|
||||
let host =
|
||||
StandaloneHost::start_with_model_client(launch, ScriptedClient::new(Vec::new())).await?;
|
||||
let worker_id = host.worker_id();
|
||||
host.shutdown().await?;
|
||||
let worker_dir = state_dir.join(worker_id.to_string());
|
||||
std::fs::write(
|
||||
worker_dir.join("lease.json"),
|
||||
serde_json::to_vec(&serde_json::json!({
|
||||
"lease_id": uuid::Uuid::now_v7(),
|
||||
"pid": std::process::id(),
|
||||
"acquired_at_unix_ms": 1
|
||||
}))?,
|
||||
)?;
|
||||
|
||||
let store = StandaloneWorkerStore::open(&state_dir)?;
|
||||
assert!(matches!(
|
||||
store.acquire_lease(worker_id, StaleLeasePolicy::Recover),
|
||||
Err(StandaloneStoreError::LeaseLivenessUnknown(id)) if id == worker_id
|
||||
));
|
||||
let restore = StandaloneHost::restore_with_model_client(
|
||||
state_dir,
|
||||
worker_id,
|
||||
ScriptedClient::new(Vec::new()),
|
||||
)
|
||||
.await;
|
||||
assert!(matches!(
|
||||
restore,
|
||||
Err(StandaloneStartupError::LeaseLivenessUnknown)
|
||||
));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn standalone_metadata_fails_closed_on_incomplete_or_newer_records() -> TestResult {
|
||||
let temp = tempfile::tempdir()?;
|
||||
let state_dir = temp.path().join("state");
|
||||
let launch = StandaloneLaunchConfig::new(
|
||||
temp.path(),
|
||||
&state_dir,
|
||||
manifest::ProfileSelector::Default,
|
||||
"standalone-schema-test",
|
||||
)
|
||||
.resolve()?;
|
||||
let host =
|
||||
StandaloneHost::start_with_model_client(launch, ScriptedClient::new(Vec::new())).await?;
|
||||
let worker_id = host.worker_id();
|
||||
host.shutdown().await?;
|
||||
let store = StandaloneWorkerStore::open(&state_dir)?;
|
||||
let worker_dir = state_dir.join(worker_id.to_string());
|
||||
std::fs::write(worker_dir.join("commit.pending"), b"interrupted\n")?;
|
||||
assert!(matches!(
|
||||
store.load(worker_id),
|
||||
Err(StandaloneStoreError::IncompleteCommit(id)) if id == worker_id
|
||||
));
|
||||
std::fs::remove_file(worker_dir.join("commit.pending"))?;
|
||||
let record_path = worker_dir.join("record.json");
|
||||
let mut record: serde_json::Value = serde_json::from_slice(&std::fs::read(&record_path)?)?;
|
||||
record["schema_version"] = serde_json::json!(u32::MAX);
|
||||
std::fs::write(&record_path, serde_json::to_vec_pretty(&record)?)?;
|
||||
assert!(matches!(
|
||||
store.load(worker_id),
|
||||
Err(StandaloneStoreError::NewerSchema { id, .. }) if id == worker_id
|
||||
));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn wait_for_run_end(client: &mut Client<InProcessSocket>) -> TestResult {
|
||||
tokio::time::timeout(Duration::from_secs(10), async {
|
||||
loop {
|
||||
if matches!(client.next_event().await, Ok(Some(Event::RunEnd { .. }))) {
|
||||
break;
|
||||
}
|
||||
}
|
||||
})
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,122 @@
|
||||
-- Canonical standalone Ticket schema. Workspace Server composes stricter cross-domain authority.
|
||||
CREATE TABLE typed_ticket_artifacts (
|
||||
workspace_id TEXT NOT NULL, ticket_id TEXT NOT NULL, relative_path TEXT NOT NULL, content BLOB NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, relative_path),
|
||||
FOREIGN KEY (workspace_id, ticket_id) REFERENCES typed_tickets(workspace_id, ticket_id) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE typed_ticket_event_attributes (
|
||||
workspace_id TEXT NOT NULL, ticket_id TEXT NOT NULL, event_index INTEGER NOT NULL, key TEXT NOT NULL, value TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, event_index, key),
|
||||
FOREIGN KEY (workspace_id, ticket_id, event_index) REFERENCES typed_ticket_events(workspace_id, ticket_id, event_index) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE typed_ticket_event_references (
|
||||
workspace_id TEXT NOT NULL, ticket_id TEXT NOT NULL, event_index INTEGER NOT NULL, ordinal INTEGER NOT NULL, kind TEXT NOT NULL, target TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, event_index, ordinal),
|
||||
FOREIGN KEY (workspace_id, ticket_id, event_index) REFERENCES typed_ticket_events(workspace_id, ticket_id, event_index) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE typed_ticket_events (
|
||||
workspace_id TEXT NOT NULL,
|
||||
ticket_id TEXT NOT NULL,
|
||||
event_index INTEGER NOT NULL,
|
||||
kind TEXT NOT NULL,
|
||||
author TEXT,
|
||||
at TEXT,
|
||||
status TEXT,
|
||||
from_state TEXT,
|
||||
to_state TEXT,
|
||||
reason TEXT,
|
||||
state_field TEXT,
|
||||
heading TEXT,
|
||||
body TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, event_index),
|
||||
FOREIGN KEY (workspace_id, ticket_id) REFERENCES typed_tickets(workspace_id, ticket_id) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE typed_ticket_labels (
|
||||
workspace_id TEXT NOT NULL, ticket_id TEXT NOT NULL, ordinal INTEGER NOT NULL, label TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, ordinal),
|
||||
FOREIGN KEY (workspace_id, ticket_id) REFERENCES typed_tickets(workspace_id, ticket_id) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE typed_ticket_orchestration_plans (
|
||||
workspace_id TEXT NOT NULL,
|
||||
ticket_id TEXT NOT NULL,
|
||||
record_id TEXT NOT NULL,
|
||||
kind TEXT NOT NULL,
|
||||
related_ticket TEXT,
|
||||
note TEXT,
|
||||
accepted_summary TEXT,
|
||||
accepted_branch TEXT,
|
||||
accepted_worktree TEXT,
|
||||
accepted_role_plan TEXT,
|
||||
author TEXT NOT NULL,
|
||||
at TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, record_id),
|
||||
FOREIGN KEY (workspace_id, ticket_id) REFERENCES typed_tickets(workspace_id, ticket_id) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE typed_ticket_raw_frontmatter (
|
||||
workspace_id TEXT NOT NULL, ticket_id TEXT NOT NULL, key TEXT NOT NULL, value TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, key),
|
||||
FOREIGN KEY (workspace_id, ticket_id) REFERENCES typed_tickets(workspace_id, ticket_id) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE typed_ticket_relations (
|
||||
workspace_id TEXT NOT NULL, ticket_id TEXT NOT NULL, kind TEXT NOT NULL, target TEXT NOT NULL, note TEXT, author TEXT NOT NULL, at TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, kind, target),
|
||||
FOREIGN KEY (workspace_id, ticket_id) REFERENCES typed_tickets(workspace_id, ticket_id) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE typed_ticket_risk_flags (
|
||||
workspace_id TEXT NOT NULL, ticket_id TEXT NOT NULL, ordinal INTEGER NOT NULL, risk_flag TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, ordinal),
|
||||
FOREIGN KEY (workspace_id, ticket_id) REFERENCES typed_tickets(workspace_id, ticket_id) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE typed_tickets (
|
||||
workspace_id TEXT NOT NULL,
|
||||
ticket_id TEXT NOT NULL,
|
||||
slug TEXT NOT NULL,
|
||||
title TEXT NOT NULL,
|
||||
status TEXT NOT NULL,
|
||||
kind TEXT NOT NULL,
|
||||
priority TEXT NOT NULL,
|
||||
body TEXT NOT NULL,
|
||||
created_at TEXT,
|
||||
updated_at TEXT,
|
||||
assignee TEXT,
|
||||
readiness TEXT,
|
||||
workflow_state TEXT NOT NULL,
|
||||
workflow_state_explicit INTEGER NOT NULL,
|
||||
queued_by TEXT,
|
||||
queued_at TEXT,
|
||||
resolution TEXT, repository_id TEXT, ref_selector TEXT,
|
||||
PRIMARY KEY (workspace_id, ticket_id)
|
||||
);
|
||||
CREATE TABLE "workspace_resource_key_counters" (
|
||||
workspace_id TEXT NOT NULL,
|
||||
resource_kind TEXT NOT NULL CHECK (resource_kind IN ('ticket', 'objective', 'worker')),
|
||||
next_sequence INTEGER NOT NULL CHECK (next_sequence > 0),
|
||||
PRIMARY KEY (workspace_id, resource_kind)
|
||||
);
|
||||
CREATE TABLE "workspace_resource_keys" (
|
||||
workspace_id TEXT NOT NULL,
|
||||
resource_kind TEXT NOT NULL CHECK (resource_kind IN ('ticket', 'objective', 'worker')),
|
||||
resource_id TEXT NOT NULL,
|
||||
sequence INTEGER NOT NULL CHECK (sequence > 0),
|
||||
resource_key TEXT NOT NULL,
|
||||
allocated_at TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, resource_kind, resource_id),
|
||||
UNIQUE (workspace_id, resource_kind, sequence),
|
||||
UNIQUE (workspace_id, resource_key)
|
||||
);
|
||||
CREATE INDEX idx_workspace_resource_keys_reverse
|
||||
ON workspace_resource_keys(workspace_id, resource_kind, resource_key);
|
||||
CREATE INDEX typed_ticket_events_workspace_kind_ticket
|
||||
ON typed_ticket_events(workspace_id, kind, ticket_id, event_index);
|
||||
CREATE INDEX typed_ticket_relations_workspace_source_kind
|
||||
ON typed_ticket_relations(workspace_id, ticket_id, kind, target);
|
||||
CREATE INDEX typed_ticket_relations_workspace_target_kind
|
||||
ON typed_ticket_relations(workspace_id, target, kind, ticket_id);
|
||||
CREATE INDEX typed_tickets_workspace_created
|
||||
ON typed_tickets(workspace_id, created_at DESC, ticket_id);
|
||||
CREATE INDEX typed_tickets_workspace_state_updated
|
||||
ON typed_tickets(workspace_id, workflow_state, updated_at DESC, ticket_id);
|
||||
CREATE INDEX typed_tickets_workspace_title
|
||||
ON typed_tickets(workspace_id, title COLLATE NOCASE, ticket_id);
|
||||
CREATE INDEX typed_tickets_workspace_updated
|
||||
ON typed_tickets(workspace_id, updated_at DESC, ticket_id);
|
||||
@@ -26,11 +26,7 @@ pub mod config;
|
||||
mod sqlite_schema;
|
||||
pub mod tool;
|
||||
|
||||
pub use sqlite_schema::{
|
||||
LATEST_SQLITE_TICKET_SCHEMA_VERSION, migrate_sqlite_ticket_resource_key_schema_in_transaction,
|
||||
migrate_sqlite_ticket_schema, migrate_sqlite_ticket_schema_through,
|
||||
verify_sqlite_ticket_schema,
|
||||
};
|
||||
pub use sqlite_schema::{migrate_sqlite_ticket_schema, verify_sqlite_ticket_schema};
|
||||
|
||||
const REQUIRED_FIELDS: [&str; 4] = ["title", "state", "created_at", "updated_at"];
|
||||
const MAX_STATE_CHANGE_REASON_BYTES: usize = 1024;
|
||||
@@ -489,6 +485,7 @@ pub struct NewTicket {
|
||||
pub workflow_state: Option<TicketWorkflowState>,
|
||||
pub queued_by: Option<String>,
|
||||
pub queued_at: Option<String>,
|
||||
#[serde(rename = "repository_key")]
|
||||
pub repository_id: Option<String>,
|
||||
pub ref_selector: Option<String>,
|
||||
}
|
||||
@@ -519,6 +516,7 @@ impl NewTicket {
|
||||
#[serde(tag = "action", rename_all = "snake_case")]
|
||||
pub enum TicketTargetEdit {
|
||||
Set {
|
||||
#[serde(rename = "repository_key")]
|
||||
repository_id: String,
|
||||
ref_selector: Option<String>,
|
||||
},
|
||||
@@ -1610,6 +1608,7 @@ pub struct TicketMeta {
|
||||
pub workflow_state_explicit: bool,
|
||||
pub queued_by: Option<String>,
|
||||
pub queued_at: Option<String>,
|
||||
#[serde(rename = "repository_key")]
|
||||
pub repository_id: Option<String>,
|
||||
pub ref_selector: Option<String>,
|
||||
pub raw: BTreeMap<String, String>,
|
||||
@@ -2573,7 +2572,7 @@ impl SqliteTicketBackend {
|
||||
}
|
||||
}
|
||||
|
||||
/// Opens a standalone Ticket backend, applying all Ticket-owned migrations once.
|
||||
/// Opens a standalone Ticket backend at the current canonical schema baseline.
|
||||
pub fn open(db_path: impl Into<PathBuf>, workspace_id: impl Into<String>) -> Result<Self> {
|
||||
let backend = Self::configured(db_path, workspace_id);
|
||||
let connection = backend.connect()?;
|
||||
|
||||
@@ -7,7 +7,7 @@ use crate::{Result, TicketError, sqlite_err};
|
||||
|
||||
const MIGRATION_TABLE: &str = "ticket_schema_migrations";
|
||||
const MAX_SCHEMA_DIAGNOSTICS: usize = 32;
|
||||
pub const LATEST_SQLITE_TICKET_SCHEMA_VERSION: i64 = 6;
|
||||
const LATEST_SQLITE_TICKET_SCHEMA_VERSION: i64 = 6;
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
struct Migration {
|
||||
@@ -16,38 +16,11 @@ struct Migration {
|
||||
apply: fn(&Connection) -> Result<()>,
|
||||
}
|
||||
|
||||
const MIGRATIONS: &[Migration] = &[
|
||||
Migration {
|
||||
version: 1,
|
||||
name: "create_typed_ticket_tables",
|
||||
apply: create_typed_ticket_tables,
|
||||
},
|
||||
Migration {
|
||||
version: 2,
|
||||
name: "add_ticket_repository_target",
|
||||
apply: add_ticket_repository_target,
|
||||
},
|
||||
Migration {
|
||||
version: 3,
|
||||
name: "convert_legacy_reviews_to_comments",
|
||||
apply: retire_legacy_ticket_review_events,
|
||||
},
|
||||
Migration {
|
||||
version: 4,
|
||||
name: "add_ticket_query_indexes",
|
||||
apply: add_ticket_query_indexes,
|
||||
},
|
||||
Migration {
|
||||
version: 5,
|
||||
name: "add_workspace_human_keys",
|
||||
apply: add_workspace_human_keys,
|
||||
},
|
||||
Migration {
|
||||
version: 6,
|
||||
name: "rename_workspace_resource_keys",
|
||||
apply: rename_workspace_resource_keys,
|
||||
},
|
||||
];
|
||||
const MIGRATIONS: &[Migration] = &[Migration {
|
||||
version: LATEST_SQLITE_TICKET_SCHEMA_VERSION,
|
||||
name: "ticket schema baseline",
|
||||
apply: create_latest_ticket_schema,
|
||||
}];
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
struct ExpectedColumn {
|
||||
@@ -258,30 +231,12 @@ const fn column(
|
||||
}
|
||||
}
|
||||
|
||||
/// Applies the Ticket crate's SQLite migrations and verifies the resulting schema.
|
||||
/// Creates and verifies the Ticket crate's latest SQLite schema.
|
||||
///
|
||||
/// This is a startup/standalone-open operation. Normal Ticket request handling must
|
||||
/// use [`verify_sqlite_ticket_schema`] instead, so request paths never acquire DDL
|
||||
/// authority.
|
||||
pub fn migrate_sqlite_ticket_schema(connection: &Connection) -> Result<()> {
|
||||
migrate_sqlite_ticket_schema_through(connection, LATEST_SQLITE_TICKET_SCHEMA_VERSION)
|
||||
}
|
||||
|
||||
/// Applies Ticket migrations only through `target_version`.
|
||||
///
|
||||
/// This exists for the Workspace Server's ordered migration bridge: older Server
|
||||
/// migrations must materialize the Ticket schema shape they were written against
|
||||
/// before the current Ticket migration is applied at the matching Server version.
|
||||
#[doc(hidden)]
|
||||
pub fn migrate_sqlite_ticket_schema_through(
|
||||
connection: &Connection,
|
||||
target_version: i64,
|
||||
) -> Result<()> {
|
||||
if !(1..=LATEST_SQLITE_TICKET_SCHEMA_VERSION).contains(&target_version) {
|
||||
return Err(TicketError::Sqlite(format!(
|
||||
"unsupported Ticket schema migration target {target_version}"
|
||||
)));
|
||||
}
|
||||
connection
|
||||
.busy_timeout(Duration::from_secs(5))
|
||||
.map_err(sqlite_err)?;
|
||||
@@ -302,25 +257,10 @@ pub fn migrate_sqlite_ticket_schema_through(
|
||||
verify_table(connection, MIGRATION_TABLE, MIGRATION_COLUMNS, &[], false)?;
|
||||
|
||||
let applied = load_applied_migrations(connection)?;
|
||||
validate_applied_migrations(&applied)?;
|
||||
|
||||
if let Some(version) = applied
|
||||
.keys()
|
||||
.copied()
|
||||
.find(|version| *version > target_version)
|
||||
{
|
||||
return Err(TicketError::Sqlite(format!(
|
||||
"Ticket schema version {version} is newer than requested migration target {target_version}"
|
||||
)));
|
||||
}
|
||||
|
||||
for migration in MIGRATIONS
|
||||
.iter()
|
||||
.filter(|migration| migration.version <= target_version)
|
||||
{
|
||||
if applied.contains_key(&migration.version) {
|
||||
continue;
|
||||
}
|
||||
if applied.is_empty() {
|
||||
let migration = MIGRATIONS
|
||||
.first()
|
||||
.ok_or_else(|| TicketError::Sqlite("Ticket migration catalog is empty".into()))?;
|
||||
(migration.apply)(connection)?;
|
||||
connection
|
||||
.execute(
|
||||
@@ -333,24 +273,11 @@ pub fn migrate_sqlite_ticket_schema_through(
|
||||
],
|
||||
)
|
||||
.map_err(sqlite_err)?;
|
||||
} else {
|
||||
validate_applied_migrations(&applied)?;
|
||||
}
|
||||
|
||||
if target_version == LATEST_SQLITE_TICKET_SCHEMA_VERSION {
|
||||
verify_sqlite_ticket_schema(connection)
|
||||
} else {
|
||||
let applied = load_applied_migrations(connection)?;
|
||||
let expected = MIGRATIONS
|
||||
.iter()
|
||||
.filter(|migration| migration.version <= target_version)
|
||||
.map(|migration| (migration.version, migration.name.to_string()))
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
if applied != expected {
|
||||
return Err(TicketError::Sqlite(format!(
|
||||
"Ticket schema migration history does not match target version {target_version}"
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
verify_sqlite_ticket_schema(connection)
|
||||
})();
|
||||
|
||||
match result {
|
||||
@@ -362,47 +289,6 @@ pub fn migrate_sqlite_ticket_schema_through(
|
||||
}
|
||||
}
|
||||
|
||||
/// Applies the resource-key Ticket migration inside a transaction owned by the
|
||||
/// Workspace Server. The caller must provide an active transaction; this function
|
||||
/// deliberately does not begin or commit one so the Ticket and Server migration
|
||||
/// markers can be persisted atomically.
|
||||
#[doc(hidden)]
|
||||
pub fn migrate_sqlite_ticket_resource_key_schema_in_transaction(
|
||||
connection: &Connection,
|
||||
) -> Result<()> {
|
||||
connection
|
||||
.execute_batch(
|
||||
"CREATE TABLE IF NOT EXISTS ticket_schema_migrations (
|
||||
version INTEGER PRIMARY KEY,
|
||||
name TEXT NOT NULL,
|
||||
applied_at TEXT NOT NULL
|
||||
);",
|
||||
)
|
||||
.map_err(sqlite_err)?;
|
||||
let applied = load_applied_migrations(connection)?;
|
||||
validate_applied_migrations(&applied)?;
|
||||
if applied.contains_key(&LATEST_SQLITE_TICKET_SCHEMA_VERSION) {
|
||||
return verify_sqlite_ticket_schema(connection);
|
||||
}
|
||||
let expected_previous = LATEST_SQLITE_TICKET_SCHEMA_VERSION - 1;
|
||||
if applied.len() != expected_previous as usize || !applied.contains_key(&expected_previous) {
|
||||
return Err(TicketError::Sqlite(format!(
|
||||
"Ticket schema must be at version {expected_previous} before the resource-key migration"
|
||||
)));
|
||||
}
|
||||
let migration = MIGRATIONS
|
||||
.last()
|
||||
.ok_or_else(|| TicketError::Sqlite("Ticket migration catalog is empty".to_string()))?;
|
||||
(migration.apply)(connection)?;
|
||||
connection
|
||||
.execute(
|
||||
"INSERT INTO ticket_schema_migrations (version, name, applied_at) VALUES (?1, ?2, datetime('now'))",
|
||||
params![migration.version, migration.name],
|
||||
)
|
||||
.map_err(sqlite_err)?;
|
||||
verify_sqlite_ticket_schema(connection)
|
||||
}
|
||||
|
||||
/// Verifies the current Ticket-owned SQLite schema without executing DDL.
|
||||
pub fn verify_sqlite_ticket_schema(connection: &Connection) -> Result<()> {
|
||||
let mut diagnostics = Vec::new();
|
||||
@@ -539,238 +425,9 @@ pub fn verify_sqlite_ticket_schema(connection: &Connection) -> Result<()> {
|
||||
}
|
||||
}
|
||||
|
||||
fn create_typed_ticket_tables(connection: &Connection) -> Result<()> {
|
||||
fn create_latest_ticket_schema(connection: &Connection) -> Result<()> {
|
||||
connection
|
||||
.execute_batch(
|
||||
r#"
|
||||
CREATE TABLE IF NOT EXISTS typed_tickets (
|
||||
workspace_id TEXT NOT NULL,
|
||||
ticket_id TEXT NOT NULL,
|
||||
slug TEXT NOT NULL,
|
||||
title TEXT NOT NULL,
|
||||
status TEXT NOT NULL,
|
||||
kind TEXT NOT NULL,
|
||||
priority TEXT NOT NULL,
|
||||
body TEXT NOT NULL,
|
||||
created_at TEXT,
|
||||
updated_at TEXT,
|
||||
assignee TEXT,
|
||||
readiness TEXT,
|
||||
workflow_state TEXT NOT NULL,
|
||||
workflow_state_explicit INTEGER NOT NULL,
|
||||
queued_by TEXT,
|
||||
queued_at TEXT,
|
||||
resolution TEXT,
|
||||
PRIMARY KEY (workspace_id, ticket_id)
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS typed_ticket_labels (
|
||||
workspace_id TEXT NOT NULL, ticket_id TEXT NOT NULL, ordinal INTEGER NOT NULL, label TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, ordinal),
|
||||
FOREIGN KEY (workspace_id, ticket_id) REFERENCES typed_tickets(workspace_id, ticket_id) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS typed_ticket_risk_flags (
|
||||
workspace_id TEXT NOT NULL, ticket_id TEXT NOT NULL, ordinal INTEGER NOT NULL, risk_flag TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, ordinal),
|
||||
FOREIGN KEY (workspace_id, ticket_id) REFERENCES typed_tickets(workspace_id, ticket_id) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS typed_ticket_raw_frontmatter (
|
||||
workspace_id TEXT NOT NULL, ticket_id TEXT NOT NULL, key TEXT NOT NULL, value TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, key),
|
||||
FOREIGN KEY (workspace_id, ticket_id) REFERENCES typed_tickets(workspace_id, ticket_id) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS typed_ticket_events (
|
||||
workspace_id TEXT NOT NULL,
|
||||
ticket_id TEXT NOT NULL,
|
||||
event_index INTEGER NOT NULL,
|
||||
kind TEXT NOT NULL,
|
||||
author TEXT,
|
||||
at TEXT,
|
||||
status TEXT,
|
||||
from_state TEXT,
|
||||
to_state TEXT,
|
||||
reason TEXT,
|
||||
state_field TEXT,
|
||||
heading TEXT,
|
||||
body TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, event_index),
|
||||
FOREIGN KEY (workspace_id, ticket_id) REFERENCES typed_tickets(workspace_id, ticket_id) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS typed_ticket_event_references (
|
||||
workspace_id TEXT NOT NULL, ticket_id TEXT NOT NULL, event_index INTEGER NOT NULL, ordinal INTEGER NOT NULL, kind TEXT NOT NULL, target TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, event_index, ordinal),
|
||||
FOREIGN KEY (workspace_id, ticket_id, event_index) REFERENCES typed_ticket_events(workspace_id, ticket_id, event_index) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS typed_ticket_event_attributes (
|
||||
workspace_id TEXT NOT NULL, ticket_id TEXT NOT NULL, event_index INTEGER NOT NULL, key TEXT NOT NULL, value TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, event_index, key),
|
||||
FOREIGN KEY (workspace_id, ticket_id, event_index) REFERENCES typed_ticket_events(workspace_id, ticket_id, event_index) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS typed_ticket_relations (
|
||||
workspace_id TEXT NOT NULL, ticket_id TEXT NOT NULL, kind TEXT NOT NULL, target TEXT NOT NULL, note TEXT, author TEXT NOT NULL, at TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, kind, target),
|
||||
FOREIGN KEY (workspace_id, ticket_id) REFERENCES typed_tickets(workspace_id, ticket_id) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS typed_ticket_orchestration_plans (
|
||||
workspace_id TEXT NOT NULL,
|
||||
ticket_id TEXT NOT NULL,
|
||||
record_id TEXT NOT NULL,
|
||||
kind TEXT NOT NULL,
|
||||
related_ticket TEXT,
|
||||
note TEXT,
|
||||
accepted_summary TEXT,
|
||||
accepted_branch TEXT,
|
||||
accepted_worktree TEXT,
|
||||
accepted_role_plan TEXT,
|
||||
author TEXT NOT NULL,
|
||||
at TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, record_id),
|
||||
FOREIGN KEY (workspace_id, ticket_id) REFERENCES typed_tickets(workspace_id, ticket_id) ON DELETE CASCADE
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS typed_ticket_artifacts (
|
||||
workspace_id TEXT NOT NULL, ticket_id TEXT NOT NULL, relative_path TEXT NOT NULL, content BLOB NOT NULL,
|
||||
PRIMARY KEY (workspace_id, ticket_id, relative_path),
|
||||
FOREIGN KEY (workspace_id, ticket_id) REFERENCES typed_tickets(workspace_id, ticket_id) ON DELETE CASCADE
|
||||
);
|
||||
"#,
|
||||
)
|
||||
.map_err(sqlite_err)
|
||||
}
|
||||
|
||||
fn add_ticket_repository_target(connection: &Connection) -> Result<()> {
|
||||
add_column_if_missing(connection, "typed_tickets", "repository_id", "TEXT")?;
|
||||
add_column_if_missing(connection, "typed_tickets", "ref_selector", "TEXT")
|
||||
}
|
||||
|
||||
fn retire_legacy_ticket_review_events(connection: &Connection) -> Result<()> {
|
||||
// Historical prose remains visible for audit, but it is explicitly converted to a
|
||||
// non-authoritative comment. Approval authority now lives only in Merge Requests.
|
||||
connection
|
||||
.execute_batch(
|
||||
r#"
|
||||
INSERT OR REPLACE INTO typed_ticket_event_attributes
|
||||
(workspace_id, ticket_id, event_index, key, value)
|
||||
SELECT workspace_id, ticket_id, event_index, 'legacy_event_kind', 'review'
|
||||
FROM typed_ticket_events WHERE kind = 'review';
|
||||
UPDATE typed_ticket_events
|
||||
SET kind = 'comment', status = NULL, heading = 'Legacy review (non-authoritative)'
|
||||
WHERE kind = 'review';
|
||||
DELETE FROM typed_ticket_event_attributes
|
||||
WHERE key IN ('result', 'review_result', 'status')
|
||||
AND EXISTS (
|
||||
SELECT 1 FROM typed_ticket_events event
|
||||
WHERE event.workspace_id = typed_ticket_event_attributes.workspace_id
|
||||
AND event.ticket_id = typed_ticket_event_attributes.ticket_id
|
||||
AND event.event_index = typed_ticket_event_attributes.event_index
|
||||
AND event.heading = 'Legacy review (non-authoritative)'
|
||||
);
|
||||
"#,
|
||||
)
|
||||
.map_err(sqlite_err)
|
||||
}
|
||||
|
||||
fn add_ticket_query_indexes(connection: &Connection) -> Result<()> {
|
||||
connection
|
||||
.execute_batch(
|
||||
r#"
|
||||
CREATE INDEX IF NOT EXISTS typed_tickets_workspace_state_updated
|
||||
ON typed_tickets(workspace_id, workflow_state, updated_at DESC, ticket_id);
|
||||
CREATE INDEX IF NOT EXISTS typed_tickets_workspace_updated
|
||||
ON typed_tickets(workspace_id, updated_at DESC, ticket_id);
|
||||
CREATE INDEX IF NOT EXISTS typed_tickets_workspace_created
|
||||
ON typed_tickets(workspace_id, created_at DESC, ticket_id);
|
||||
CREATE INDEX IF NOT EXISTS typed_tickets_workspace_title
|
||||
ON typed_tickets(workspace_id, title COLLATE NOCASE, ticket_id);
|
||||
CREATE INDEX IF NOT EXISTS typed_ticket_events_workspace_kind_ticket
|
||||
ON typed_ticket_events(workspace_id, kind, ticket_id, event_index);
|
||||
CREATE INDEX IF NOT EXISTS typed_ticket_relations_workspace_source_kind
|
||||
ON typed_ticket_relations(workspace_id, ticket_id, kind, target);
|
||||
CREATE INDEX IF NOT EXISTS typed_ticket_relations_workspace_target_kind
|
||||
ON typed_ticket_relations(workspace_id, target, kind, ticket_id);
|
||||
"#,
|
||||
)
|
||||
.map_err(sqlite_err)
|
||||
}
|
||||
|
||||
fn add_workspace_human_keys(connection: &Connection) -> Result<()> {
|
||||
connection
|
||||
.execute_batch(
|
||||
r#"
|
||||
CREATE TABLE IF NOT EXISTS workspace_resource_human_keys (
|
||||
workspace_id TEXT NOT NULL,
|
||||
resource_kind TEXT NOT NULL CHECK (resource_kind IN ('ticket', 'objective', 'worker')),
|
||||
resource_id TEXT NOT NULL,
|
||||
sequence INTEGER NOT NULL CHECK (sequence > 0),
|
||||
human_key TEXT NOT NULL,
|
||||
allocated_at TEXT NOT NULL,
|
||||
PRIMARY KEY (workspace_id, resource_kind, resource_id),
|
||||
UNIQUE (workspace_id, resource_kind, sequence),
|
||||
UNIQUE (workspace_id, human_key)
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS workspace_resource_human_key_counters (
|
||||
workspace_id TEXT NOT NULL,
|
||||
resource_kind TEXT NOT NULL CHECK (resource_kind IN ('ticket', 'objective', 'worker')),
|
||||
next_sequence INTEGER NOT NULL CHECK (next_sequence > 0),
|
||||
PRIMARY KEY (workspace_id, resource_kind)
|
||||
);
|
||||
|
||||
INSERT OR IGNORE INTO workspace_resource_human_keys (
|
||||
workspace_id, resource_kind, resource_id, sequence, human_key, allocated_at
|
||||
)
|
||||
SELECT workspace_id,
|
||||
'ticket',
|
||||
ticket_id,
|
||||
ROW_NUMBER() OVER (
|
||||
PARTITION BY workspace_id ORDER BY created_at ASC, ticket_id ASC
|
||||
),
|
||||
'T-' || ROW_NUMBER() OVER (
|
||||
PARTITION BY workspace_id ORDER BY created_at ASC, ticket_id ASC
|
||||
),
|
||||
COALESCE(created_at, updated_at)
|
||||
FROM typed_tickets;
|
||||
|
||||
INSERT INTO workspace_resource_human_key_counters (
|
||||
workspace_id, resource_kind, next_sequence
|
||||
)
|
||||
SELECT workspace_id, 'ticket', MAX(sequence) + 1
|
||||
FROM workspace_resource_human_keys
|
||||
WHERE resource_kind = 'ticket'
|
||||
GROUP BY workspace_id
|
||||
ON CONFLICT(workspace_id, resource_kind) DO UPDATE SET
|
||||
next_sequence = MAX(next_sequence, excluded.next_sequence);
|
||||
"#,
|
||||
)
|
||||
.map_err(sqlite_err)
|
||||
}
|
||||
|
||||
fn rename_workspace_resource_keys(connection: &Connection) -> Result<()> {
|
||||
connection
|
||||
.execute_batch(
|
||||
r#"
|
||||
ALTER TABLE workspace_resource_human_keys RENAME TO workspace_resource_keys;
|
||||
ALTER TABLE workspace_resource_keys RENAME COLUMN human_key TO resource_key;
|
||||
ALTER TABLE workspace_resource_human_key_counters RENAME TO workspace_resource_key_counters;
|
||||
DROP INDEX IF EXISTS idx_workspace_resource_human_keys_reverse;
|
||||
CREATE INDEX idx_workspace_resource_keys_reverse
|
||||
ON workspace_resource_keys(workspace_id, resource_kind, resource_key);
|
||||
"#,
|
||||
)
|
||||
.map_err(sqlite_err)
|
||||
}
|
||||
|
||||
fn add_column_if_missing(
|
||||
connection: &Connection,
|
||||
table: &str,
|
||||
column: &str,
|
||||
declaration: &str,
|
||||
) -> Result<()> {
|
||||
let columns = load_columns(connection, table)?;
|
||||
if columns.iter().any(|found| found.name == column) {
|
||||
return Ok(());
|
||||
}
|
||||
connection
|
||||
.execute_batch(&format!(
|
||||
"ALTER TABLE {table} ADD COLUMN {column} {declaration}"
|
||||
))
|
||||
.execute_batch(include_str!("latest_schema.sql"))
|
||||
.map_err(sqlite_err)
|
||||
}
|
||||
|
||||
@@ -796,33 +453,17 @@ fn load_applied_migrations(connection: &Connection) -> Result<BTreeMap<i64, Stri
|
||||
}
|
||||
|
||||
fn validate_applied_migrations(applied: &BTreeMap<i64, String>) -> Result<()> {
|
||||
for (&version, name) in applied {
|
||||
let Some(expected) = MIGRATIONS
|
||||
.iter()
|
||||
.find(|migration| migration.version == version)
|
||||
else {
|
||||
return Err(TicketError::Sqlite(format!(
|
||||
"unsupported Ticket schema migration version {version}; latest supported version is {LATEST_SQLITE_TICKET_SCHEMA_VERSION}"
|
||||
)));
|
||||
};
|
||||
if name != expected.name {
|
||||
return Err(TicketError::Sqlite(format!(
|
||||
"Ticket schema migration {version} is named {name:?}, expected {:?}",
|
||||
expected.name
|
||||
)));
|
||||
}
|
||||
let expected = BTreeMap::from([(
|
||||
LATEST_SQLITE_TICKET_SCHEMA_VERSION,
|
||||
MIGRATIONS[0].name.to_string(),
|
||||
)]);
|
||||
if applied == &expected {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(TicketError::Sqlite(format!(
|
||||
"Ticket schema migration history must contain only the canonical version {LATEST_SQLITE_TICKET_SCHEMA_VERSION} baseline marker"
|
||||
)))
|
||||
}
|
||||
for migration in MIGRATIONS {
|
||||
if applied.keys().any(|version| *version > migration.version)
|
||||
&& !applied.contains_key(&migration.version)
|
||||
{
|
||||
return Err(TicketError::Sqlite(format!(
|
||||
"Ticket schema migration history has a gap at version {}",
|
||||
migration.version
|
||||
)));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
@@ -1189,223 +830,16 @@ mod tests {
|
||||
verify_sqlite_ticket_schema(&connection).unwrap();
|
||||
|
||||
let versions = load_applied_migrations(&connection).unwrap();
|
||||
assert_eq!(versions.len(), 6);
|
||||
assert_eq!(
|
||||
versions.get(&LATEST_SQLITE_TICKET_SCHEMA_VERSION),
|
||||
Some(&"rename_workspace_resource_keys".to_string())
|
||||
versions,
|
||||
BTreeMap::from([(
|
||||
LATEST_SQLITE_TICKET_SCHEMA_VERSION,
|
||||
"ticket schema baseline".to_string(),
|
||||
)])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn adopts_existing_current_schema_without_losing_data() {
|
||||
let connection = Connection::open_in_memory().unwrap();
|
||||
create_typed_ticket_tables(&connection).unwrap();
|
||||
add_ticket_repository_target(&connection).unwrap();
|
||||
connection
|
||||
.execute(
|
||||
"INSERT INTO typed_tickets (
|
||||
workspace_id, ticket_id, slug, title, status, kind, priority, body,
|
||||
workflow_state, workflow_state_explicit, repository_id, ref_selector
|
||||
) VALUES ('workspace-1', 'ticket-1', 'ticket-1', 'kept', 'open',
|
||||
'task', 'medium', 'body', 'ready', 1, 'main', 'develop')",
|
||||
[],
|
||||
)
|
||||
.unwrap();
|
||||
connection
|
||||
.execute_batch(
|
||||
"INSERT INTO typed_ticket_events (
|
||||
workspace_id, ticket_id, event_index, kind, author, at, heading, body
|
||||
) VALUES (
|
||||
'workspace-1', 'ticket-1', 0, 'comment', 'hare',
|
||||
'2026-08-10T00:00:00Z', 'Evidence', 'event kept'
|
||||
);
|
||||
INSERT INTO typed_ticket_event_references (
|
||||
workspace_id, ticket_id, event_index, ordinal, kind, target
|
||||
) VALUES ('workspace-1', 'ticket-1', 0, 0, 'commit', 'abc123');
|
||||
INSERT INTO typed_ticket_relations (
|
||||
workspace_id, ticket_id, kind, target, note, author, at
|
||||
) VALUES (
|
||||
'workspace-1', 'ticket-1', 'related', 'ticket-2', 'relation kept',
|
||||
'hare', '2026-08-10T00:00:00Z'
|
||||
);
|
||||
INSERT INTO typed_ticket_orchestration_plans (
|
||||
workspace_id, ticket_id, record_id, kind, note, author, at
|
||||
) VALUES (
|
||||
'workspace-1', 'ticket-1', 'plan-1', 'waiting_capacity_note',
|
||||
'plan kept', 'hare', '2026-08-10T00:00:00Z'
|
||||
);
|
||||
INSERT INTO typed_ticket_artifacts (
|
||||
workspace_id, ticket_id, relative_path, content
|
||||
) VALUES ('workspace-1', 'ticket-1', 'evidence.txt', X'6b657074');",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
migrate_sqlite_ticket_schema(&connection).unwrap();
|
||||
|
||||
let row = connection
|
||||
.query_row(
|
||||
"SELECT title, repository_id, ref_selector FROM typed_tickets",
|
||||
[],
|
||||
|row| {
|
||||
Ok((
|
||||
row.get::<_, String>(0)?,
|
||||
row.get::<_, String>(1)?,
|
||||
row.get::<_, String>(2)?,
|
||||
))
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(row, ("kept".into(), "main".into(), "develop".into()));
|
||||
let preserved = connection
|
||||
.query_row(
|
||||
"SELECT
|
||||
(SELECT COUNT(*) FROM typed_ticket_events),
|
||||
(SELECT COUNT(*) FROM typed_ticket_event_references),
|
||||
(SELECT COUNT(*) FROM typed_ticket_relations),
|
||||
(SELECT COUNT(*) FROM typed_ticket_orchestration_plans),
|
||||
(SELECT COUNT(*) FROM typed_ticket_artifacts)",
|
||||
[],
|
||||
|row| {
|
||||
Ok((
|
||||
row.get::<_, i64>(0)?,
|
||||
row.get::<_, i64>(1)?,
|
||||
row.get::<_, i64>(2)?,
|
||||
row.get::<_, i64>(3)?,
|
||||
row.get::<_, i64>(4)?,
|
||||
))
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(preserved, (1, 1, 1, 1, 1));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn v5_backfills_ticket_keys_and_v6_preserves_them_under_resource_key_schema() {
|
||||
let connection = Connection::open_in_memory().unwrap();
|
||||
migrate_sqlite_ticket_schema_through(&connection, 4).unwrap();
|
||||
connection.execute_batch(
|
||||
"INSERT INTO typed_tickets (
|
||||
workspace_id, ticket_id, slug, title, status, kind, priority, body,
|
||||
workflow_state, workflow_state_explicit, created_at, updated_at
|
||||
) VALUES
|
||||
('workspace-1', 'later', 'later', 'Later', 'open', 'task', 'medium', '', 'ready', 1, '2026-01-02T00:00:00Z', '2026-01-02T00:00:00Z'),
|
||||
('workspace-1', 'earlier', 'earlier', 'Earlier', 'open', 'task', 'medium', '', 'ready', 1, '2026-01-01T00:00:00Z', '2026-01-01T00:00:00Z');"
|
||||
).unwrap();
|
||||
|
||||
migrate_sqlite_ticket_schema_through(&connection, 5).unwrap();
|
||||
let legacy_keys = connection
|
||||
.prepare(
|
||||
"SELECT resource_id, human_key FROM workspace_resource_human_keys
|
||||
WHERE workspace_id = 'workspace-1' AND resource_kind = 'ticket'
|
||||
ORDER BY sequence",
|
||||
)
|
||||
.unwrap()
|
||||
.query_map([], |row| {
|
||||
Ok((row.get::<_, String>(0)?, row.get::<_, String>(1)?))
|
||||
})
|
||||
.unwrap()
|
||||
.collect::<std::result::Result<Vec<_>, _>>()
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
legacy_keys,
|
||||
vec![
|
||||
("earlier".into(), "T-1".into()),
|
||||
("later".into(), "T-2".into())
|
||||
]
|
||||
);
|
||||
let next: i64 = connection
|
||||
.query_row(
|
||||
"SELECT next_sequence FROM workspace_resource_human_key_counters
|
||||
WHERE workspace_id = 'workspace-1' AND resource_kind = 'ticket'",
|
||||
[],
|
||||
|row| row.get(0),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(next, 3);
|
||||
|
||||
migrate_sqlite_ticket_schema(&connection).unwrap();
|
||||
let resource_keys = connection
|
||||
.prepare(
|
||||
"SELECT resource_id, resource_key FROM workspace_resource_keys
|
||||
WHERE workspace_id = 'workspace-1' AND resource_kind = 'ticket'
|
||||
ORDER BY sequence",
|
||||
)
|
||||
.unwrap()
|
||||
.query_map([], |row| {
|
||||
Ok((row.get::<_, String>(0)?, row.get::<_, String>(1)?))
|
||||
})
|
||||
.unwrap()
|
||||
.collect::<std::result::Result<Vec<_>, _>>()
|
||||
.unwrap();
|
||||
assert_eq!(resource_keys, legacy_keys);
|
||||
assert_eq!(
|
||||
connection
|
||||
.query_row(
|
||||
"SELECT next_sequence FROM workspace_resource_key_counters
|
||||
WHERE workspace_id = 'workspace-1' AND resource_kind = 'ticket'",
|
||||
[],
|
||||
|row| row.get::<_, i64>(0),
|
||||
)
|
||||
.unwrap(),
|
||||
3
|
||||
);
|
||||
for legacy_table in [
|
||||
"workspace_resource_human_keys",
|
||||
"workspace_resource_human_key_counters",
|
||||
] {
|
||||
assert!(
|
||||
connection
|
||||
.query_row(
|
||||
"SELECT 1 FROM sqlite_schema WHERE type = 'table' AND name = ?1",
|
||||
[legacy_table],
|
||||
|_| Ok(()),
|
||||
)
|
||||
.optional()
|
||||
.unwrap()
|
||||
.is_none(),
|
||||
"{legacy_table} still exists"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn upgrades_legacy_schema_without_repository_target_columns() {
|
||||
let connection = Connection::open_in_memory().unwrap();
|
||||
create_typed_ticket_tables(&connection).unwrap();
|
||||
connection
|
||||
.execute(
|
||||
"INSERT INTO typed_tickets (
|
||||
workspace_id, ticket_id, slug, title, status, kind, priority, body,
|
||||
workflow_state, workflow_state_explicit
|
||||
) VALUES ('workspace-1', 'ticket-1', 'ticket-1', 'legacy', 'open',
|
||||
'task', 'medium', 'body', 'ready', 1)",
|
||||
[],
|
||||
)
|
||||
.unwrap();
|
||||
connection
|
||||
.execute_batch(
|
||||
"CREATE TABLE ticket_schema_migrations (
|
||||
version INTEGER PRIMARY KEY,
|
||||
name TEXT NOT NULL,
|
||||
applied_at TEXT NOT NULL
|
||||
);
|
||||
INSERT INTO ticket_schema_migrations (version, name, applied_at)
|
||||
VALUES (1, 'create_typed_ticket_tables', '2026-08-10T00:00:00Z');",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
migrate_sqlite_ticket_schema(&connection).unwrap();
|
||||
verify_sqlite_ticket_schema(&connection).unwrap();
|
||||
|
||||
let columns = load_columns(&connection, "typed_tickets").unwrap();
|
||||
assert!(columns.iter().any(|column| column.name == "repository_id"));
|
||||
assert!(columns.iter().any(|column| column.name == "ref_selector"));
|
||||
let title = connection
|
||||
.query_row("SELECT title FROM typed_tickets", [], |row| {
|
||||
row.get::<_, String>(0)
|
||||
})
|
||||
.unwrap();
|
||||
assert_eq!(title, "legacy");
|
||||
assert_eq!(load_applied_migrations(&connection).unwrap(), versions);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -1421,12 +855,32 @@ mod tests {
|
||||
.unwrap();
|
||||
|
||||
let error = migrate_sqlite_ticket_schema(&connection).unwrap_err();
|
||||
assert!(
|
||||
error
|
||||
.to_string()
|
||||
.contains("unsupported Ticket schema migration version 99")
|
||||
);
|
||||
assert_eq!(load_applied_migrations(&connection).unwrap().len(), 7);
|
||||
assert!(error.to_string().contains(
|
||||
"migration history must contain only the canonical version 6 baseline marker"
|
||||
));
|
||||
assert_eq!(load_applied_migrations(&connection).unwrap().len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_legacy_migration_marker() {
|
||||
let connection = Connection::open_in_memory().unwrap();
|
||||
connection
|
||||
.execute_batch(
|
||||
"CREATE TABLE ticket_schema_migrations (
|
||||
version INTEGER PRIMARY KEY,
|
||||
name TEXT NOT NULL,
|
||||
applied_at TEXT NOT NULL
|
||||
);
|
||||
INSERT INTO ticket_schema_migrations (version, name, applied_at)
|
||||
VALUES (6, 'rename_workspace_resource_keys', '2026-08-10T00:00:00Z');",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let error = migrate_sqlite_ticket_schema(&connection).unwrap_err();
|
||||
assert!(error.to_string().contains(
|
||||
"migration history must contain only the canonical version 6 baseline marker"
|
||||
));
|
||||
assert!(!table_exists(&connection, "typed_tickets").unwrap());
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -1509,77 +963,6 @@ mod tests {
|
||||
verify_sqlite_ticket_schema(&connection).unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn migration_rejects_constraint_drift_and_rolls_back_version_adoption() {
|
||||
let connection = Connection::open_in_memory().unwrap();
|
||||
connection
|
||||
.execute_batch(
|
||||
"CREATE TABLE typed_tickets (
|
||||
workspace_id TEXT NOT NULL,
|
||||
ticket_id TEXT NOT NULL,
|
||||
slug TEXT NOT NULL,
|
||||
title TEXT NOT NULL,
|
||||
status TEXT NOT NULL,
|
||||
kind TEXT NOT NULL,
|
||||
priority TEXT NOT NULL,
|
||||
body TEXT NOT NULL,
|
||||
created_at TEXT,
|
||||
updated_at TEXT,
|
||||
assignee TEXT,
|
||||
readiness TEXT,
|
||||
workflow_state TEXT NOT NULL,
|
||||
workflow_state_explicit INTEGER NOT NULL,
|
||||
queued_by TEXT,
|
||||
queued_at TEXT,
|
||||
resolution TEXT,
|
||||
PRIMARY KEY (ticket_id, workspace_id)
|
||||
);",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let error = migrate_sqlite_ticket_schema(&connection).unwrap_err();
|
||||
assert!(error.to_string().contains("primary-key position"));
|
||||
let migration_table_exists = connection
|
||||
.query_row(
|
||||
"SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = 'ticket_schema_migrations'",
|
||||
[],
|
||||
|_| Ok(()),
|
||||
)
|
||||
.optional()
|
||||
.unwrap()
|
||||
.is_some();
|
||||
assert!(!migration_table_exists);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn legacy_review_upgrade_preserves_prose_as_non_authoritative_comment() {
|
||||
let connection = Connection::open_in_memory().unwrap();
|
||||
migrate_sqlite_ticket_schema(&connection).unwrap();
|
||||
connection.execute("INSERT INTO typed_tickets (workspace_id,ticket_id,slug,title,status,kind,priority,body,workflow_state,workflow_state_explicit) VALUES ('workspace-1','ticket-1','ticket-1','title','open','task','medium','body','inprogress',1)",[]).unwrap();
|
||||
connection.execute("INSERT INTO typed_ticket_events (workspace_id,ticket_id,event_index,kind,author,at,status,heading,body) VALUES ('workspace-1','ticket-1',0,'review','reviewer','2026-08-11T00:00:00Z','approve','Review','legacy evidence')",[]).unwrap();
|
||||
connection.execute("INSERT INTO typed_ticket_event_attributes (workspace_id,ticket_id,event_index,key,value) VALUES ('workspace-1','ticket-1',0,'result','approve')",[]).unwrap();
|
||||
connection
|
||||
.execute_batch(
|
||||
"DROP TABLE workspace_resource_key_counters;
|
||||
DROP TABLE workspace_resource_keys;
|
||||
DELETE FROM ticket_schema_migrations WHERE version >= 3;",
|
||||
)
|
||||
.unwrap();
|
||||
migrate_sqlite_ticket_schema(&connection).unwrap();
|
||||
let (kind,status,heading,body):(String,Option<String>,Option<String>,Option<String>)=connection.query_row("SELECT kind,status,heading,body FROM typed_ticket_events WHERE workspace_id='workspace-1' AND ticket_id='ticket-1' AND event_index=0",[],|row|Ok((row.get(0)?,row.get(1)?,row.get(2)?,row.get(3)?))).unwrap();
|
||||
assert_eq!(kind, "comment");
|
||||
assert_eq!(status, None);
|
||||
assert_eq!(
|
||||
heading.as_deref(),
|
||||
Some("Legacy review (non-authoritative)")
|
||||
);
|
||||
assert_eq!(body.as_deref(), Some("legacy evidence"));
|
||||
let attributes:i64=connection.query_row("SELECT COUNT(*) FROM typed_ticket_event_attributes WHERE workspace_id='workspace-1' AND ticket_id='ticket-1'",[],|row|row.get(0)).unwrap();
|
||||
assert_eq!(attributes, 1);
|
||||
let legacy:String=connection.query_row("SELECT value FROM typed_ticket_event_attributes WHERE workspace_id='workspace-1' AND ticket_id='ticket-1' AND key='legacy_event_kind'",[],|row|row.get(0)).unwrap();
|
||||
assert_eq!(legacy, "review");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn concurrent_migrators_converge_on_one_version_history() {
|
||||
let directory = tempdir().unwrap();
|
||||
@@ -1602,6 +985,6 @@ mod tests {
|
||||
|
||||
let connection = Connection::open(database).unwrap();
|
||||
verify_sqlite_ticket_schema(&connection).unwrap();
|
||||
assert_eq!(load_applied_migrations(&connection).unwrap().len(), 6);
|
||||
assert_eq!(load_applied_migrations(&connection).unwrap().len(), 1);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -402,8 +402,8 @@ struct TicketCreateParams {
|
||||
queued_at: Option<String>,
|
||||
/// Optional target Workspace repository id.
|
||||
#[serde(default)]
|
||||
repository_id: Option<String>,
|
||||
/// Optional target Git ref selector. Requires `repository_id`.
|
||||
repository_key: Option<String>,
|
||||
/// Optional target Git ref selector. Requires `repository_key`.
|
||||
#[serde(default)]
|
||||
ref_selector: Option<String>,
|
||||
}
|
||||
@@ -944,7 +944,7 @@ impl Tool for TicketCreateTool {
|
||||
input.workflow_state = params.state.map(TicketWorkflowStateParam::into_state);
|
||||
input.queued_by = None;
|
||||
input.queued_at = params.queued_at;
|
||||
input.repository_id = params.repository_id;
|
||||
input.repository_id = params.repository_key;
|
||||
input.ref_selector = params.ref_selector;
|
||||
|
||||
let created = self
|
||||
@@ -1173,7 +1173,7 @@ impl Tool for TicketMarkReadyTool {
|
||||
json!({
|
||||
"ticket": ticket.meta.id,
|
||||
"state": ticket.meta.workflow_state.as_str(),
|
||||
"repository_id": ticket.meta.repository_id,
|
||||
"repository_key": ticket.meta.repository_id,
|
||||
"ref_selector": ticket.meta.ref_selector,
|
||||
"ok": true
|
||||
}),
|
||||
@@ -1206,7 +1206,7 @@ impl Tool for TicketIntakeReadyTool {
|
||||
json!({
|
||||
"ticket": ticket.meta.id,
|
||||
"state": ticket.meta.workflow_state.as_str(),
|
||||
"repository_id": ticket.meta.repository_id,
|
||||
"repository_key": ticket.meta.repository_id,
|
||||
"ref_selector": ticket.meta.ref_selector,
|
||||
"ok": true
|
||||
}),
|
||||
@@ -1546,7 +1546,7 @@ fn model_ticket_reference(
|
||||
match ticket.meta.resource_key {
|
||||
Some(resource_key) if is_canonical_ticket_resource_key(&resource_key) => Ok(resource_key),
|
||||
Some(_) => Err(ToolError::ExecutionFailed(format!(
|
||||
"{tool_name} failed: required Ticket human key is unavailable"
|
||||
"{tool_name} failed: required Ticket key is unavailable"
|
||||
))),
|
||||
None => Ok(ticket.meta.id),
|
||||
}
|
||||
@@ -1940,11 +1940,11 @@ mod tests {
|
||||
fn resolve_target(
|
||||
&self,
|
||||
_workspace_id: &str,
|
||||
repository_id: Option<&str>,
|
||||
repository_key: Option<&str>,
|
||||
ref_selector: Option<&str>,
|
||||
) -> crate::Result<crate::ResolvedTicketTarget> {
|
||||
Ok(crate::ResolvedTicketTarget {
|
||||
repository_id: repository_id.unwrap_or("main").to_owned(),
|
||||
repository_id: repository_key.unwrap_or("main").to_owned(),
|
||||
ref_selector: ref_selector.unwrap_or("develop").to_owned(),
|
||||
})
|
||||
}
|
||||
|
||||
+134
-6
@@ -21,6 +21,7 @@ struct BashParams {
|
||||
|
||||
pub(crate) struct BashTool {
|
||||
session: WorkdirSessionHandle,
|
||||
output_dir: PathBuf,
|
||||
state: Arc<Mutex<BashExecutionState>>,
|
||||
}
|
||||
|
||||
@@ -117,6 +118,7 @@ impl Tool for BashTool {
|
||||
command: params.command,
|
||||
timeout_secs,
|
||||
output_limit: INLINE_BYTE_BUDGET,
|
||||
spill_dir: Some(self.output_dir.clone()),
|
||||
tool_call_id: Some(call_id.clone()),
|
||||
})
|
||||
.await
|
||||
@@ -183,10 +185,15 @@ impl Tool for BashTool {
|
||||
let content = if output.content.is_empty() {
|
||||
None
|
||||
} else if output.truncated {
|
||||
Some(format!(
|
||||
"[showing bounded WorkdirSession command output; additional output was truncated]\n{}",
|
||||
output.content
|
||||
))
|
||||
let notice = match output.output_path {
|
||||
Some(path) => format!(
|
||||
"[showing bounded WorkdirSession command output; full output saved to {}]",
|
||||
path.display()
|
||||
),
|
||||
None => "[showing bounded WorkdirSession command output; additional output was truncated]"
|
||||
.to_owned(),
|
||||
};
|
||||
Some(format!("{notice}\n{}", output.content))
|
||||
} else {
|
||||
Some(output.content)
|
||||
};
|
||||
@@ -259,16 +266,137 @@ fn truncate_for_summary(command: &str) -> String {
|
||||
summary
|
||||
}
|
||||
|
||||
pub fn bash_tool(session: WorkdirSessionHandle, _output_dir: PathBuf) -> ToolDefinition {
|
||||
pub fn bash_tool(session: WorkdirSessionHandle, output_dir: PathBuf) -> ToolDefinition {
|
||||
Arc::new(move || {
|
||||
let schema = schemars::schema_for!(BashParams);
|
||||
let meta = ToolMeta::new("Bash")
|
||||
.description("Execute a shell command in the bound Workdir. Process start, bounded output, timeout and cancellation are owned by the WorkdirSession provider. This is not a sandbox.")
|
||||
.description("Execute a shell command in the bound Workdir. Process start, bounded inline output, full-output spill, timeout and cancellation are owned by the WorkdirSession provider. This is not a sandbox.")
|
||||
.input_schema(serde_json::to_value(schema).expect("Bash schema serialization"));
|
||||
let tool: Arc<dyn Tool> = Arc::new(BashTool {
|
||||
session: session.clone(),
|
||||
output_dir: output_dir.clone(),
|
||||
state: Arc::new(Mutex::new(BashExecutionState::default())),
|
||||
});
|
||||
(meta, tool)
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use manifest::{Permission, Scope, ScopeConfig, ScopeRule};
|
||||
use tempfile::TempDir;
|
||||
use workdir::{LocalWorkdirSession, WorkdirSessionHandle};
|
||||
|
||||
use super::bash_tool;
|
||||
use crate::{grep::grep_tool, read::read_tool, tracker::Tracker};
|
||||
|
||||
fn session_with_output_scope(root: &TempDir, output: &TempDir) -> WorkdirSessionHandle {
|
||||
let scope = Scope::from_config(&ScopeConfig {
|
||||
allow: vec![
|
||||
ScopeRule {
|
||||
target: root.path().to_path_buf(),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
},
|
||||
ScopeRule {
|
||||
target: output.path().to_path_buf(),
|
||||
permission: Permission::Read,
|
||||
recursive: true,
|
||||
},
|
||||
],
|
||||
deny: Vec::new(),
|
||||
})
|
||||
.unwrap();
|
||||
Arc::new(LocalWorkdirSession::new(scope, root.path().to_path_buf()))
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn long_output_is_spilled_and_available_to_read_and_grep() {
|
||||
let root = TempDir::new().unwrap();
|
||||
let output = TempDir::new().unwrap();
|
||||
let session = session_with_output_scope(&root, &output);
|
||||
let (_, bash) = bash_tool(session.clone(), output.path().to_path_buf())();
|
||||
let command = "i=0; while [ $i -lt 2000 ]; do printf 'line-%04d\\n' \"$i\"; i=$((i+1)); done; printf 'FINAL-NEEDLE\\n'";
|
||||
let result = bash
|
||||
.execute(
|
||||
&serde_json::json!({ "command": command }).to_string(),
|
||||
Default::default(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let rendered = result.content.expect("bounded Bash output");
|
||||
let artifact = std::fs::read_dir(output.path())
|
||||
.unwrap()
|
||||
.next()
|
||||
.expect("artifact entry")
|
||||
.unwrap()
|
||||
.path();
|
||||
|
||||
assert!(rendered.contains("full output saved to"));
|
||||
assert!(rendered.contains(&artifact.display().to_string()));
|
||||
let retained = std::fs::read_to_string(&artifact).unwrap();
|
||||
assert!(retained.starts_with("line-0000\n"));
|
||||
assert!(retained.ends_with("FINAL-NEEDLE\n"));
|
||||
assert_eq!(retained.lines().count(), 2001);
|
||||
|
||||
let (_, read) = read_tool(session.clone(), Tracker::new())();
|
||||
let read_result = read
|
||||
.execute(
|
||||
&serde_json::json!({
|
||||
"file_path": artifact,
|
||||
"offset": 2000,
|
||||
"limit": 1,
|
||||
})
|
||||
.to_string(),
|
||||
Default::default(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(
|
||||
read_result
|
||||
.content
|
||||
.expect("Read content")
|
||||
.contains("FINAL-NEEDLE")
|
||||
);
|
||||
|
||||
let (_, grep) = grep_tool(session)();
|
||||
let grep_result = grep
|
||||
.execute(
|
||||
&serde_json::json!({
|
||||
"pattern": "FINAL-NEEDLE",
|
||||
"path": artifact,
|
||||
"output_mode": "content",
|
||||
})
|
||||
.to_string(),
|
||||
Default::default(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let grep_content = grep_result.content.expect("Grep content");
|
||||
assert!(
|
||||
grep_content.contains("FINAL-NEEDLE"),
|
||||
"unexpected Grep content: {grep_content:?}"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn short_output_does_not_leave_a_spill_artifact() {
|
||||
let root = TempDir::new().unwrap();
|
||||
let output = TempDir::new().unwrap();
|
||||
let session = session_with_output_scope(&root, &output);
|
||||
let (_, bash) = bash_tool(session, output.path().to_path_buf())();
|
||||
|
||||
let result = bash
|
||||
.execute(
|
||||
&serde_json::json!({ "command": "printf short" }).to_string(),
|
||||
Default::default(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(result.content.as_deref(), Some("short"));
|
||||
assert_eq!(std::fs::read_dir(output.path()).unwrap().count(), 0);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -42,6 +42,7 @@ impl From<ToolsError> for ToolError {
|
||||
workdir::WorkdirError::NotFound(_)
|
||||
| workdir::WorkdirError::Io { .. }
|
||||
| workdir::WorkdirError::Unavailable(_)
|
||||
| workdir::WorkdirError::OperationFailed
|
||||
| workdir::WorkdirError::Transport(_),
|
||||
) => ToolError::ExecutionFailed(err.to_string()),
|
||||
ToolsError::FileSystem(_)
|
||||
|
||||
@@ -22,7 +22,7 @@ enum OutputMode {
|
||||
#[derive(Debug, Deserialize, JsonSchema)]
|
||||
struct GrepParams {
|
||||
pattern: String,
|
||||
/// Logical Workdir-relative path to search. Defaults to the Workdir root.
|
||||
/// Workdir-relative path, or an absolute path covered by readable scope. Defaults to the Workdir root.
|
||||
#[serde(default)]
|
||||
path: Option<String>,
|
||||
#[serde(default)]
|
||||
@@ -61,7 +61,7 @@ impl Tool for GrepTool {
|
||||
let params: GrepParams = serde_json::from_str(input_json)
|
||||
.map_err(|error| ToolError::InvalidArgument(format!("invalid Grep input: {error}")))?;
|
||||
let path = match params.path {
|
||||
Some(path) => WorkdirPath::new(&path).map_err(ToolsError::from)?,
|
||||
Some(path) => WorkdirPath::new_scoped(&path).map_err(ToolsError::from)?,
|
||||
None => WorkdirPath::root(),
|
||||
};
|
||||
let mode = match params.output_mode.unwrap_or_default() {
|
||||
@@ -129,7 +129,7 @@ pub fn grep_tool(session: WorkdirSessionHandle) -> ToolDefinition {
|
||||
Arc::new(move || {
|
||||
let schema = schemars::schema_for!(GrepParams);
|
||||
let meta = ToolMeta::new("Grep")
|
||||
.description("Search Workdir file contents with a regex. Glob/Grep traversal executes inside the WorkdirSession provider. Results are bounded and Workdir-relative.")
|
||||
.description("Search a Workdir file or directory with a regex. Content results group lines by file; `>` marks matching lines and unmarked lines are context. Directory traversal executes inside the WorkdirSession provider. Results are bounded and Workdir-relative.")
|
||||
.input_schema(serde_json::to_value(schema).expect("Grep schema serialization"));
|
||||
let tool: Arc<dyn Tool> = Arc::new(GrepTool {
|
||||
session: session.clone(),
|
||||
|
||||
@@ -13,14 +13,14 @@ use workdir::{ReadRequest, WorkdirPath, WorkdirSessionHandle};
|
||||
const DESCRIPTION: &str = "Read a text file from the local filesystem. \
|
||||
Supports offset/limit for large files. Returns line-numbered output (1-based). \
|
||||
Directories cannot be read. The file must be read before Write or Edit can \
|
||||
modify it. Paths are relative to the bound Workdir.";
|
||||
modify it. Paths are Workdir-relative unless an absolute path is explicitly readable.";
|
||||
|
||||
const DEFAULT_LIMIT: usize = 2000;
|
||||
const PROVIDER_BYTE_LIMIT: usize = 256 * 1024;
|
||||
|
||||
#[derive(Debug, Deserialize, schemars::JsonSchema)]
|
||||
pub(crate) struct ReadParams {
|
||||
/// Logical path relative to the bound Workdir root.
|
||||
/// Workdir-relative path, or an absolute path covered by readable scope.
|
||||
pub file_path: String,
|
||||
/// 0-based line offset from the start. Defaults to 0.
|
||||
#[serde(default)]
|
||||
@@ -47,7 +47,7 @@ impl Tool for ReadTool {
|
||||
let offset = params.offset.unwrap_or(0);
|
||||
let limit = params.limit.unwrap_or(DEFAULT_LIMIT).max(1);
|
||||
|
||||
let path = WorkdirPath::new(¶ms.file_path).map_err(ToolsError::from)?;
|
||||
let path = WorkdirPath::new_scoped(¶ms.file_path).map_err(ToolsError::from)?;
|
||||
tracing::debug!(path = %path, offset, limit, "Read");
|
||||
|
||||
let result = self
|
||||
|
||||
@@ -224,20 +224,23 @@ async fn very_long_single_line() {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn absolute_path_is_rejected() {
|
||||
let (dir, _spill, reg) = setup();
|
||||
async fn absolute_path_requires_matching_read_scope() {
|
||||
let (_dir, _spill, reg) = setup();
|
||||
let outside = tempfile::tempdir().unwrap();
|
||||
let outside_file = outside.path().join("outside.txt");
|
||||
std::fs::write(&outside_file, "secret").unwrap();
|
||||
let read = reg.get("Read");
|
||||
let err = read
|
||||
.execute(
|
||||
&json!({ "file_path": dir.path().join("outside.txt") }).to_string(),
|
||||
&json!({ "file_path": outside_file }).to_string(),
|
||||
Default::default(),
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
let msg = format!("{err}");
|
||||
assert!(
|
||||
msg.contains("invalid logical filesystem path"),
|
||||
"absolute path was not rejected as invalid: {msg}"
|
||||
msg.contains("outside allowed scope"),
|
||||
"absolute path escaped readable scope: {msg}"
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -394,14 +394,21 @@ async fn bash_inherits_workdir_cwd() {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn bash_provider_output_does_not_expose_internal_paths() {
|
||||
async fn bash_provider_output_exposes_readable_retained_path() {
|
||||
let (_dir, spill, reg) = setup();
|
||||
let bash = reg.get("Bash");
|
||||
let out = call(&bash, json!({ "command": "printf 'x%.0s' {1..20480}" })).await;
|
||||
let body = out.content.unwrap();
|
||||
assert!(body.contains("bounded WorkdirSession command output"));
|
||||
assert!(!body.contains(spill.path().to_str().unwrap()));
|
||||
assert_eq!(std::fs::read_dir(spill.path()).unwrap().count(), 0);
|
||||
assert!(body.contains("full output saved to"));
|
||||
assert!(body.contains(spill.path().to_str().unwrap()));
|
||||
let artifact = std::fs::read_dir(spill.path())
|
||||
.unwrap()
|
||||
.next()
|
||||
.expect("retained output")
|
||||
.unwrap()
|
||||
.path();
|
||||
assert_eq!(std::fs::metadata(artifact).unwrap().len(), 20_480);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
@@ -10,11 +10,13 @@ e2e-test = []
|
||||
|
||||
[dependencies]
|
||||
client = { workspace = true }
|
||||
standalone = { workspace = true }
|
||||
thiserror.workspace = true
|
||||
protocol = { workspace = true }
|
||||
ratatui = { version = "0.30.0", features = ["scrolling-regions"] }
|
||||
base64 = "0.22.1"
|
||||
crossterm = "0.28"
|
||||
tokio = { workspace = true, features = ["rt-multi-thread", "macros", "net", "io-util", "sync", "time", "process"] }
|
||||
tokio = { workspace = true, features = ["rt-multi-thread", "macros", "sync", "time"] }
|
||||
serde_json = { workspace = true }
|
||||
unicode-width = "0.2.2"
|
||||
uuid = { workspace = true }
|
||||
@@ -22,12 +24,11 @@ toml = { workspace = true }
|
||||
manifest = { workspace = true }
|
||||
secrets = { workspace = true }
|
||||
session-store = { workspace = true }
|
||||
fs4 = { workspace = true }
|
||||
ticket = { workspace = true }
|
||||
serde = { workspace = true, features = ["derive"] }
|
||||
worker = { path = "../worker" }
|
||||
pulldown-cmark = { version = "0.13.3", default-features = false }
|
||||
agen.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
async-trait.workspace = true
|
||||
tempfile = { workspace = true }
|
||||
|
||||
+224
-197
@@ -249,6 +249,9 @@ pub struct App {
|
||||
pub running: bool,
|
||||
/// True while the Worker is in `WorkerStatus::Paused`.
|
||||
pub paused: bool,
|
||||
/// Local observation time for the current run. Used only for live UI
|
||||
/// elapsed time and spinner animation; it is not persisted in history.
|
||||
pub run_started_at: Option<Instant>,
|
||||
pub run_requests: usize,
|
||||
/// Sum of `input_tokens - cache_read_input_tokens` across the
|
||||
/// current turn's LLM requests — i.e. the net tokens this turn
|
||||
@@ -281,6 +284,9 @@ pub struct App {
|
||||
/// records the instant; a second press within the timeout exits the
|
||||
/// TUI (the Worker itself stays alive).
|
||||
pub quit_confirm: Option<std::time::Instant>,
|
||||
/// Independent 2-tap guard for `Ctrl-X` when the Worker is idle or
|
||||
/// stopped. A second press within the timeout shuts down the Worker.
|
||||
pub shutdown_confirm: Option<std::time::Instant>,
|
||||
/// Full display history in render order.
|
||||
pub blocks: Vec<Block>,
|
||||
/// Turn/protocol errors retained when a real `SegmentStart` replaces the
|
||||
@@ -352,6 +358,7 @@ impl App {
|
||||
worker_status: WorkerStatus::Idle,
|
||||
running: false,
|
||||
paused: false,
|
||||
run_started_at: None,
|
||||
run_requests: 0,
|
||||
run_upload_tokens: 0,
|
||||
run_output_tokens: 0,
|
||||
@@ -369,6 +376,7 @@ impl App {
|
||||
command_completion_selected: None,
|
||||
quit: false,
|
||||
quit_confirm: None,
|
||||
shutdown_confirm: None,
|
||||
blocks: Vec::new(),
|
||||
run_error_messages: Vec::new(),
|
||||
internal_workers: Vec::new(),
|
||||
@@ -553,11 +561,18 @@ impl App {
|
||||
}
|
||||
|
||||
pub fn set_worker_status(&mut self, status: WorkerStatus) {
|
||||
let was_running = self.running;
|
||||
self.worker_status = status;
|
||||
self.running = status == WorkerStatus::Running;
|
||||
self.paused = status == WorkerStatus::Paused;
|
||||
if self.running {
|
||||
if !was_running {
|
||||
self.run_started_at = Some(Instant::now());
|
||||
}
|
||||
self.quit_confirm = None;
|
||||
self.shutdown_confirm = None;
|
||||
} else {
|
||||
self.run_started_at = None;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -763,9 +778,23 @@ impl App {
|
||||
Some(self.method_for_run(segments))
|
||||
}
|
||||
|
||||
pub fn restore_unsent_run(&mut self, method: &Method) {
|
||||
let Method::Run { input } = method else {
|
||||
return;
|
||||
};
|
||||
self.pending_submit_rollback = None;
|
||||
if self.input.is_empty() {
|
||||
self.input.replace_with_segments(input);
|
||||
self.completion = None;
|
||||
} else {
|
||||
self.queued_inputs
|
||||
.push_front(QueuedInput::new(input.clone()));
|
||||
}
|
||||
}
|
||||
|
||||
fn method_for_run(&mut self, segments: Vec<Segment>) -> Method {
|
||||
// TurnHeader / UserMessage blocks are pushed only after the Worker
|
||||
// emits `Event::UserMessage` from a committed `LogEntry::UserInput`.
|
||||
// emits `Event::UserMessage` from a committed `LogEntry::AnnotatedUserInput`.
|
||||
// Locally we only clear the input buffer and forward the method,
|
||||
// while remembering enough local state to undo the visible submit if
|
||||
// the accepted run produced no assistant output and was rolled back.
|
||||
@@ -913,6 +942,10 @@ impl App {
|
||||
Some(self.method_for_run(queued.segments))
|
||||
}
|
||||
|
||||
pub fn clear_actionbar_notice(&mut self) {
|
||||
self.actionbar_notice = None;
|
||||
}
|
||||
|
||||
pub fn push_error(&mut self, message: impl Into<String>) {
|
||||
self.blocks.push(Block::Alert {
|
||||
level: AlertLevel::Error,
|
||||
@@ -1098,10 +1131,9 @@ impl App {
|
||||
self.blocks.push(Block::UserMessage { segments });
|
||||
self.assistant_streaming = false;
|
||||
}
|
||||
Event::SegmentRotated { entry } => {
|
||||
Event::SegmentRotated { session } => {
|
||||
let retained_run_errors = self.run_error_messages.clone();
|
||||
self.reset_for_rotation();
|
||||
self.apply_log_entry_raw(&entry);
|
||||
self.restore_session(&session, self.greeting.clone());
|
||||
for message in retained_run_errors {
|
||||
self.blocks.push(Block::Alert {
|
||||
level: AlertLevel::Error,
|
||||
@@ -1122,11 +1154,13 @@ impl App {
|
||||
self.latest_llm_wait_event = None;
|
||||
self.assistant_streaming = false;
|
||||
}
|
||||
// UI consumers of Invoke / LlmCall semantics are out of scope
|
||||
// for `tickets/invoke-turn-llmcall-semantics.md`; events flow
|
||||
// through to subscribers but the TUI currently derives its
|
||||
// turn header from `UserMessage` / `SystemItem` arrivals.
|
||||
Event::InvokeStart { .. } | Event::LlmCallStart { .. } | Event::LlmCallEnd { .. } => {
|
||||
Event::InvokeStart { .. } => {
|
||||
self.set_worker_status(WorkerStatus::Running);
|
||||
}
|
||||
// UI consumers of per-attempt LlmCall semantics remain out of scope;
|
||||
// the run-level status starts at InvokeStart and TurnStart counts each
|
||||
// LLM request within that run.
|
||||
Event::LlmCallStart { .. } | Event::LlmCallEnd { .. } => {
|
||||
self.latest_llm_wait_event = None;
|
||||
}
|
||||
Event::LlmRetry {
|
||||
@@ -1408,14 +1442,14 @@ impl App {
|
||||
self.latest_memory_worker_event = Some(event.message);
|
||||
}
|
||||
Event::Snapshot {
|
||||
entries,
|
||||
session,
|
||||
greeting,
|
||||
status,
|
||||
in_flight,
|
||||
internal_workers,
|
||||
} => {
|
||||
self.rewind_refresh_fence = false;
|
||||
self.restore_snapshot(&entries, greeting, in_flight);
|
||||
self.restore_snapshot(&session, greeting, in_flight);
|
||||
self.replace_internal_worker_snapshots(internal_workers);
|
||||
self.set_worker_status(status);
|
||||
}
|
||||
@@ -1455,11 +1489,11 @@ impl App {
|
||||
}
|
||||
}
|
||||
Event::RewindApplied {
|
||||
entries,
|
||||
session,
|
||||
input,
|
||||
summary,
|
||||
} => {
|
||||
self.restore_rewind_snapshot(&entries);
|
||||
self.restore_rewind_snapshot(&session);
|
||||
self.rewind_refresh_fence = true;
|
||||
let restored_composer = if self.input.is_empty() {
|
||||
self.input.replace_with_segments(&input);
|
||||
@@ -2173,7 +2207,7 @@ impl App {
|
||||
) -> InternalWorkerView {
|
||||
let mut app = App::new(snapshot.worker.name.clone());
|
||||
app.mode = mode;
|
||||
app.restore_entries(&snapshot.entries, None);
|
||||
app.restore_session(&snapshot.session, None);
|
||||
app.apply_in_flight_snapshot(snapshot.in_flight);
|
||||
app.set_worker_status(snapshot.status);
|
||||
if let Some(error) = snapshot.error {
|
||||
@@ -2254,14 +2288,14 @@ impl App {
|
||||
|
||||
fn restore_snapshot(
|
||||
&mut self,
|
||||
entries: &[serde_json::Value],
|
||||
session: &protocol::SessionSnapshot,
|
||||
greeting: protocol::Greeting,
|
||||
in_flight: InFlightSnapshot,
|
||||
) {
|
||||
self.greeting = Some(greeting.clone());
|
||||
self.context_window = greeting.context_window;
|
||||
self.session_context_tokens = greeting.context_tokens;
|
||||
self.restore_entries(entries, Some(greeting));
|
||||
self.restore_session(session, Some(greeting));
|
||||
self.apply_in_flight_snapshot(in_flight);
|
||||
}
|
||||
|
||||
@@ -2270,7 +2304,7 @@ impl App {
|
||||
/// session tail; always clear/replay from it even if this TUI instance has
|
||||
/// somehow lost connect-time greeting metadata. Skipping the restore in
|
||||
/// that case would leave old post-target output visible after success.
|
||||
fn restore_rewind_snapshot(&mut self, entries: &[serde_json::Value]) {
|
||||
fn restore_rewind_snapshot(&mut self, session: &protocol::SessionSnapshot) {
|
||||
let greeting = self.greeting.clone().or_else(|| {
|
||||
self.blocks.iter().find_map(|b| match b {
|
||||
Block::Greeting(g) => Some(g.clone()),
|
||||
@@ -2283,7 +2317,7 @@ impl App {
|
||||
self.session_context_tokens = greeting.context_tokens;
|
||||
}
|
||||
let missing_greeting = greeting.is_none();
|
||||
self.restore_entries(entries, greeting);
|
||||
self.restore_session(session, greeting);
|
||||
if missing_greeting {
|
||||
self.blocks.push(Block::Alert {
|
||||
level: AlertLevel::Warn,
|
||||
@@ -2293,9 +2327,9 @@ impl App {
|
||||
}
|
||||
}
|
||||
|
||||
fn restore_entries(
|
||||
fn restore_session(
|
||||
&mut self,
|
||||
entries: &[serde_json::Value],
|
||||
session: &protocol::SessionSnapshot,
|
||||
greeting: Option<protocol::Greeting>,
|
||||
) {
|
||||
self.run_error_messages.clear();
|
||||
@@ -2309,137 +2343,90 @@ impl App {
|
||||
}
|
||||
self.assistant_streaming = false;
|
||||
|
||||
for entry in entries {
|
||||
self.apply_log_entry_raw(entry);
|
||||
for entry in &session.entries {
|
||||
use protocol::{SessionContentPart, SessionMessageRole, SessionSnapshotEntryData};
|
||||
match &entry.data {
|
||||
SessionSnapshotEntryData::UserInput { segments } => {
|
||||
self.turn_index += 1;
|
||||
self.blocks.push(Block::TurnHeader {
|
||||
turn: self.turn_index,
|
||||
});
|
||||
if !segments.is_empty() {
|
||||
self.blocks.push(Block::UserMessage {
|
||||
segments: segments.clone(),
|
||||
});
|
||||
}
|
||||
}
|
||||
SessionSnapshotEntryData::Message { role, content } => {
|
||||
let role = match role {
|
||||
SessionMessageRole::User => agen::Role::User,
|
||||
SessionMessageRole::Assistant => agen::Role::Assistant,
|
||||
};
|
||||
let item = agen::Item::Message {
|
||||
id: None,
|
||||
role,
|
||||
content: content
|
||||
.iter()
|
||||
.map(|part| match part {
|
||||
SessionContentPart::Text { text } => {
|
||||
agen::ContentPart::Text { text: text.clone() }
|
||||
}
|
||||
SessionContentPart::Refusal { refusal } => {
|
||||
agen::ContentPart::Refusal {
|
||||
refusal: refusal.clone(),
|
||||
}
|
||||
}
|
||||
})
|
||||
.collect(),
|
||||
status: None,
|
||||
};
|
||||
let value = serde_json::to_value(item).expect("Item is Serialize");
|
||||
self.push_history_item(&value);
|
||||
}
|
||||
SessionSnapshotEntryData::ToolCall {
|
||||
call_id,
|
||||
name,
|
||||
arguments,
|
||||
} => {
|
||||
let item =
|
||||
agen::Item::tool_call(call_id.clone(), name.clone(), arguments.clone());
|
||||
let value = serde_json::to_value(item).expect("Item is Serialize");
|
||||
self.push_history_item(&value);
|
||||
}
|
||||
SessionSnapshotEntryData::ToolResult {
|
||||
call_id,
|
||||
summary,
|
||||
content,
|
||||
is_error,
|
||||
..
|
||||
} => {
|
||||
let item = agen::Item::tool_result_item(
|
||||
call_id.clone(),
|
||||
summary.clone(),
|
||||
content.clone(),
|
||||
*is_error,
|
||||
);
|
||||
let value = serde_json::to_value(item).expect("Item is Serialize");
|
||||
self.push_history_item(&value);
|
||||
}
|
||||
SessionSnapshotEntryData::SystemItem { data, .. } => {
|
||||
if let Some(data) = data {
|
||||
self.apply_system_item(data);
|
||||
}
|
||||
}
|
||||
SessionSnapshotEntryData::RunError { message } => {
|
||||
self.push_run_error(message.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
self.mark_orphan_tool_calls_incomplete_pass();
|
||||
}
|
||||
|
||||
/// Drop the derived view in preparation for replaying a new
|
||||
/// `SegmentStart` (compaction / fork). Greeting is preserved
|
||||
/// because the Worker identity hasn't changed.
|
||||
fn reset_for_rotation(&mut self) {
|
||||
let greeting = self.blocks.iter().find_map(|b| match b {
|
||||
Block::Greeting(g) => Some(g.clone()),
|
||||
_ => None,
|
||||
});
|
||||
self.turn_index = 0;
|
||||
self.blocks.clear();
|
||||
self.cache = FileCache::new();
|
||||
self.task_store = TaskStore::new();
|
||||
self.task_pane_scroll = 0;
|
||||
if let Some(g) = greeting {
|
||||
self.greeting = Some(g.clone());
|
||||
self.blocks.push(Block::Greeting(g));
|
||||
}
|
||||
}
|
||||
|
||||
/// Walk a single `LogEntry` JSON value and translate it into blocks
|
||||
/// the live event path would have produced. Shared between
|
||||
/// `restore_snapshot` (replay path) and `apply_log_entry` (live
|
||||
/// path).
|
||||
fn apply_log_entry_raw(&mut self, value: &serde_json::Value) {
|
||||
let Ok(entry) = serde_json::from_value::<session_store::LogEntry>(value.clone()) else {
|
||||
return;
|
||||
};
|
||||
match entry {
|
||||
session_store::LogEntry::SegmentStart { history, .. } => {
|
||||
for logged in history {
|
||||
let item: agen::Item = logged.into();
|
||||
let item_value = serde_json::to_value(&item).expect("Item is Serialize");
|
||||
self.push_history_item(&item_value);
|
||||
}
|
||||
}
|
||||
session_store::LogEntry::UserInput { segments, .. } => {
|
||||
self.turn_index += 1;
|
||||
self.blocks.push(Block::TurnHeader {
|
||||
turn: self.turn_index,
|
||||
});
|
||||
if !segments.is_empty() {
|
||||
self.blocks.push(Block::UserMessage { segments });
|
||||
}
|
||||
}
|
||||
session_store::LogEntry::AssistantItem { item, .. }
|
||||
| session_store::LogEntry::ToolResult { item, .. } => {
|
||||
let it: agen::Item = item.into();
|
||||
let item_value = serde_json::to_value(&it).expect("Item is Serialize");
|
||||
self.push_history_item(&item_value);
|
||||
}
|
||||
session_store::LogEntry::SystemItem { item, .. } => {
|
||||
let value = serde_json::to_value(&item).expect("SystemItem is Serialize");
|
||||
self.apply_system_item(&value);
|
||||
}
|
||||
session_store::LogEntry::Extension {
|
||||
domain, payload, ..
|
||||
} if domain == "yoi.compaction" => {
|
||||
self.apply_compaction_extension(&payload);
|
||||
}
|
||||
session_store::LogEntry::RunErrored { message, .. } => {
|
||||
self.push_run_error(message);
|
||||
}
|
||||
// Non-history-bearing variants don't affect the block view.
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
/// Dispatch one `SystemItem` JSON value into the appropriate block.
|
||||
///
|
||||
/// Kind-based routing replaces the old free-text `[Notification]` /
|
||||
/// `[File: …]` parsing path: each kind maps directly to a typed
|
||||
/// block (`Block::Notify`, `Block::WorkerEvent`, …).
|
||||
fn apply_compaction_extension(&mut self, payload: &serde_json::Value) {
|
||||
if payload.get("kind").and_then(|value| value.as_str()) != Some("compaction_block") {
|
||||
return;
|
||||
}
|
||||
match payload.get("state").and_then(|value| value.as_str()) {
|
||||
Some("running") => {
|
||||
if self.last_streaming_compact_mut().is_none() {
|
||||
self.blocks.push(Block::Compact(CompactEvent::Streaming {
|
||||
started_at: Instant::now(),
|
||||
}));
|
||||
}
|
||||
}
|
||||
Some("done") => {
|
||||
let new_segment_id = payload
|
||||
.get("new_segment_id")
|
||||
.and_then(|value| value.as_str())
|
||||
.and_then(|value| value.parse::<uuid::Uuid>().ok())
|
||||
.unwrap_or_else(uuid::Uuid::nil);
|
||||
if let Some(evt) = self.last_streaming_compact_mut() {
|
||||
*evt = CompactEvent::Done {
|
||||
new_segment_id,
|
||||
elapsed_secs: None,
|
||||
};
|
||||
} else {
|
||||
self.blocks.push(Block::Compact(CompactEvent::Done {
|
||||
new_segment_id,
|
||||
elapsed_secs: None,
|
||||
}));
|
||||
}
|
||||
}
|
||||
Some("failed") => {
|
||||
let error = payload
|
||||
.get("error")
|
||||
.and_then(|value| value.as_str())
|
||||
.unwrap_or("compact failed")
|
||||
.to_string();
|
||||
if let Some(evt) = self.last_streaming_compact_mut() {
|
||||
*evt = CompactEvent::Failed {
|
||||
error,
|
||||
elapsed_secs: None,
|
||||
};
|
||||
} else {
|
||||
self.blocks.push(Block::Compact(CompactEvent::Failed {
|
||||
error,
|
||||
elapsed_secs: None,
|
||||
}));
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
fn apply_system_item(&mut self, value: &serde_json::Value) {
|
||||
let Ok(item) = serde_json::from_value::<session_store::SystemItem>(value.clone()) else {
|
||||
// Unknown / forward-compat shape: fall back to rendering the
|
||||
@@ -2542,6 +2529,15 @@ fn fmt_millis(ms: u64) -> String {
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
fn public_session(values: Vec<serde_json::Value>) -> protocol::SessionSnapshot {
|
||||
let entries = values
|
||||
.into_iter()
|
||||
.map(|value| serde_json::from_value(value).expect("LogEntry deserializes"))
|
||||
.collect::<Vec<session_store::LogEntry>>();
|
||||
session_store::public_snapshot::project_current_session_snapshot(&entries)
|
||||
}
|
||||
|
||||
fn message_text(item: &serde_json::Value) -> String {
|
||||
item["content"]
|
||||
.as_array()
|
||||
@@ -2685,7 +2681,7 @@ mod rewind_refresh_tests {
|
||||
});
|
||||
|
||||
app.handle_worker_event(Event::RewindApplied {
|
||||
entries: vec![],
|
||||
session: protocol::SessionSnapshot { entries: vec![] },
|
||||
input: vec![Segment::text("selected rewind input")],
|
||||
summary: summary(3),
|
||||
});
|
||||
@@ -2704,7 +2700,7 @@ mod rewind_refresh_tests {
|
||||
});
|
||||
|
||||
app.handle_worker_event(Event::RewindApplied {
|
||||
entries: vec![],
|
||||
session: protocol::SessionSnapshot { entries: vec![] },
|
||||
input: vec![Segment::text("rewound input")],
|
||||
summary: summary(1),
|
||||
});
|
||||
@@ -2747,7 +2743,7 @@ mod rewind_refresh_tests {
|
||||
});
|
||||
|
||||
app.handle_worker_event(Event::RewindApplied {
|
||||
entries: vec![],
|
||||
session: protocol::SessionSnapshot { entries: vec![] },
|
||||
input: vec![Segment::text("rewound input")],
|
||||
summary: summary(2),
|
||||
});
|
||||
@@ -2976,6 +2972,17 @@ mod composer_history_persistence_tests {
|
||||
mod completion_flow_tests {
|
||||
use super::*;
|
||||
|
||||
fn annotated(item: agen::Item) -> session_store::LoggedHistoryEntry {
|
||||
session_store::LoggedHistoryEntry {
|
||||
item: session_store::LoggedItem::from(item),
|
||||
metadata: session_store::LoggedSessionHistoryMetadata {
|
||||
entry_id: session_store::LoggedSessionHistoryEntryId::new(),
|
||||
origin: session_store::LoggedSessionHistoryOrigin::LegacyUnknown,
|
||||
derivation: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn typing_at_creates_completion_state_and_emits_query() {
|
||||
let mut app = App::new("test".into());
|
||||
@@ -3278,7 +3285,7 @@ mod completion_flow_tests {
|
||||
#[test]
|
||||
fn committed_user_message_survives_fresh_segment_rotation() {
|
||||
let mut app = App::new("test".into());
|
||||
let start = session_store::LogEntry::SegmentStart {
|
||||
let start = session_store::LogEntry::AnnotatedSegmentStart {
|
||||
ts: session_store::segment_log::now_millis(),
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -3289,7 +3296,9 @@ mod completion_flow_tests {
|
||||
};
|
||||
|
||||
app.handle_worker_event(Event::SegmentRotated {
|
||||
entry: serde_json::to_value(start).expect("LogEntry is Serialize"),
|
||||
session: public_session(vec![
|
||||
serde_json::to_value(start).expect("LogEntry is Serialize"),
|
||||
]),
|
||||
});
|
||||
app.handle_worker_event(Event::UserMessage {
|
||||
segments: vec![Segment::text("first persisted message")],
|
||||
@@ -3403,6 +3412,17 @@ mod completion_flow_tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn running_status_starts_and_stops_live_run_clock() {
|
||||
let mut app = App::new("test".into());
|
||||
|
||||
app.set_worker_status(WorkerStatus::Running);
|
||||
assert!(app.run_started_at.is_some());
|
||||
|
||||
app.set_worker_status(WorkerStatus::Idle);
|
||||
assert!(app.run_started_at.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn running_submit_is_queued_locally_and_clears_composer() {
|
||||
let mut app = App::new("test".into());
|
||||
@@ -3533,23 +3553,23 @@ mod completion_flow_tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn snapshot_renders_system_message_block_from_session_start() {
|
||||
fn snapshot_excludes_system_prompt_history_from_public_blocks() {
|
||||
let mut app = App::new("test".into());
|
||||
let session_start = session_store::LogEntry::SegmentStart {
|
||||
let session_start = session_store::LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
config: Default::default(),
|
||||
history: vec![session_store::LoggedItem::from(
|
||||
&agen::Item::system_message("[File: src/main.rs]\nfn main() {}"),
|
||||
)],
|
||||
history: vec![annotated(agen::Item::system_message(
|
||||
"[File: src/main.rs]\nfn main() {}",
|
||||
))],
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
};
|
||||
let session_start_value = serde_json::to_value(&session_start).unwrap();
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
entries: vec![session_start_value],
|
||||
session: public_session(vec![session_start_value]),
|
||||
status: WorkerStatus::Running,
|
||||
in_flight: Default::default(),
|
||||
internal_workers: Vec::new(),
|
||||
@@ -3557,10 +3577,8 @@ mod completion_flow_tests {
|
||||
|
||||
assert!(matches!(app.worker_status, WorkerStatus::Running));
|
||||
assert!(app.running);
|
||||
assert!(matches!(
|
||||
app.blocks.get(1),
|
||||
Some(Block::SystemMessage { text }) if text == "[File: src/main.rs]\nfn main() {}"
|
||||
));
|
||||
assert_eq!(app.blocks.len(), 1);
|
||||
assert!(matches!(app.blocks.first(), Some(Block::Greeting(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -3595,7 +3613,7 @@ mod completion_flow_tests {
|
||||
};
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
entries: vec![serde_json::to_value(run_errored).unwrap()],
|
||||
session: public_session(vec![serde_json::to_value(run_errored).unwrap()]),
|
||||
status: WorkerStatus::Idle,
|
||||
in_flight: Default::default(),
|
||||
internal_workers: Vec::new(),
|
||||
@@ -3623,7 +3641,7 @@ mod completion_flow_tests {
|
||||
code: ErrorCode::ProviderError,
|
||||
message: "provider unavailable".into(),
|
||||
});
|
||||
let segment_start = session_store::LogEntry::SegmentStart {
|
||||
let segment_start = session_store::LogEntry::AnnotatedSegmentStart {
|
||||
ts: 5,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -3633,7 +3651,7 @@ mod completion_flow_tests {
|
||||
compacted_from: None,
|
||||
};
|
||||
app.handle_worker_event(Event::SegmentRotated {
|
||||
entry: serde_json::to_value(segment_start).unwrap(),
|
||||
session: public_session(vec![serde_json::to_value(segment_start).unwrap()]),
|
||||
});
|
||||
|
||||
let errors = app
|
||||
@@ -3656,7 +3674,9 @@ mod completion_flow_tests {
|
||||
let mut app = App::new("test".into());
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
entries: Vec::new(),
|
||||
session: protocol::SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
status: WorkerStatus::Running,
|
||||
in_flight: InFlightSnapshot {
|
||||
blocks: vec![
|
||||
@@ -3762,7 +3782,9 @@ mod completion_flow_tests {
|
||||
},
|
||||
revision,
|
||||
status: WorkerStatus::Idle,
|
||||
entries: Vec::new(),
|
||||
session: protocol::SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
in_flight: protocol::InFlightSnapshot::default(),
|
||||
error: None,
|
||||
internal_workers: Vec::new(),
|
||||
@@ -3977,7 +3999,9 @@ mod completion_flow_tests {
|
||||
assert_eq!(app.selected_worker_view().worker_name, "parent");
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
entries: Vec::new(),
|
||||
session: protocol::SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
status: WorkerStatus::Idle,
|
||||
in_flight: Default::default(),
|
||||
internal_workers: Vec::new(),
|
||||
@@ -4026,7 +4050,9 @@ mod completion_flow_tests {
|
||||
});
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
entries: Vec::new(),
|
||||
session: protocol::SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
status: WorkerStatus::Idle,
|
||||
in_flight: Default::default(),
|
||||
internal_workers: vec![InternalWorkerSnapshot {
|
||||
@@ -4037,7 +4063,9 @@ mod completion_flow_tests {
|
||||
kind: protocol::InternalWorkerKind::SubWorker,
|
||||
},
|
||||
revision: 4,
|
||||
entries: Vec::new(),
|
||||
session: protocol::SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
status: WorkerStatus::Running,
|
||||
error: None,
|
||||
in_flight: Default::default(),
|
||||
@@ -4193,7 +4221,9 @@ mod completion_flow_tests {
|
||||
greeting.context_tokens = 45_000;
|
||||
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
entries: Vec::new(),
|
||||
session: protocol::SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
greeting,
|
||||
status: WorkerStatus::Idle,
|
||||
in_flight: Default::default(),
|
||||
@@ -4363,40 +4393,37 @@ mod completion_flow_tests {
|
||||
});
|
||||
|
||||
let assistant_item_entries = vec![
|
||||
serde_json::json!({
|
||||
"kind": "assistant_item",
|
||||
"ts": 1,
|
||||
"item": {
|
||||
"kind": "tool_call",
|
||||
"call_id": "c1",
|
||||
"name": "TaskCreate",
|
||||
"arguments": r#"{"subject":"a","description":"A"}"#,
|
||||
},
|
||||
}),
|
||||
serde_json::json!({
|
||||
"kind": "assistant_item",
|
||||
"ts": 2,
|
||||
"item": {
|
||||
"kind": "tool_call",
|
||||
"call_id": "c2",
|
||||
"name": "TaskCreate",
|
||||
"arguments": r#"{"subject":"b","description":"B"}"#,
|
||||
},
|
||||
}),
|
||||
serde_json::json!({
|
||||
"kind": "assistant_item",
|
||||
"ts": 3,
|
||||
"item": {
|
||||
"kind": "tool_call",
|
||||
"call_id": "u1",
|
||||
"name": "TaskUpdate",
|
||||
"arguments": r#"{"taskid":2,"status":"inprogress"}"#,
|
||||
},
|
||||
}),
|
||||
serde_json::to_value(session_store::LogEntry::AnnotatedAssistantItem {
|
||||
ts: 1,
|
||||
entry: annotated(agen::Item::tool_call(
|
||||
"c1",
|
||||
"TaskCreate",
|
||||
r#"{"subject":"a","description":"A"}"#,
|
||||
)),
|
||||
})
|
||||
.unwrap(),
|
||||
serde_json::to_value(session_store::LogEntry::AnnotatedAssistantItem {
|
||||
ts: 2,
|
||||
entry: annotated(agen::Item::tool_call(
|
||||
"c2",
|
||||
"TaskCreate",
|
||||
r#"{"subject":"b","description":"B"}"#,
|
||||
)),
|
||||
})
|
||||
.unwrap(),
|
||||
serde_json::to_value(session_store::LogEntry::AnnotatedAssistantItem {
|
||||
ts: 3,
|
||||
entry: annotated(agen::Item::tool_call(
|
||||
"u1",
|
||||
"TaskUpdate",
|
||||
r#"{"taskid":2,"status":"inprogress"}"#,
|
||||
)),
|
||||
})
|
||||
.unwrap(),
|
||||
];
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
entries: assistant_item_entries,
|
||||
session: public_session(assistant_item_entries),
|
||||
status: WorkerStatus::Running,
|
||||
in_flight: Default::default(),
|
||||
internal_workers: Vec::new(),
|
||||
|
||||
@@ -7,15 +7,15 @@ use client::{
|
||||
list_backend_workers, restore_backend_worker,
|
||||
};
|
||||
use crossterm::event::{self, Event as TermEvent, KeyCode, KeyEventKind, KeyModifiers};
|
||||
use ratatui::backend::CrosstermBackend;
|
||||
use ratatui::Frame;
|
||||
use ratatui::layout::{Constraint, Layout};
|
||||
use ratatui::style::{Color, Modifier, Style};
|
||||
use ratatui::text::{Line, Span};
|
||||
use ratatui::widgets::Paragraph;
|
||||
use ratatui::{Frame, Terminal, TerminalOptions, Viewport};
|
||||
|
||||
use crate::backend_workspace_picker::select_backend_workspace;
|
||||
use crate::console;
|
||||
use crate::inline_terminal::with_inline_terminal;
|
||||
|
||||
const MAX_ROWS: usize = 10;
|
||||
const VIEWPORT_LINES: u16 = MAX_ROWS as u16 + 4;
|
||||
@@ -127,31 +127,32 @@ fn pick_worker(
|
||||
workers.truncate(MAX_ROWS);
|
||||
|
||||
let mut state = BackendWorkerPickerState::new(target, workers);
|
||||
let mut terminal = make_inline_terminal()?;
|
||||
loop {
|
||||
terminal.draw(|frame| draw(frame, &state))?;
|
||||
match poll_event()? {
|
||||
None => continue,
|
||||
Some(Action::Up) => state.previous(),
|
||||
Some(Action::Down) => state.next(),
|
||||
Some(Action::Submit) => {
|
||||
close_viewport(&mut terminal)?;
|
||||
return Ok(WorkerPickerResult::Selected(
|
||||
state.selected_worker().clone(),
|
||||
));
|
||||
with_inline_terminal(
|
||||
VIEWPORT_LINES,
|
||||
|terminal| -> Result<_, Box<dyn std::error::Error>> {
|
||||
loop {
|
||||
terminal.draw(|frame| draw(frame, &state))?;
|
||||
match poll_event()? {
|
||||
None => continue,
|
||||
Some(Action::Up) => state.previous(),
|
||||
Some(Action::Down) => state.next(),
|
||||
Some(Action::Submit) => {
|
||||
return Ok(WorkerPickerResult::Selected(
|
||||
state.selected_worker().clone(),
|
||||
));
|
||||
}
|
||||
Some(Action::SwitchWorkspace) => {
|
||||
return Ok(WorkerPickerResult::SwitchWorkspace);
|
||||
}
|
||||
Some(Action::Cancel) => {
|
||||
return Err(Box::new(io::Error::other(
|
||||
"Backend worker picker cancelled",
|
||||
)));
|
||||
}
|
||||
}
|
||||
}
|
||||
Some(Action::SwitchWorkspace) => {
|
||||
close_viewport(&mut terminal)?;
|
||||
return Ok(WorkerPickerResult::SwitchWorkspace);
|
||||
}
|
||||
Some(Action::Cancel) => {
|
||||
close_viewport(&mut terminal)?;
|
||||
return Err(Box::new(io::Error::other(
|
||||
"Backend worker picker cancelled",
|
||||
)));
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
struct BackendWorkerPickerState {
|
||||
@@ -184,27 +185,6 @@ impl BackendWorkerPickerState {
|
||||
}
|
||||
}
|
||||
|
||||
fn make_inline_terminal() -> io::Result<Terminal<CrosstermBackend<io::Stdout>>> {
|
||||
let backend = CrosstermBackend::new(io::stdout());
|
||||
Terminal::with_options(
|
||||
backend,
|
||||
TerminalOptions {
|
||||
viewport: Viewport::Inline(VIEWPORT_LINES),
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
fn close_viewport(terminal: &mut Terminal<CrosstermBackend<io::Stdout>>) -> io::Result<()> {
|
||||
let area = terminal.get_frame().area();
|
||||
let last_row = area.bottom().saturating_sub(1);
|
||||
terminal.set_cursor_position((0, last_row))?;
|
||||
use std::io::Write;
|
||||
let mut out = io::stdout();
|
||||
out.write_all(b"\r\n")?;
|
||||
out.flush()?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
enum Action {
|
||||
Up,
|
||||
Down,
|
||||
@@ -370,7 +350,7 @@ fn working_directory_text(worker: &BackendWorkerSummary) -> String {
|
||||
let cleanliness = wd.cleanliness.as_deref().unwrap_or("unknown");
|
||||
format!(
|
||||
"wd:{}:{} {} {}",
|
||||
wd.repository_id, wd.working_directory_id, wd.status, cleanliness
|
||||
wd.repository_key, wd.working_directory_id, wd.status, cleanliness
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
+886
-369
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -1,784 +0,0 @@
|
||||
use super::*;
|
||||
|
||||
pub(super) fn draw(frame: &mut Frame<'_>, app: &mut DashboardApp) {
|
||||
let area = frame.area();
|
||||
let input_content_width = area.width.saturating_sub(2).max(1);
|
||||
let mut input_render = app.input.render(input_content_width);
|
||||
let input_height = input_area_height(&input_render, area.height);
|
||||
app.input
|
||||
.apply_cursor_viewport(&mut input_render, input_height);
|
||||
let layout = dashboard_layout(area, input_height);
|
||||
|
||||
draw_title(frame, app, layout.title);
|
||||
draw_list(frame, app, layout.list);
|
||||
draw_separator(frame, layout.boundary);
|
||||
draw_target_status(frame, app, layout.target_status);
|
||||
draw_input(frame, &input_render, layout.input);
|
||||
draw_actionbar(frame, app, layout.actionbar);
|
||||
if app.panel_diagnostic_open {
|
||||
render_panel_diagnostic(frame, app, area);
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn panel_diagnostic_area(area: Rect) -> Rect {
|
||||
let width = if area.width <= 20 {
|
||||
area.width
|
||||
} else {
|
||||
area.width.saturating_sub(4).min(100).max(20)
|
||||
};
|
||||
let height = if area.height <= 8 {
|
||||
area.height
|
||||
} else {
|
||||
area.height.saturating_sub(4).min(24).max(8)
|
||||
};
|
||||
let x = area.x + area.width.saturating_sub(width) / 2;
|
||||
let y = area.y + area.height.saturating_sub(height) / 2;
|
||||
Rect::new(x, y, width, height)
|
||||
}
|
||||
|
||||
pub(super) fn render_panel_diagnostic(frame: &mut Frame<'_>, app: &DashboardApp, area: Rect) {
|
||||
let Some(diagnostic) = app.panel_diagnostic.as_ref() else {
|
||||
return;
|
||||
};
|
||||
let popup_area = panel_diagnostic_area(area);
|
||||
let title = format!(" {} ", diagnostic.title);
|
||||
let text = format!("{}\n\nF2/Esc: close", diagnostic.details);
|
||||
let paragraph = Paragraph::new(text)
|
||||
.block(Block::default().title(title).borders(Borders::ALL))
|
||||
.wrap(Wrap { trim: false });
|
||||
frame.render_widget(Clear, popup_area);
|
||||
frame.render_widget(paragraph, popup_area);
|
||||
}
|
||||
|
||||
pub(super) fn input_area_height(render: &crate::input::InputRender, terminal_height: u16) -> u16 {
|
||||
let needed = render.lines.len().max(1) as u16;
|
||||
let cap = (terminal_height / 3).max(1).min(10);
|
||||
needed.clamp(1, cap)
|
||||
}
|
||||
|
||||
pub(super) fn draw_title(frame: &mut Frame<'_>, app: &DashboardApp, area: Rect) {
|
||||
frame.render_widget(Paragraph::new(title_line(app)), area);
|
||||
}
|
||||
|
||||
pub(super) fn title_line(app: &DashboardApp) -> Line<'static> {
|
||||
let mut spans = vec![Span::styled(
|
||||
"workspace dashboard",
|
||||
Style::default().add_modifier(Modifier::BOLD),
|
||||
)];
|
||||
if let Some(companion) = &app.panel.header.companion {
|
||||
spans.push(Span::styled(
|
||||
" · companion ",
|
||||
Style::default().fg(Color::DarkGray),
|
||||
));
|
||||
spans.push(Span::styled(
|
||||
companion.status.label(),
|
||||
companion_status_style(companion.status),
|
||||
));
|
||||
if let Some(detail) = companion.detail.as_deref() {
|
||||
spans.push(Span::styled(
|
||||
format!(" ({detail})"),
|
||||
Style::default().fg(Color::DarkGray),
|
||||
));
|
||||
}
|
||||
}
|
||||
if let Some(orchestrator) = &app.panel.header.orchestrator {
|
||||
spans.push(Span::styled(
|
||||
" · orchestrator ",
|
||||
Style::default().fg(Color::DarkGray),
|
||||
));
|
||||
spans.push(Span::styled(
|
||||
orchestrator.status.label(),
|
||||
orchestrator_status_style(orchestrator.status),
|
||||
));
|
||||
}
|
||||
Line::from(spans)
|
||||
}
|
||||
|
||||
pub(super) fn companion_status_style(status: CompanionPanelStatus) -> Style {
|
||||
match status {
|
||||
CompanionPanelStatus::Live
|
||||
| CompanionPanelStatus::Restored
|
||||
| CompanionPanelStatus::Spawned => Style::default().fg(Color::Green),
|
||||
CompanionPanelStatus::Stopped | CompanionPanelStatus::Missing => {
|
||||
Style::default().fg(Color::Yellow)
|
||||
}
|
||||
CompanionPanelStatus::Unavailable => Style::default().fg(Color::Red),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn orchestrator_status_style(status: OrchestratorPanelStatus) -> Style {
|
||||
match status {
|
||||
OrchestratorPanelStatus::Live
|
||||
| OrchestratorPanelStatus::Restored
|
||||
| OrchestratorPanelStatus::Spawned => Style::default().fg(Color::Green),
|
||||
OrchestratorPanelStatus::Stopped | OrchestratorPanelStatus::Missing => {
|
||||
Style::default().fg(Color::Yellow)
|
||||
}
|
||||
OrchestratorPanelStatus::Unavailable => Style::default().fg(Color::Red),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn draw_list(frame: &mut Frame<'_>, app: &mut DashboardApp, area: Rect) {
|
||||
if area.width == 0 || area.height == 0 {
|
||||
app.row_hit_boxes.clear();
|
||||
return;
|
||||
}
|
||||
let rows = list_rows(app, area.width, area.height);
|
||||
app.set_row_hit_boxes(&rows, area);
|
||||
let lines = rows.into_iter().map(|row| row.line).collect::<Vec<_>>();
|
||||
Paragraph::new(lines).render(area, frame.buffer_mut());
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(super) struct PanelListRow {
|
||||
pub(super) line: Line<'static>,
|
||||
pub(super) key: Option<PanelRowKey>,
|
||||
}
|
||||
|
||||
impl PanelListRow {
|
||||
fn inert(line: Line<'static>) -> Self {
|
||||
Self { line, key: None }
|
||||
}
|
||||
|
||||
fn selectable(line: Line<'static>, key: PanelRowKey) -> Self {
|
||||
Self {
|
||||
line,
|
||||
key: Some(key),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(super) fn list_lines(app: &DashboardApp, width: u16, height: u16) -> Vec<Line<'static>> {
|
||||
list_rows(app, width, height)
|
||||
.into_iter()
|
||||
.map(|row| row.line)
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(super) fn list_rows(app: &DashboardApp, width: u16, height: u16) -> Vec<PanelListRow> {
|
||||
let sections = sectioned_entries(&app.list);
|
||||
let selected = app.selected_row.as_ref();
|
||||
let diagnostic_rows = panel_diagnostic_lines(&app.panel, width)
|
||||
.into_iter()
|
||||
.map(PanelListRow::inert)
|
||||
.collect::<Vec<_>>();
|
||||
let action_rows = panel_action_rows(&app.panel, selected, width);
|
||||
let live_rows = sections
|
||||
.iter()
|
||||
.filter(|section| section.kind != DashboardSectionKind::Closed)
|
||||
.flat_map(|section| section_rows(&app.list, section, selected, width))
|
||||
.collect::<Vec<_>>();
|
||||
let closed_rows = sections
|
||||
.iter()
|
||||
.find(|section| section.kind == DashboardSectionKind::Closed)
|
||||
.map(|section| section_rows(&app.list, section, selected, width))
|
||||
.unwrap_or_default();
|
||||
|
||||
let available = height as usize;
|
||||
let diagnostic_len = diagnostic_rows.len().min(available);
|
||||
let remaining_after_diagnostics = available.saturating_sub(diagnostic_len);
|
||||
let action_len = action_rows.len().min(remaining_after_diagnostics);
|
||||
let remaining_after_actions = remaining_after_diagnostics.saturating_sub(action_len);
|
||||
let closed_len = closed_rows.len().min(remaining_after_actions);
|
||||
let live_len = live_rows
|
||||
.len()
|
||||
.min(remaining_after_actions.saturating_sub(closed_len));
|
||||
let spacer_len = available.saturating_sub(diagnostic_len + action_len + live_len + closed_len);
|
||||
|
||||
let mut rows = Vec::with_capacity(available);
|
||||
rows.extend(diagnostic_rows.into_iter().take(diagnostic_len));
|
||||
rows.extend(action_rows.into_iter().take(action_len));
|
||||
rows.extend(live_rows.into_iter().take(live_len));
|
||||
rows.extend(
|
||||
std::iter::repeat_with(|| PanelListRow::inert(Line::from(Span::raw("")))).take(spacer_len),
|
||||
);
|
||||
rows.extend(closed_rows.into_iter().take(closed_len));
|
||||
rows
|
||||
}
|
||||
|
||||
pub(super) fn row_hit_boxes(rows: &[PanelListRow], area: Rect) -> Vec<PanelRowHitBox> {
|
||||
if area.width == 0 || area.height == 0 {
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
let mut hit_boxes: Vec<PanelRowHitBox> = Vec::new();
|
||||
for (offset, row) in rows.iter().enumerate() {
|
||||
let Some(key) = row.key.clone() else {
|
||||
continue;
|
||||
};
|
||||
let Some(y) = area.y.checked_add(offset as u16) else {
|
||||
continue;
|
||||
};
|
||||
if y >= area.y.saturating_add(area.height) {
|
||||
continue;
|
||||
}
|
||||
if let Some(last) = hit_boxes.last_mut() {
|
||||
if last.key == key
|
||||
&& last.rect.x == area.x
|
||||
&& last.rect.width == area.width
|
||||
&& last.rect.y.saturating_add(last.rect.height) == y
|
||||
{
|
||||
last.rect.height = last.rect.height.saturating_add(1);
|
||||
continue;
|
||||
}
|
||||
}
|
||||
hit_boxes.push(PanelRowHitBox {
|
||||
rect: Rect::new(area.x, y, area.width, 1),
|
||||
key,
|
||||
});
|
||||
}
|
||||
hit_boxes
|
||||
}
|
||||
|
||||
pub(super) fn panel_diagnostic_lines(
|
||||
panel: &WorkspacePanelViewModel,
|
||||
width: u16,
|
||||
) -> Vec<Line<'static>> {
|
||||
panel
|
||||
.header
|
||||
.diagnostics
|
||||
.iter()
|
||||
.map(|diagnostic| {
|
||||
Line::from(vec![
|
||||
Span::styled("⚠ ", Style::default().fg(Color::Yellow)),
|
||||
Span::styled(
|
||||
truncate_with_ellipsis(diagnostic, width.saturating_sub(2) as usize),
|
||||
Style::default().fg(Color::Yellow),
|
||||
),
|
||||
])
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(super) fn panel_action_rows(
|
||||
panel: &WorkspacePanelViewModel,
|
||||
selected: Option<&PanelRowKey>,
|
||||
width: u16,
|
||||
) -> Vec<PanelListRow> {
|
||||
let rows = panel
|
||||
.rows
|
||||
.iter()
|
||||
.filter(|row| row.is_ticket_section_row())
|
||||
.collect::<Vec<_>>();
|
||||
if rows.is_empty() {
|
||||
return Vec::new();
|
||||
}
|
||||
let mut lines = Vec::with_capacity((rows.len() * 2) + 1);
|
||||
lines.push(PanelListRow::inert(panel_action_header_line(
|
||||
rows.len(),
|
||||
width,
|
||||
)));
|
||||
for row in rows {
|
||||
for line in panel_row_lines(row, selected == Some(&row.key), width) {
|
||||
lines.push(PanelListRow::selectable(line, row.key.clone()));
|
||||
}
|
||||
}
|
||||
lines
|
||||
}
|
||||
|
||||
pub(super) fn panel_action_header_line(total: usize, width: u16) -> Line<'static> {
|
||||
let detail = if total == 1 {
|
||||
" 1 row".to_string()
|
||||
} else {
|
||||
format!(" {total} rows")
|
||||
};
|
||||
let text = truncate_with_ellipsis(&format!("--tickets{detail}---"), width as usize);
|
||||
Line::from(Span::styled(
|
||||
text,
|
||||
Style::default()
|
||||
.fg(Color::DarkGray)
|
||||
.add_modifier(Modifier::BOLD),
|
||||
))
|
||||
}
|
||||
|
||||
pub(super) const TICKET_STATE_COLUMN_WIDTH: usize = 10;
|
||||
pub(super) const POD_STATUS_COLUMN_WIDTH: usize = 18;
|
||||
|
||||
pub(super) fn panel_row_lines(row: &PanelRow, selected: bool, width: u16) -> Vec<Line<'static>> {
|
||||
if row.kind == PanelRowKind::TicketIntakeWorker {
|
||||
vec![panel_intake_child_line(row, selected, width)]
|
||||
} else {
|
||||
vec![
|
||||
panel_row_title_line(row, selected, width),
|
||||
panel_row_detail_line(row, selected, width),
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn panel_row_title_line(row: &PanelRow, selected: bool, width: u16) -> Line<'static> {
|
||||
let title_style = if selected {
|
||||
Style::default()
|
||||
.fg(Color::Magenta)
|
||||
.add_modifier(Modifier::BOLD)
|
||||
} else {
|
||||
Style::default().fg(Color::Magenta)
|
||||
};
|
||||
let mut spans = Vec::new();
|
||||
let mut remaining = width as usize;
|
||||
|
||||
push_ticket_primary_marker_span(&mut spans, selected, &mut remaining);
|
||||
push_column_span(
|
||||
&mut spans,
|
||||
&row.status,
|
||||
TICKET_STATE_COLUMN_WIDTH,
|
||||
panel_priority_style(row.priority),
|
||||
&mut remaining,
|
||||
);
|
||||
push_bounded_span(&mut spans, row.title.as_str(), title_style, &mut remaining);
|
||||
|
||||
Line::from(spans)
|
||||
}
|
||||
|
||||
pub(super) fn panel_intake_child_line(row: &PanelRow, selected: bool, width: u16) -> Line<'static> {
|
||||
let title_style = if selected {
|
||||
Style::default()
|
||||
.fg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD)
|
||||
} else {
|
||||
Style::default().fg(Color::Cyan)
|
||||
};
|
||||
let mut spans = Vec::new();
|
||||
let mut remaining = width as usize;
|
||||
|
||||
push_intake_child_marker_span(&mut spans, selected, &mut remaining);
|
||||
push_column_span(
|
||||
&mut spans,
|
||||
&row.status,
|
||||
TICKET_STATE_COLUMN_WIDTH,
|
||||
intake_status_style(&row.status),
|
||||
&mut remaining,
|
||||
);
|
||||
push_bounded_span(&mut spans, row.title.as_str(), title_style, &mut remaining);
|
||||
|
||||
Line::from(spans)
|
||||
}
|
||||
|
||||
pub(super) fn panel_row_detail_line(row: &PanelRow, selected: bool, width: u16) -> Line<'static> {
|
||||
let mut spans = Vec::new();
|
||||
let mut remaining = width as usize;
|
||||
|
||||
push_ticket_detail_marker_span(&mut spans, selected, &mut remaining);
|
||||
push_bounded_span(
|
||||
&mut spans,
|
||||
"meta ",
|
||||
Style::default().fg(Color::DarkGray),
|
||||
&mut remaining,
|
||||
);
|
||||
push_bounded_span(
|
||||
&mut spans,
|
||||
&panel_ticket_detail(row),
|
||||
ticket_detail_style(row),
|
||||
&mut remaining,
|
||||
);
|
||||
|
||||
Line::from(spans)
|
||||
}
|
||||
|
||||
pub(super) fn push_ticket_primary_marker_span(
|
||||
spans: &mut Vec<Span<'static>>,
|
||||
selected: bool,
|
||||
remaining: &mut usize,
|
||||
) {
|
||||
let (marker, style) = if selected {
|
||||
(
|
||||
"▶ ",
|
||||
Style::default()
|
||||
.fg(Color::Magenta)
|
||||
.add_modifier(Modifier::BOLD),
|
||||
)
|
||||
} else {
|
||||
(" ", Style::default().fg(Color::DarkGray))
|
||||
};
|
||||
push_bounded_span(spans, marker, style, remaining);
|
||||
}
|
||||
|
||||
pub(super) fn push_ticket_detail_marker_span(
|
||||
spans: &mut Vec<Span<'static>>,
|
||||
selected: bool,
|
||||
remaining: &mut usize,
|
||||
) {
|
||||
let (marker, style) = if selected {
|
||||
(
|
||||
"│ ",
|
||||
Style::default()
|
||||
.fg(Color::Magenta)
|
||||
.add_modifier(Modifier::BOLD),
|
||||
)
|
||||
} else {
|
||||
(" ", Style::default().fg(Color::DarkGray))
|
||||
};
|
||||
push_bounded_span(spans, marker, style, remaining);
|
||||
}
|
||||
|
||||
pub(super) fn push_intake_child_marker_span(
|
||||
spans: &mut Vec<Span<'static>>,
|
||||
selected: bool,
|
||||
remaining: &mut usize,
|
||||
) {
|
||||
let (marker, style) = if selected {
|
||||
(
|
||||
" ▶ ",
|
||||
Style::default()
|
||||
.fg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD),
|
||||
)
|
||||
} else {
|
||||
(" └ ", Style::default().fg(Color::DarkGray))
|
||||
};
|
||||
push_bounded_span(spans, marker, style, remaining);
|
||||
}
|
||||
|
||||
pub(super) fn panel_ticket_detail(row: &PanelRow) -> String {
|
||||
if row.kind == PanelRowKind::InvalidTicket {
|
||||
let mut parts = vec![panel_ticket_reference(row), "Gate: unavailable".to_string()];
|
||||
if let Some(reason) = panel_ticket_reason(row) {
|
||||
parts.push(format!("Reason: {reason}"));
|
||||
}
|
||||
return parts.join(" · ");
|
||||
}
|
||||
|
||||
if row.kind == PanelRowKind::TicketIntakeWorker {
|
||||
let mut parts = row
|
||||
.subtitle
|
||||
.as_ref()
|
||||
.map(|subtitle| vec![subtitle.clone()])
|
||||
.unwrap_or_else(|| vec![panel_ticket_reference(row)]);
|
||||
if let Some(action) = row.next_action {
|
||||
parts.push(format!("Action: {}", action.label()));
|
||||
}
|
||||
if let Some(reason) = panel_ticket_reason(row) {
|
||||
parts.push(format!("Reason: {reason}"));
|
||||
}
|
||||
return parts.join(" · ");
|
||||
}
|
||||
|
||||
let mut parts = vec![panel_ticket_reference(row)];
|
||||
if let Some(overlay_detail) = panel_ticket_overlay_detail(row) {
|
||||
parts.push(overlay_detail);
|
||||
}
|
||||
if let Some(blocked_reason) = row
|
||||
.ticket
|
||||
.as_ref()
|
||||
.and_then(|ticket| ticket.blocked_reason.as_deref())
|
||||
{
|
||||
parts.push(format!("Dependencies: {blocked_reason}"));
|
||||
} else {
|
||||
parts.push("Gate: clear".to_string());
|
||||
}
|
||||
if let Some(action) = row.next_action {
|
||||
parts.push(format!(
|
||||
"Action: {}",
|
||||
panel_ticket_action_label(row, action)
|
||||
));
|
||||
}
|
||||
if let Some(reason) = panel_ticket_reason(row) {
|
||||
parts.push(format!("Reason: {reason}"));
|
||||
}
|
||||
parts.join(" · ")
|
||||
}
|
||||
|
||||
pub(super) fn panel_ticket_action_label(row: &PanelRow, action: NextUserAction) -> &'static str {
|
||||
if action == NextUserAction::Wait
|
||||
&& row
|
||||
.ticket
|
||||
.as_ref()
|
||||
.and_then(|ticket| ticket.blocked_reason.as_ref())
|
||||
.is_some()
|
||||
{
|
||||
"queue disabled"
|
||||
} else {
|
||||
action.label()
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn panel_ticket_overlay_detail(row: &PanelRow) -> Option<String> {
|
||||
let ticket = row.ticket.as_ref()?;
|
||||
let overlay = ticket.orchestration_overlay.as_ref()?;
|
||||
let mut detail = format!(
|
||||
"Overlay: local {} · {} {}",
|
||||
ticket.workflow_state.as_str(),
|
||||
overlay.source,
|
||||
overlay.workflow_state.as_str()
|
||||
);
|
||||
if matches!(
|
||||
overlay.workflow_state,
|
||||
TicketWorkflowState::Done | TicketWorkflowState::Closed
|
||||
) {
|
||||
detail.push_str(" · merge pending");
|
||||
}
|
||||
Some(detail)
|
||||
}
|
||||
|
||||
pub(super) fn panel_ticket_reason(row: &PanelRow) -> Option<&str> {
|
||||
row.disabled_reason
|
||||
.as_deref()
|
||||
.or_else(|| row.key_hint.as_deref())
|
||||
}
|
||||
|
||||
pub(super) fn ticket_detail_style(row: &PanelRow) -> Style {
|
||||
if row.kind == PanelRowKind::InvalidTicket {
|
||||
return Style::default().fg(Color::Yellow);
|
||||
}
|
||||
if row
|
||||
.ticket
|
||||
.as_ref()
|
||||
.and_then(|ticket| ticket.blocked_reason.as_ref())
|
||||
.is_some()
|
||||
{
|
||||
Style::default().fg(Color::Yellow)
|
||||
} else {
|
||||
Style::default().fg(Color::DarkGray)
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn panel_ticket_reference(row: &PanelRow) -> String {
|
||||
row.ticket
|
||||
.as_ref()
|
||||
.map(|ticket| {
|
||||
ticket
|
||||
.resource_key
|
||||
.clone()
|
||||
.unwrap_or_else(|| "resource key unavailable".to_string())
|
||||
})
|
||||
.unwrap_or_else(|| match &row.key {
|
||||
PanelRowKey::Ticket(id) | PanelRowKey::InvalidTicket(id) => id.clone(),
|
||||
PanelRowKey::TicketIntakeWorker { ticket_id, .. } => ticket_id.clone(),
|
||||
PanelRowKey::Worker(name) => name.clone(),
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn push_column_span(
|
||||
spans: &mut Vec<Span<'static>>,
|
||||
value: &str,
|
||||
column_width: usize,
|
||||
style: Style,
|
||||
remaining: &mut usize,
|
||||
) {
|
||||
if *remaining == 0 {
|
||||
return;
|
||||
}
|
||||
let mut content = padded_cell(value, column_width);
|
||||
content.push(' ');
|
||||
push_bounded_span(spans, &content, style, remaining);
|
||||
}
|
||||
|
||||
pub(super) fn push_bounded_span(
|
||||
spans: &mut Vec<Span<'static>>,
|
||||
value: &str,
|
||||
style: Style,
|
||||
remaining: &mut usize,
|
||||
) {
|
||||
if *remaining == 0 || value.is_empty() {
|
||||
return;
|
||||
}
|
||||
let content = truncate_with_ellipsis(value, *remaining);
|
||||
*remaining = remaining.saturating_sub(content.width());
|
||||
spans.push(Span::styled(content, style));
|
||||
}
|
||||
|
||||
pub(super) fn padded_cell(value: &str, width: usize) -> String {
|
||||
let mut cell = truncate_with_ellipsis(value, width);
|
||||
let padding = width.saturating_sub(cell.width());
|
||||
cell.extend(std::iter::repeat_n(' ', padding));
|
||||
cell
|
||||
}
|
||||
|
||||
pub(super) fn panel_priority_style(priority: ActionPriority) -> Style {
|
||||
match priority {
|
||||
ActionPriority::ReadyForQueue => Style::default().fg(Color::Green),
|
||||
ActionPriority::ActiveWork => Style::default().fg(Color::Cyan),
|
||||
ActionPriority::Background => Style::default().fg(Color::DarkGray),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn intake_status_style(status: &str) -> Style {
|
||||
match status {
|
||||
"live" => Style::default().fg(Color::Green),
|
||||
"restorable" => Style::default().fg(Color::Yellow),
|
||||
"stale" => Style::default().fg(Color::DarkGray),
|
||||
_ => Style::default().fg(Color::Cyan),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn section_rows(
|
||||
list: &WorkerList,
|
||||
section: &DashboardSection,
|
||||
selected: Option<&PanelRowKey>,
|
||||
width: u16,
|
||||
) -> Vec<PanelListRow> {
|
||||
let visible = visible_section_indices(section);
|
||||
if visible.is_empty() {
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
let mut rows = Vec::with_capacity(visible.len() + 1);
|
||||
rows.push(PanelListRow::inert(section_header_line(
|
||||
section.kind,
|
||||
section.entries.len(),
|
||||
section.hidden_count(),
|
||||
width,
|
||||
)));
|
||||
for index in visible {
|
||||
if let Some(entry) = list.entries.get(index) {
|
||||
let key = PanelRowKey::Worker(entry.name.clone());
|
||||
let selected = selected == Some(&key);
|
||||
rows.push(PanelListRow::selectable(
|
||||
row_line(entry, selected, width),
|
||||
key,
|
||||
));
|
||||
}
|
||||
}
|
||||
rows
|
||||
}
|
||||
|
||||
pub(super) fn row_line(entry: &WorkerListEntry, selected: bool, width: u16) -> Line<'static> {
|
||||
let marker = if selected { "▶ " } else { " " };
|
||||
let name_style = if selected {
|
||||
Style::default()
|
||||
.fg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD)
|
||||
} else {
|
||||
Style::default().fg(Color::Cyan)
|
||||
};
|
||||
let (status, status_style) = row_status_label(entry);
|
||||
let mut spans = Vec::new();
|
||||
let mut remaining = width as usize;
|
||||
|
||||
push_bounded_span(
|
||||
&mut spans,
|
||||
marker,
|
||||
if selected {
|
||||
Style::default()
|
||||
.fg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD)
|
||||
} else {
|
||||
Style::default().fg(Color::DarkGray)
|
||||
},
|
||||
&mut remaining,
|
||||
);
|
||||
push_column_span(
|
||||
&mut spans,
|
||||
status,
|
||||
POD_STATUS_COLUMN_WIDTH,
|
||||
status_style,
|
||||
&mut remaining,
|
||||
);
|
||||
push_bounded_span(&mut spans, entry.name.as_str(), name_style, &mut remaining);
|
||||
|
||||
Line::from(spans)
|
||||
}
|
||||
|
||||
pub(super) fn draw_separator(frame: &mut Frame<'_>, area: Rect) {
|
||||
frame.render_widget(
|
||||
Paragraph::new(Line::from(Span::styled(
|
||||
"─".repeat(area.width as usize),
|
||||
Style::default().fg(Color::DarkGray),
|
||||
))),
|
||||
area,
|
||||
);
|
||||
}
|
||||
|
||||
pub(super) fn draw_target_status(frame: &mut Frame<'_>, app: &DashboardApp, area: Rect) {
|
||||
frame.render_widget(Paragraph::new(target_status_line(app)), area);
|
||||
}
|
||||
|
||||
pub(super) fn target_status_line(_app: &DashboardApp) -> Line<'static> {
|
||||
Line::from(Span::raw(""))
|
||||
}
|
||||
|
||||
pub(super) fn draw_input(frame: &mut Frame<'_>, render: &crate::input::InputRender, area: Rect) {
|
||||
let mut lines: Vec<Line<'static>> = Vec::with_capacity(render.lines.len());
|
||||
for (i, src) in render.lines.iter().enumerate() {
|
||||
let absolute_row = render.viewport_start_row as usize + i;
|
||||
let prefix = if absolute_row == 0 { "> " } else { " " };
|
||||
let mut spans = vec![Span::styled(prefix, Style::default().fg(Color::DarkGray))];
|
||||
spans.extend(src.spans.iter().cloned());
|
||||
lines.push(Line::from(spans));
|
||||
}
|
||||
frame.render_widget(Paragraph::new(lines), area);
|
||||
|
||||
let cursor_x = area.x + 2 + render.cursor_col;
|
||||
let cursor_y = area.y + render.cursor_row;
|
||||
if cursor_y < area.y + area.height {
|
||||
frame.set_cursor_position(Position::new(cursor_x, cursor_y));
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn actionbar_left_text(app: &DashboardApp) -> String {
|
||||
if app.sending && app.composer_target() == ComposerTarget::TicketIntake {
|
||||
"launching Ticket Intake…".to_string()
|
||||
} else if app.sending {
|
||||
"working…".to_string()
|
||||
} else if app.refreshing {
|
||||
match app.notice.as_deref() {
|
||||
Some(notice) if notice.contains("Refreshing") || notice.contains("refreshing") => {
|
||||
notice.to_string()
|
||||
}
|
||||
Some(notice) => format!("{notice} Refreshing workspace…"),
|
||||
None => "Refreshing workspace…".to_string(),
|
||||
}
|
||||
} else if let Some(notice) = app.notice.as_deref() {
|
||||
notice.to_string()
|
||||
} else {
|
||||
String::new()
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn actionbar_right_text(app: &DashboardApp) -> &'static str {
|
||||
if app.panel_diagnostic_open {
|
||||
"F2/Esc close details"
|
||||
} else if app.panel_diagnostic.is_some() {
|
||||
"F2 details"
|
||||
} else {
|
||||
""
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn draw_actionbar(frame: &mut Frame<'_>, app: &DashboardApp, area: Rect) {
|
||||
let left = actionbar_left_text(app);
|
||||
let right = actionbar_right_text(app);
|
||||
let left_width = area
|
||||
.width
|
||||
.saturating_sub(right.width() as u16)
|
||||
.saturating_sub(2) as usize;
|
||||
frame.render_widget(
|
||||
Paragraph::new(Line::from(Span::styled(
|
||||
truncate_with_ellipsis(&left, left_width),
|
||||
Style::default().fg(Color::DarkGray),
|
||||
))),
|
||||
area,
|
||||
);
|
||||
frame.render_widget(
|
||||
Paragraph::new(Line::from(Span::styled(
|
||||
right,
|
||||
Style::default().fg(Color::DarkGray),
|
||||
)))
|
||||
.alignment(ratatui::layout::Alignment::Right),
|
||||
area,
|
||||
);
|
||||
}
|
||||
|
||||
pub(super) fn truncate_with_ellipsis(s: &str, max_width: usize) -> String {
|
||||
if max_width == 0 {
|
||||
return String::new();
|
||||
}
|
||||
if s.width() <= max_width {
|
||||
return s.to_string();
|
||||
}
|
||||
if max_width == 1 {
|
||||
return "…".to_string();
|
||||
}
|
||||
let mut out = String::new();
|
||||
let mut width = 0usize;
|
||||
for c in s.chars() {
|
||||
let cw = unicode_width::UnicodeWidthChar::width(c).unwrap_or(0);
|
||||
if width + cw > max_width - 1 {
|
||||
break;
|
||||
}
|
||||
out.push(c);
|
||||
width += cw;
|
||||
}
|
||||
out.push('…');
|
||||
out
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,109 @@
|
||||
use std::io::{self, Stdout, Write};
|
||||
|
||||
use ratatui::Terminal;
|
||||
use ratatui::backend::CrosstermBackend;
|
||||
use ratatui::{TerminalOptions, Viewport};
|
||||
|
||||
pub(crate) type InlineTerminal = Terminal<CrosstermBackend<Stdout>>;
|
||||
|
||||
struct InlineTerminalGuard {
|
||||
terminal: InlineTerminal,
|
||||
closed: bool,
|
||||
}
|
||||
|
||||
impl InlineTerminalGuard {
|
||||
fn open(height: u16) -> io::Result<Self> {
|
||||
let terminal = Terminal::with_options(
|
||||
CrosstermBackend::new(io::stdout()),
|
||||
TerminalOptions {
|
||||
viewport: Viewport::Inline(height),
|
||||
},
|
||||
)?;
|
||||
Ok(Self {
|
||||
terminal,
|
||||
closed: false,
|
||||
})
|
||||
}
|
||||
|
||||
fn close(&mut self) -> io::Result<()> {
|
||||
if self.closed {
|
||||
return Ok(());
|
||||
}
|
||||
self.closed = true;
|
||||
|
||||
let area = self.terminal.get_frame().area();
|
||||
let last_row = area.bottom().saturating_sub(1);
|
||||
let cursor_result = self.terminal.set_cursor_position((0, last_row));
|
||||
let output_result = write_viewport_terminator(&mut io::stdout());
|
||||
cursor_result?;
|
||||
output_result
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for InlineTerminalGuard {
|
||||
fn drop(&mut self) {
|
||||
let _ = self.close();
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn with_inline_terminal<T, E>(
|
||||
height: u16,
|
||||
run: impl FnOnce(&mut InlineTerminal) -> Result<T, E>,
|
||||
) -> Result<T, E>
|
||||
where
|
||||
E: From<io::Error>,
|
||||
{
|
||||
let mut guard = InlineTerminalGuard::open(height).map_err(E::from)?;
|
||||
let result = run(&mut guard.terminal);
|
||||
let close_result = guard.close();
|
||||
match result {
|
||||
Ok(value) => {
|
||||
close_result.map_err(E::from)?;
|
||||
Ok(value)
|
||||
}
|
||||
Err(error) => Err(error),
|
||||
}
|
||||
}
|
||||
|
||||
fn write_viewport_terminator(output: &mut impl Write) -> io::Result<()> {
|
||||
output.write_all(b"\r\n")?;
|
||||
output.flush()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn viewport_terminator_moves_following_output_to_a_fresh_line() {
|
||||
let mut output = Vec::new();
|
||||
|
||||
write_viewport_terminator(&mut output).unwrap();
|
||||
|
||||
assert_eq!(output, b"\r\n");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn inline_viewport_construction_is_owned_by_this_module() {
|
||||
fn assert_shared_owner(path: &std::path::Path) {
|
||||
for entry in std::fs::read_dir(path).unwrap() {
|
||||
let path = entry.unwrap().path();
|
||||
if path.is_dir() {
|
||||
assert_shared_owner(&path);
|
||||
} else if path.extension().and_then(|value| value.to_str()) == Some("rs")
|
||||
&& path.file_name().and_then(|value| value.to_str())
|
||||
!= Some("inline_terminal.rs")
|
||||
{
|
||||
let source = std::fs::read_to_string(&path).unwrap();
|
||||
assert!(
|
||||
!source.contains("Viewport::Inline"),
|
||||
"{} constructs an inline viewport outside its shared owner",
|
||||
path.display()
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
assert_shared_owner(&std::path::Path::new(env!("CARGO_MANIFEST_DIR")).join("src"));
|
||||
}
|
||||
}
|
||||
+388
-74
@@ -15,6 +15,64 @@ use ratatui::style::{Color, Style};
|
||||
use ratatui::text::{Line, Span};
|
||||
use unicode_width::UnicodeWidthChar;
|
||||
|
||||
pub const MAX_PLAIN_TEXT_PASTE_CHARS: usize = 50;
|
||||
pub const MAX_PLAIN_TEXT_PASTE_LOGICAL_LINES: usize = 3;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct PasteMeasurement {
|
||||
pub chars: usize,
|
||||
pub logical_lines: usize,
|
||||
}
|
||||
|
||||
impl PasteMeasurement {
|
||||
pub fn presentation(self) -> PastePresentation {
|
||||
if self.chars <= MAX_PLAIN_TEXT_PASTE_CHARS
|
||||
&& self.logical_lines <= MAX_PLAIN_TEXT_PASTE_LOGICAL_LINES
|
||||
{
|
||||
PastePresentation::Text
|
||||
} else {
|
||||
PastePresentation::Chip
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum PastePresentation {
|
||||
Text,
|
||||
Chip,
|
||||
}
|
||||
|
||||
pub fn measure_paste(content: &str) -> PasteMeasurement {
|
||||
PasteMeasurement {
|
||||
chars: content.chars().count(),
|
||||
logical_lines: logical_line_count(content),
|
||||
}
|
||||
}
|
||||
|
||||
/// Empty content has zero logical lines. Otherwise LF, lone CR, and CRLF each
|
||||
/// advance one line; a CRLF pair is one break rather than two.
|
||||
pub fn logical_line_count(content: &str) -> usize {
|
||||
if content.is_empty() {
|
||||
return 0;
|
||||
}
|
||||
|
||||
let mut lines = 1;
|
||||
let mut chars = content.chars().peekable();
|
||||
while let Some(ch) = chars.next() {
|
||||
match ch {
|
||||
'\r' => {
|
||||
if chars.peek() == Some(&'\n') {
|
||||
chars.next();
|
||||
}
|
||||
lines += 1;
|
||||
}
|
||||
'\n' => lines += 1,
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
lines
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct PasteRef {
|
||||
pub id: u32,
|
||||
@@ -61,6 +119,7 @@ impl FlowRefAtom {
|
||||
pub enum Atom {
|
||||
Char(char),
|
||||
Paste(PasteRef),
|
||||
PasteArtifact(protocol::PasteArtifactRef),
|
||||
FileRef(FileRefAtom),
|
||||
FlowRef(FlowRefAtom),
|
||||
}
|
||||
@@ -72,6 +131,18 @@ impl Atom {
|
||||
match self {
|
||||
Atom::Char(_) => None,
|
||||
Atom::Paste(p) => Some((Style::default().fg(Color::Magenta), p.label())),
|
||||
Atom::PasteArtifact(artifact) => Some((
|
||||
Style::default().fg(Color::Magenta),
|
||||
format!(
|
||||
"[Paste artifact {} | {} chars, {} lines, {}, {}, created {} ms]",
|
||||
artifact.artifact_id,
|
||||
artifact.char_count,
|
||||
artifact.line_count,
|
||||
artifact.media_type.as_str(),
|
||||
artifact.availability.as_str(),
|
||||
artifact.created_at_ms
|
||||
),
|
||||
)),
|
||||
Atom::FileRef(r) => Some((Style::default().fg(Color::Cyan), r.label())),
|
||||
Atom::FlowRef(r) => Some((Style::default().fg(Color::Yellow), r.label())),
|
||||
}
|
||||
@@ -102,7 +173,9 @@ enum WordKind {
|
||||
fn atom_class(atom: &Atom) -> AtomClass {
|
||||
match atom {
|
||||
Atom::Char(c) => char_class(*c),
|
||||
Atom::Paste(_) | Atom::FileRef(_) | Atom::FlowRef(_) => AtomClass::Chip,
|
||||
Atom::Paste(_) | Atom::PasteArtifact(_) | Atom::FileRef(_) | Atom::FlowRef(_) => {
|
||||
AtomClass::Chip
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -190,6 +263,16 @@ impl InputBuffer {
|
||||
content: content.clone(),
|
||||
}));
|
||||
}
|
||||
protocol::Segment::PasteArtifact { artifact } => {
|
||||
self.atoms.push(Atom::PasteArtifact(artifact.clone()));
|
||||
}
|
||||
protocol::Segment::UploadedFile { file } => {
|
||||
self.atoms.extend(
|
||||
format!("[Attached file: {}]", file.file_name)
|
||||
.chars()
|
||||
.map(Atom::Char),
|
||||
);
|
||||
}
|
||||
protocol::Segment::FileRef { path } => {
|
||||
self.atoms
|
||||
.push(Atom::FileRef(FileRefAtom { path: path.clone() }));
|
||||
@@ -225,6 +308,13 @@ impl InputBuffer {
|
||||
match atom {
|
||||
Atom::Char(c) => text.push(*c),
|
||||
Atom::Paste(paste) => text.push_str(&paste.content),
|
||||
Atom::PasteArtifact(artifact) => {
|
||||
text.push_str(&protocol::Segment::flatten_to_text(&[
|
||||
protocol::Segment::PasteArtifact {
|
||||
artifact: artifact.clone(),
|
||||
},
|
||||
]))
|
||||
}
|
||||
Atom::FileRef(file) => text.push_str(&file.path),
|
||||
Atom::FlowRef(flow) => text.push_str(&flow.selector),
|
||||
}
|
||||
@@ -237,16 +327,20 @@ impl InputBuffer {
|
||||
}
|
||||
|
||||
pub fn insert_paste(&mut self, content: String) {
|
||||
let measurement = measure_paste(&content);
|
||||
if measurement.presentation() == PastePresentation::Text {
|
||||
self.insert_str(&content);
|
||||
return;
|
||||
}
|
||||
|
||||
let id = self.next_paste_id;
|
||||
self.next_paste_id = self.next_paste_id.wrapping_add(1);
|
||||
let chars = content.chars().count();
|
||||
let lines = content.lines().count().max(1);
|
||||
self.atoms.insert(
|
||||
self.cursor,
|
||||
Atom::Paste(PasteRef {
|
||||
id,
|
||||
chars,
|
||||
lines,
|
||||
chars: measurement.chars,
|
||||
lines: measurement.logical_lines,
|
||||
content,
|
||||
}),
|
||||
);
|
||||
@@ -395,80 +489,78 @@ impl InputBuffer {
|
||||
self.cursor = 0;
|
||||
}
|
||||
|
||||
pub fn move_home(&mut self) {
|
||||
while self.cursor > 0 {
|
||||
if matches!(self.atoms[self.cursor - 1], Atom::Char('\n')) {
|
||||
break;
|
||||
}
|
||||
self.cursor -= 1;
|
||||
fn logical_line_ranges(&self) -> Vec<(usize, usize)> {
|
||||
let mut ranges = Vec::new();
|
||||
let mut start = 0;
|
||||
let mut index = 0;
|
||||
while index < self.atoms.len() {
|
||||
let break_len = match self.atoms[index] {
|
||||
Atom::Char('\r') => {
|
||||
if matches!(self.atoms.get(index + 1), Some(Atom::Char('\n'))) {
|
||||
2
|
||||
} else {
|
||||
1
|
||||
}
|
||||
}
|
||||
Atom::Char('\n') => 1,
|
||||
_ => {
|
||||
index += 1;
|
||||
continue;
|
||||
}
|
||||
};
|
||||
ranges.push((start, index));
|
||||
index += break_len;
|
||||
start = index;
|
||||
}
|
||||
ranges.push((start, self.atoms.len()));
|
||||
ranges
|
||||
}
|
||||
|
||||
fn logical_line_and_col(&self) -> (Vec<(usize, usize)>, usize, usize) {
|
||||
let ranges = self.logical_line_ranges();
|
||||
for (line, &(start, end)) in ranges.iter().enumerate() {
|
||||
if self.cursor <= end {
|
||||
return (ranges, line, self.cursor.saturating_sub(start));
|
||||
}
|
||||
if let Some(&(next_start, _)) = ranges.get(line + 1)
|
||||
&& self.cursor < next_start
|
||||
{
|
||||
return (ranges, line + 1, 0);
|
||||
}
|
||||
}
|
||||
let line = ranges.len().saturating_sub(1);
|
||||
let col = self.cursor.saturating_sub(ranges[line].0);
|
||||
(ranges, line, col)
|
||||
}
|
||||
|
||||
pub fn move_home(&mut self) {
|
||||
let (ranges, line, _) = self.logical_line_and_col();
|
||||
self.cursor = ranges[line].0;
|
||||
}
|
||||
|
||||
pub fn move_end(&mut self) {
|
||||
while self.cursor < self.atoms.len() {
|
||||
if matches!(self.atoms[self.cursor], Atom::Char('\n')) {
|
||||
break;
|
||||
}
|
||||
self.cursor += 1;
|
||||
}
|
||||
let (ranges, line, _) = self.logical_line_and_col();
|
||||
self.cursor = ranges[line].1;
|
||||
}
|
||||
|
||||
/// Move one logical line up, preserving column (atom count from
|
||||
/// current line start). No-op if already on the first line.
|
||||
pub fn move_up(&mut self) {
|
||||
let (line_start, col) = self.line_start_and_col();
|
||||
if line_start == 0 {
|
||||
let (ranges, line, col) = self.logical_line_and_col();
|
||||
if line == 0 {
|
||||
return;
|
||||
}
|
||||
// `atoms[line_start - 1]` is the '\n' that opens the current
|
||||
// line; find the previous line's start.
|
||||
let prev_end = line_start - 1;
|
||||
let mut prev_start = 0;
|
||||
for i in (0..prev_end).rev() {
|
||||
if matches!(self.atoms[i], Atom::Char('\n')) {
|
||||
prev_start = i + 1;
|
||||
break;
|
||||
}
|
||||
}
|
||||
let prev_len = prev_end - prev_start;
|
||||
self.cursor = prev_start + col.min(prev_len);
|
||||
let (start, end) = ranges[line - 1];
|
||||
self.cursor = start + col.min(end - start);
|
||||
}
|
||||
|
||||
/// Move one logical line down, preserving column.
|
||||
pub fn move_down(&mut self) {
|
||||
let (line_start, col) = self.line_start_and_col();
|
||||
// End of current line.
|
||||
let mut cur_end = self.atoms.len();
|
||||
for i in line_start..self.atoms.len() {
|
||||
if matches!(self.atoms[i], Atom::Char('\n')) {
|
||||
cur_end = i;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if cur_end == self.atoms.len() {
|
||||
return; // no next line
|
||||
}
|
||||
let next_start = cur_end + 1;
|
||||
let mut next_end = self.atoms.len();
|
||||
for i in next_start..self.atoms.len() {
|
||||
if matches!(self.atoms[i], Atom::Char('\n')) {
|
||||
next_end = i;
|
||||
break;
|
||||
}
|
||||
}
|
||||
let next_len = next_end - next_start;
|
||||
self.cursor = next_start + col.min(next_len);
|
||||
}
|
||||
|
||||
fn line_start_and_col(&self) -> (usize, usize) {
|
||||
let mut start = 0;
|
||||
for i in (0..self.cursor).rev() {
|
||||
if matches!(self.atoms[i], Atom::Char('\n')) {
|
||||
start = i + 1;
|
||||
break;
|
||||
}
|
||||
}
|
||||
(start, self.cursor - start)
|
||||
let (ranges, line, col) = self.logical_line_and_col();
|
||||
let Some(&(start, end)) = ranges.get(line + 1) else {
|
||||
return;
|
||||
};
|
||||
self.cursor = start + col.min(end - start);
|
||||
}
|
||||
|
||||
/// Build the typed `Vec<Segment>` sent over the protocol. Adjacent
|
||||
@@ -497,6 +589,12 @@ impl InputBuffer {
|
||||
content: p.content.clone(),
|
||||
});
|
||||
}
|
||||
Atom::PasteArtifact(artifact) => {
|
||||
flush_text(&mut buf, &mut out);
|
||||
out.push(protocol::Segment::PasteArtifact {
|
||||
artifact: artifact.clone(),
|
||||
});
|
||||
}
|
||||
Atom::FileRef(r) => {
|
||||
flush_text(&mut buf, &mut out);
|
||||
out.push(protocol::Segment::FileRef {
|
||||
@@ -535,6 +633,7 @@ impl InputBuffer {
|
||||
let mut cursor_row: u16 = 0;
|
||||
let mut cursor_col: u16 = 0;
|
||||
let mut cursor_set = false;
|
||||
let mut previous_was_cr = false;
|
||||
|
||||
// Record cursor once, at the point right before `atom` would be
|
||||
// placed — accounting for a wrap that the atom itself will cause.
|
||||
@@ -558,7 +657,7 @@ impl InputBuffer {
|
||||
for (i, atom) in self.atoms.iter().enumerate() {
|
||||
if !cursor_set && i == self.cursor {
|
||||
let leading = match atom {
|
||||
Atom::Char('\n') => 0,
|
||||
Atom::Char('\n' | '\r') => 0,
|
||||
Atom::Char(c) => UnicodeWidthChar::width(*c).unwrap_or(0),
|
||||
other => other
|
||||
.chip()
|
||||
@@ -573,6 +672,21 @@ impl InputBuffer {
|
||||
}
|
||||
|
||||
match atom {
|
||||
Atom::Char('\r') => {
|
||||
flush_pending(
|
||||
&mut pending,
|
||||
&mut pending_width,
|
||||
pending_style,
|
||||
&mut rows,
|
||||
&mut row_width,
|
||||
);
|
||||
rows.push(Vec::new());
|
||||
row_width = 0;
|
||||
previous_was_cr = true;
|
||||
}
|
||||
Atom::Char('\n') if previous_was_cr => {
|
||||
previous_was_cr = false;
|
||||
}
|
||||
Atom::Char('\n') => {
|
||||
flush_pending(
|
||||
&mut pending,
|
||||
@@ -583,8 +697,10 @@ impl InputBuffer {
|
||||
);
|
||||
rows.push(Vec::new());
|
||||
row_width = 0;
|
||||
previous_was_cr = false;
|
||||
}
|
||||
Atom::Char(c) => {
|
||||
previous_was_cr = false;
|
||||
let cw = UnicodeWidthChar::width(*c).unwrap_or(0);
|
||||
if pending_style != text_style && !pending.is_empty() {
|
||||
flush_pending(
|
||||
@@ -608,6 +724,7 @@ impl InputBuffer {
|
||||
);
|
||||
}
|
||||
other => {
|
||||
previous_was_cr = false;
|
||||
let (chip_style, label) = other.chip().expect("non-char atom has a chip");
|
||||
if pending_style != chip_style && !pending.is_empty() {
|
||||
flush_pending(
|
||||
@@ -848,6 +965,161 @@ mod render_viewport_tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod paste_policy_tests {
|
||||
use super::*;
|
||||
use protocol::Segment;
|
||||
use serde::Deserialize;
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct Fixture {
|
||||
max_plain_text_chars: usize,
|
||||
max_plain_text_logical_lines: usize,
|
||||
cases: Vec<FixtureCase>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct FixtureCase {
|
||||
name: String,
|
||||
parts: Vec<FixturePart>,
|
||||
char_count: usize,
|
||||
logical_line_count: usize,
|
||||
presentation: FixturePresentation,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct FixturePart {
|
||||
value: String,
|
||||
repeat: usize,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
enum FixturePresentation {
|
||||
Text,
|
||||
Chip,
|
||||
}
|
||||
|
||||
fn fixture() -> Fixture {
|
||||
serde_json::from_str(include_str!(
|
||||
"../../../tests/fixtures/composer-paste-policy.json"
|
||||
))
|
||||
.expect("shared composer paste policy fixture must be valid")
|
||||
}
|
||||
|
||||
fn fixture_content(case: &FixtureCase) -> String {
|
||||
case.parts
|
||||
.iter()
|
||||
.map(|part| part.value.repeat(part.repeat))
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tui_follows_shared_paste_presentation_contract() {
|
||||
let fixture = fixture();
|
||||
assert_eq!(fixture.max_plain_text_chars, MAX_PLAIN_TEXT_PASTE_CHARS);
|
||||
assert_eq!(
|
||||
fixture.max_plain_text_logical_lines,
|
||||
MAX_PLAIN_TEXT_PASTE_LOGICAL_LINES
|
||||
);
|
||||
|
||||
for case in fixture.cases {
|
||||
let content = fixture_content(&case);
|
||||
let measurement = measure_paste(&content);
|
||||
let expected_presentation = match case.presentation {
|
||||
FixturePresentation::Text => PastePresentation::Text,
|
||||
FixturePresentation::Chip => PastePresentation::Chip,
|
||||
};
|
||||
assert_eq!(measurement.chars, case.char_count, "{} chars", case.name);
|
||||
assert_eq!(
|
||||
measurement.logical_lines, case.logical_line_count,
|
||||
"{} logical lines",
|
||||
case.name
|
||||
);
|
||||
assert_eq!(
|
||||
measurement.presentation(),
|
||||
expected_presentation,
|
||||
"{} presentation",
|
||||
case.name
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn short_paste_is_editable_text_at_the_cursor() {
|
||||
let mut buffer = InputBuffer::new();
|
||||
buffer.insert_str("ac");
|
||||
buffer.move_left();
|
||||
buffer.insert_paste("b".to_owned());
|
||||
|
||||
assert_eq!(buffer.plain_text(), "abc");
|
||||
assert!(
|
||||
buffer
|
||||
.atoms
|
||||
.iter()
|
||||
.all(|atom| matches!(atom, Atom::Char(_)))
|
||||
);
|
||||
assert_eq!(
|
||||
buffer.submit_segments(),
|
||||
vec![Segment::text("abc".to_owned())]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn short_multiline_paste_preserves_original_line_endings_as_text() {
|
||||
let content = "ab\r\ncd\ref";
|
||||
let mut buffer = InputBuffer::new();
|
||||
buffer.insert_paste(content.to_owned());
|
||||
|
||||
assert_eq!(buffer.plain_text(), content);
|
||||
assert!(
|
||||
buffer
|
||||
.atoms
|
||||
.iter()
|
||||
.all(|atom| matches!(atom, Atom::Char(_)))
|
||||
);
|
||||
assert_eq!(
|
||||
buffer.submit_segments(),
|
||||
vec![Segment::text(content.to_owned())]
|
||||
);
|
||||
|
||||
let rendered: Vec<String> = buffer
|
||||
.render(80)
|
||||
.lines
|
||||
.iter()
|
||||
.map(|line| {
|
||||
line.spans
|
||||
.iter()
|
||||
.map(|span| span.content.as_ref())
|
||||
.collect()
|
||||
})
|
||||
.collect();
|
||||
assert_eq!(rendered, vec!["ab", "cd", "ef"]);
|
||||
|
||||
buffer.move_up();
|
||||
assert_eq!(buffer.cursor, 6);
|
||||
buffer.move_up();
|
||||
assert_eq!(buffer.cursor, 2);
|
||||
buffer.move_down();
|
||||
assert_eq!(buffer.cursor, 6);
|
||||
buffer.move_home();
|
||||
assert_eq!(buffer.cursor, 4);
|
||||
buffer.move_end();
|
||||
assert_eq!(buffer.cursor, 6);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn empty_paste_is_a_noop() {
|
||||
let mut buffer = InputBuffer::new();
|
||||
buffer.insert_str("unchanged");
|
||||
let paste_id = buffer.next_paste_id;
|
||||
buffer.insert_paste(String::new());
|
||||
|
||||
assert_eq!(buffer.plain_text(), "unchanged");
|
||||
assert_eq!(buffer.next_paste_id, paste_id);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod submit_segments_tests {
|
||||
use super::*;
|
||||
@@ -873,7 +1145,8 @@ mod submit_segments_tests {
|
||||
for c in "see ".chars() {
|
||||
buf.insert_char(c);
|
||||
}
|
||||
buf.insert_paste("line1\nline2".into());
|
||||
let pasted = "line1\nline2\nline3\nline4";
|
||||
buf.insert_paste(pasted.into());
|
||||
for c in " end".chars() {
|
||||
buf.insert_char(c);
|
||||
}
|
||||
@@ -890,9 +1163,9 @@ mod submit_segments_tests {
|
||||
content,
|
||||
..
|
||||
} => {
|
||||
assert_eq!(content, "line1\nline2");
|
||||
assert_eq!(*chars, "line1\nline2".chars().count() as u32);
|
||||
assert_eq!(*lines, 2);
|
||||
assert_eq!(content, pasted);
|
||||
assert_eq!(*chars, pasted.chars().count() as u32);
|
||||
assert_eq!(*lines, 4);
|
||||
}
|
||||
other => panic!("expected Paste, got {other:?}"),
|
||||
}
|
||||
@@ -902,6 +1175,45 @@ mod submit_segments_tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn restored_direct_paste_remains_a_typed_segment_without_reclassification() {
|
||||
let original = Segment::Paste {
|
||||
id: 7,
|
||||
chars: 1,
|
||||
lines: 1,
|
||||
content: "x".to_owned(),
|
||||
};
|
||||
let mut buf = InputBuffer::new();
|
||||
buf.replace_with_segments(std::slice::from_ref(&original));
|
||||
|
||||
assert_eq!(buf.submit_segments(), vec![original]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn restored_paste_artifact_remains_a_typed_segment() {
|
||||
let artifact = protocol::PasteArtifactRef {
|
||||
artifact_id: "019ca7c8-57b6-7f05-8edf-524147aba7b2".to_string(),
|
||||
created_at_ms: 1_700_000_000_000,
|
||||
media_type: protocol::PasteArtifactMediaType::TextPlainUtf8,
|
||||
availability: protocol::PasteArtifactAvailability::Available,
|
||||
byte_len: 65_536,
|
||||
char_count: 65_530,
|
||||
line_count: 200,
|
||||
sha256: "a".repeat(64),
|
||||
source_entry_id: "entry-1".to_string(),
|
||||
};
|
||||
let original = Segment::PasteArtifact {
|
||||
artifact: artifact.clone(),
|
||||
};
|
||||
let mut buf = InputBuffer::new();
|
||||
buf.replace_with_segments(std::slice::from_ref(&original));
|
||||
|
||||
assert_eq!(
|
||||
buf.submit_segments(),
|
||||
vec![Segment::PasteArtifact { artifact }]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn empty_buffer_yields_empty_segments() {
|
||||
let buf = InputBuffer::new();
|
||||
@@ -911,7 +1223,7 @@ mod submit_segments_tests {
|
||||
#[test]
|
||||
fn leading_paste_does_not_emit_empty_text() {
|
||||
let mut buf = InputBuffer::new();
|
||||
buf.insert_paste("X".into());
|
||||
buf.insert_paste("X".repeat(MAX_PLAIN_TEXT_PASTE_CHARS + 1));
|
||||
let segs = buf.submit_segments();
|
||||
assert_eq!(segs.len(), 1);
|
||||
assert!(matches!(segs[0], Segment::Paste { .. }));
|
||||
@@ -1011,7 +1323,7 @@ mod completion_prefix_tests {
|
||||
#[test]
|
||||
fn trigger_after_chip_atom() {
|
||||
let mut buf = InputBuffer::new();
|
||||
buf.insert_paste("X".into());
|
||||
buf.insert_paste("X".repeat(MAX_PLAIN_TEXT_PASTE_CHARS + 1));
|
||||
for c in "@sr".chars() {
|
||||
buf.insert_char(c);
|
||||
}
|
||||
@@ -1120,7 +1432,7 @@ mod word_motion_tests {
|
||||
for c in "foo ".chars() {
|
||||
buf.insert_char(c);
|
||||
}
|
||||
buf.insert_paste("anything".into());
|
||||
buf.insert_paste("anything".repeat(MAX_PLAIN_TEXT_PASTE_CHARS + 1));
|
||||
for c in " bar".chars() {
|
||||
buf.insert_char(c);
|
||||
}
|
||||
@@ -1219,7 +1531,9 @@ mod word_motion_tests {
|
||||
for a in &buf.atoms {
|
||||
match a {
|
||||
Atom::Char(c) => out.push(*c),
|
||||
Atom::Paste(_) | Atom::FileRef(_) | Atom::FlowRef(_) => out.push_str("<P>"),
|
||||
Atom::Paste(_) | Atom::PasteArtifact(_) | Atom::FileRef(_) | Atom::FlowRef(_) => {
|
||||
out.push_str("<P>")
|
||||
}
|
||||
}
|
||||
}
|
||||
out
|
||||
@@ -1277,7 +1591,7 @@ mod word_motion_tests {
|
||||
for c in "foo ".chars() {
|
||||
buf.insert_char(c);
|
||||
}
|
||||
buf.insert_paste("anything".into());
|
||||
buf.insert_paste("anything".repeat(MAX_PLAIN_TEXT_PASTE_CHARS + 1));
|
||||
for c in " bar".chars() {
|
||||
buf.insert_char(c);
|
||||
}
|
||||
|
||||
+5
-34
@@ -1,17 +1,17 @@
|
||||
use std::io::{self, Stdout, Write};
|
||||
use std::process::ExitCode;
|
||||
use std::time::Duration;
|
||||
|
||||
use crossterm::event::{self, Event, KeyCode, KeyEventKind, KeyModifiers};
|
||||
use crossterm::terminal::{disable_raw_mode, enable_raw_mode};
|
||||
use ratatui::backend::CrosstermBackend;
|
||||
use ratatui::Frame;
|
||||
use ratatui::layout::{Constraint, Layout};
|
||||
use ratatui::style::{Color, Modifier, Style};
|
||||
use ratatui::text::{Line, Span};
|
||||
use ratatui::widgets::Paragraph;
|
||||
use ratatui::{Frame, Terminal, TerminalOptions, Viewport};
|
||||
use secrets::{SecretStore, SecretValue};
|
||||
|
||||
use crate::inline_terminal::{InlineTerminal, with_inline_terminal};
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
enum Mode {
|
||||
Normal,
|
||||
@@ -235,7 +235,6 @@ pub async fn launch() -> ExitCode {
|
||||
}
|
||||
|
||||
type UiResult<T> = Result<T, Box<dyn std::error::Error>>;
|
||||
type InlineTerminal = Terminal<CrosstermBackend<Stdout>>;
|
||||
|
||||
const MAX_ROWS: usize = 10;
|
||||
const VIEWPORT_LINES: u16 = MAX_ROWS as u16 + 5;
|
||||
@@ -270,37 +269,9 @@ impl Drop for RawModeGuard {
|
||||
fn run(store: SecretStore) -> UiResult<()> {
|
||||
enable_raw_mode()?;
|
||||
let guard = RawModeGuard::new();
|
||||
let mut terminal = make_inline_terminal()?;
|
||||
let result = run_loop(&mut terminal, store);
|
||||
let close_result = close_viewport(&mut terminal);
|
||||
drop(terminal);
|
||||
let result = with_inline_terminal(VIEWPORT_LINES, |terminal| run_loop(terminal, store));
|
||||
guard.restore();
|
||||
result?;
|
||||
close_result?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn make_inline_terminal() -> io::Result<InlineTerminal> {
|
||||
let backend = CrosstermBackend::new(io::stdout());
|
||||
Terminal::with_options(
|
||||
backend,
|
||||
TerminalOptions {
|
||||
viewport: Viewport::Inline(VIEWPORT_LINES),
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
/// Park the cursor at the very bottom of the inline viewport and emit one
|
||||
/// newline before dropping the terminal. This matches the resume picker and
|
||||
/// keeps the shell prompt (or a later inline viewport) from drawing over rows.
|
||||
fn close_viewport(terminal: &mut InlineTerminal) -> io::Result<()> {
|
||||
let area = terminal.get_frame().area();
|
||||
let last_row = area.bottom().saturating_sub(1);
|
||||
terminal.set_cursor_position((0, last_row))?;
|
||||
let mut out = io::stdout();
|
||||
out.write_all(b"\r\n")?;
|
||||
out.flush()?;
|
||||
Ok(())
|
||||
result
|
||||
}
|
||||
|
||||
fn run_loop(terminal: &mut InlineTerminal, store: SecretStore) -> UiResult<()> {
|
||||
|
||||
+82
-100
@@ -8,24 +8,21 @@ mod command;
|
||||
mod composer_history;
|
||||
mod composer_keys;
|
||||
mod console;
|
||||
mod dashboard;
|
||||
#[cfg(feature = "e2e-test")]
|
||||
mod e2e_observer;
|
||||
mod inline_terminal;
|
||||
mod input;
|
||||
pub mod keys;
|
||||
mod markdown;
|
||||
mod picker;
|
||||
mod role_session_registry;
|
||||
mod scroll;
|
||||
pub mod setup_model;
|
||||
mod spawn;
|
||||
mod standalone_picker;
|
||||
mod standalone_spawn;
|
||||
mod task;
|
||||
mod text_selection;
|
||||
mod tool;
|
||||
mod ui;
|
||||
mod view_mode;
|
||||
mod worker_list;
|
||||
mod workspace_panel;
|
||||
|
||||
use std::io;
|
||||
use std::path::PathBuf;
|
||||
@@ -34,7 +31,6 @@ use std::process::ExitCode;
|
||||
use crossterm::event::{DisableBracketedPaste, DisableMouseCapture, EnableBracketedPaste};
|
||||
use crossterm::execute;
|
||||
use crossterm::terminal::{LeaveAlternateScreen, disable_raw_mode, enable_raw_mode};
|
||||
use session_store::SegmentId;
|
||||
|
||||
use client::{Target, WorkerConnectionSelector, WorkerListRequest};
|
||||
|
||||
@@ -47,42 +43,69 @@ pub struct LaunchOptions {
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum LaunchMode {
|
||||
/// Start one client-owned in-process Standalone Worker.
|
||||
Spawn {
|
||||
worker_name: Option<String>,
|
||||
profile: Option<String>,
|
||||
},
|
||||
/// `yoi --worker <name>`: attach to a live Worker by name if possible;
|
||||
/// otherwise launch the Worker runtime command with `--worker <name>` so it
|
||||
/// resumes from name-keyed state or creates a fresh same-name Worker.
|
||||
WorkerName {
|
||||
worker_name: String,
|
||||
socket_override: Option<PathBuf>,
|
||||
},
|
||||
/// `yoi workers` / `yoi --backend <url>`: list workers through the selected
|
||||
/// connection target, then attach to the selected Worker.
|
||||
/// Restore one client-owned standalone Worker. The current cwd is the default scope;
|
||||
/// `include_all` opts into all standalone Workers under the same client data root.
|
||||
StandaloneResume { include_all: bool },
|
||||
/// List Backend Workers and attach to the selected Worker.
|
||||
Workers {
|
||||
runtime_id: Option<String>,
|
||||
include_stopped: bool,
|
||||
all: bool,
|
||||
},
|
||||
/// `yoi --backend <url> --runtime-id <id> --worker-id <id>`: open one Worker
|
||||
/// through the selected connection target.
|
||||
/// Open one Backend Worker through the selected connection target.
|
||||
OpenWorker {
|
||||
runtime_id: String,
|
||||
worker_id: String,
|
||||
},
|
||||
/// `yoi resume`: open the Worker picker, then attach to the selected live Worker
|
||||
/// or restore the selected stopped Worker by name. Without `--all`, the picker
|
||||
/// is scoped to the current runtime workspace.
|
||||
Resume { all: bool },
|
||||
/// `yoi --session <UUID>`: skip the picker, go straight to the
|
||||
/// resume name dialog with `id` baked in.
|
||||
ResumeWithSession {
|
||||
id: SegmentId,
|
||||
worker_name: Option<String>,
|
||||
},
|
||||
/// `yoi panel`: open the workspace Dashboard from the current workspace.
|
||||
Panel { include_stopped: bool },
|
||||
/// Open the Backend Workspace dashboard.
|
||||
Panel,
|
||||
}
|
||||
|
||||
struct TerminalModeGuard {
|
||||
active: bool,
|
||||
}
|
||||
|
||||
impl TerminalModeGuard {
|
||||
fn new() -> Self {
|
||||
Self { active: true }
|
||||
}
|
||||
|
||||
fn restore(&mut self) -> io::Result<()> {
|
||||
if !self.active {
|
||||
return Ok(());
|
||||
}
|
||||
self.active = false;
|
||||
let mut stdout = io::stdout();
|
||||
execute!(
|
||||
stdout,
|
||||
DisableMouseCapture,
|
||||
LeaveAlternateScreen,
|
||||
DisableBracketedPaste,
|
||||
crossterm::cursor::Show
|
||||
)?;
|
||||
disable_raw_mode()
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for TerminalModeGuard {
|
||||
fn drop(&mut self) {
|
||||
if self.active {
|
||||
let mut stdout = io::stdout();
|
||||
let _ = execute!(
|
||||
stdout,
|
||||
DisableMouseCapture,
|
||||
LeaveAlternateScreen,
|
||||
DisableBracketedPaste,
|
||||
crossterm::cursor::Show
|
||||
);
|
||||
let _ = disable_raw_mode();
|
||||
self.active = false;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn launch(options: LaunchOptions) -> ExitCode {
|
||||
@@ -109,56 +132,46 @@ pub async fn launch(options: LaunchOptions) -> ExitCode {
|
||||
eprintln!("yoi: {e}");
|
||||
return ExitCode::FAILURE;
|
||||
}
|
||||
let mut terminal_mode = TerminalModeGuard::new();
|
||||
|
||||
let result = match mode {
|
||||
LaunchMode::Spawn {
|
||||
worker_name,
|
||||
profile,
|
||||
} => match target.spawn_worker() {
|
||||
Ok(spawn) => {
|
||||
console::run_spawn(None, worker_name, profile, spawn.runtime_command).await
|
||||
}
|
||||
Err(e) => Err(Box::new(e) as Box<dyn std::error::Error>),
|
||||
} => match standalone_spawn::select(&workspace_root, worker_name, profile) {
|
||||
Ok(Some(selection)) => match target.spawn_worker() {
|
||||
Ok(spawn) => {
|
||||
console::run_standalone(
|
||||
workspace_root.clone(),
|
||||
spawn.state_dir,
|
||||
Some(selection.worker_name),
|
||||
Some(selection.profile),
|
||||
)
|
||||
.await
|
||||
}
|
||||
Err(error) => Err(Box::new(error) as Box<dyn std::error::Error>),
|
||||
},
|
||||
Ok(None) => Ok(()),
|
||||
Err(error) => Err(Box::new(error) as Box<dyn std::error::Error>),
|
||||
},
|
||||
LaunchMode::WorkerName {
|
||||
worker_name,
|
||||
socket_override,
|
||||
} => match target.worker_by_name() {
|
||||
Ok(worker_by_name) => {
|
||||
console::run_worker_name(
|
||||
worker_name,
|
||||
socket_override,
|
||||
worker_by_name.runtime_command,
|
||||
)
|
||||
.await
|
||||
LaunchMode::StandaloneResume { include_all } => {
|
||||
match standalone_picker::pick(target.as_ref(), include_all) {
|
||||
Ok(Some(intent)) => console::run_standalone_restore(intent).await,
|
||||
Ok(None) => Ok(()),
|
||||
Err(error) => Err(Box::new(error) as Box<dyn std::error::Error>),
|
||||
}
|
||||
Err(e) => Err(Box::new(e) as Box<dyn std::error::Error>),
|
||||
},
|
||||
}
|
||||
LaunchMode::Workers {
|
||||
runtime_id,
|
||||
include_stopped,
|
||||
all,
|
||||
} => match target.list_workers(if include_stopped {
|
||||
WorkerListRequest::with_stopped(runtime_id)
|
||||
} else {
|
||||
WorkerListRequest::new(runtime_id)
|
||||
}) {
|
||||
Ok(worker_list) => {
|
||||
if let Some(target) = worker_list.backend_target {
|
||||
backend_worker_picker::run(target, worker_list.include_stopped).await
|
||||
} else if let Some(runtime_command) = worker_list.local_runtime_command {
|
||||
console::run_worker_picker(
|
||||
runtime_command,
|
||||
workspace_root.clone(),
|
||||
all,
|
||||
worker_list.include_stopped,
|
||||
)
|
||||
backend_worker_picker::run(worker_list.backend_target, worker_list.include_stopped)
|
||||
.await
|
||||
} else {
|
||||
Err(Box::new(io::Error::other(
|
||||
"worker list target did not include a local or backend source",
|
||||
)) as Box<dyn std::error::Error>)
|
||||
}
|
||||
}
|
||||
Err(e) => Err(Box::new(e) as Box<dyn std::error::Error>),
|
||||
},
|
||||
@@ -169,28 +182,12 @@ pub async fn launch(options: LaunchOptions) -> ExitCode {
|
||||
Ok(connection) => console::run_backend_runtime(connection.target).await,
|
||||
Err(e) => Err(Box::new(e) as Box<dyn std::error::Error>),
|
||||
},
|
||||
LaunchMode::Resume { all } => match target.resume_worker() {
|
||||
Ok(resume) => {
|
||||
console::run_resume(resume.runtime_command, workspace_root.clone(), all).await
|
||||
LaunchMode::Panel => match target.dashboard() {
|
||||
Ok(dashboard) => {
|
||||
backend_dashboard::launch(dashboard.base_url, dashboard.workspace_id).await
|
||||
}
|
||||
Err(e) => Err(Box::new(e) as Box<dyn std::error::Error>),
|
||||
},
|
||||
LaunchMode::ResumeWithSession { id, worker_name } => match target.spawn_worker() {
|
||||
Ok(spawn) => {
|
||||
console::run_spawn(Some(id), worker_name, None, spawn.runtime_command).await
|
||||
}
|
||||
Err(e) => Err(Box::new(e) as Box<dyn std::error::Error>),
|
||||
},
|
||||
LaunchMode::Panel { include_stopped } => match target.dashboard() {
|
||||
Ok(client::Dashboard::Local { runtime_command }) => {
|
||||
dashboard::launch(runtime_command, include_stopped).await
|
||||
}
|
||||
Ok(client::Dashboard::Backend {
|
||||
base_url,
|
||||
workspace_id,
|
||||
}) => backend_dashboard::launch(base_url, workspace_id).await,
|
||||
Err(e) => Err(Box::new(e) as Box<dyn std::error::Error>),
|
||||
},
|
||||
};
|
||||
|
||||
// Always restore the terminal first so any pending eprintln below
|
||||
@@ -198,15 +195,7 @@ pub async fn launch(options: LaunchOptions) -> ExitCode {
|
||||
// alternate-screen buffer.
|
||||
#[cfg(feature = "e2e-test")]
|
||||
e2e_observer::emit("tui", "terminal_cleanup_started", serde_json::json!({}));
|
||||
let mut stdout = io::stdout();
|
||||
let _ = execute!(
|
||||
stdout,
|
||||
DisableMouseCapture,
|
||||
LeaveAlternateScreen,
|
||||
DisableBracketedPaste
|
||||
);
|
||||
let _ = disable_raw_mode();
|
||||
let _ = execute!(stdout, crossterm::cursor::Show);
|
||||
let _ = terminal_mode.restore();
|
||||
#[cfg(feature = "e2e-test")]
|
||||
e2e_observer::emit("tui", "terminal_cleanup_finished", serde_json::json!({}));
|
||||
|
||||
@@ -217,14 +206,7 @@ pub async fn launch(options: LaunchOptions) -> ExitCode {
|
||||
ExitCode::SUCCESS
|
||||
}
|
||||
Err(e) => {
|
||||
// SpawnError has already been painted into the inline
|
||||
// viewport's final frame, so it's already visible in the
|
||||
// user's scrollback — printing it again would be a noisy
|
||||
// duplicate. Other errors (worker-name failures, terminal setup
|
||||
// hiccups, etc.) need surfacing here.
|
||||
if e.downcast_ref::<spawn::SpawnError>().is_none() {
|
||||
eprintln!("yoi: {e}");
|
||||
}
|
||||
eprintln!("yoi: {e}");
|
||||
#[cfg(feature = "e2e-test")]
|
||||
e2e_observer::emit("tui", "exit", serde_json::json!({ "status": "failure" }));
|
||||
ExitCode::FAILURE
|
||||
|
||||
@@ -1,525 +0,0 @@
|
||||
//! Inline-viewport "pick a Worker to attach or restore" UX.
|
||||
//!
|
||||
//! Reads live Worker allocations from the runtime registry and stopped Worker state
|
||||
//! from the session-store worker metadata name-keyed metadata. Picking a live row attaches to
|
||||
//! its socket; picking a stopped row restores via the Worker runtime command.
|
||||
|
||||
use std::io;
|
||||
use std::path::PathBuf;
|
||||
use std::time::Duration;
|
||||
|
||||
use crossterm::event::{self, Event as TermEvent, KeyCode, KeyEventKind, KeyModifiers};
|
||||
use ratatui::Terminal;
|
||||
use ratatui::backend::CrosstermBackend;
|
||||
use ratatui::layout::{Constraint, Layout};
|
||||
use ratatui::style::{Color, Modifier, Style};
|
||||
use ratatui::text::{Line, Span};
|
||||
use ratatui::widgets::Paragraph;
|
||||
use ratatui::{Frame, TerminalOptions, Viewport};
|
||||
use session_store::FsStore;
|
||||
use session_store::FsWorkerStore;
|
||||
|
||||
use crate::worker_list::{
|
||||
LiveWorkerInfo, StoredMetadataState, StoredWorkerInfo, WorkerList, WorkerListEntry,
|
||||
WorkerVisibilitySource, live_socket_for_worker as worker_list_live_socket_for_worker,
|
||||
read_reachable_live_worker_infos, read_stored_worker_infos,
|
||||
};
|
||||
|
||||
const MAX_ROWS: usize = 10;
|
||||
const VIEWPORT_LINES: u16 = MAX_ROWS as u16 + 4;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum PickerError {
|
||||
Io(io::Error),
|
||||
Store(session_store::StoreError),
|
||||
NoWorkers { all: bool },
|
||||
}
|
||||
|
||||
impl std::fmt::Display for PickerError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::Io(e) => write!(f, "io error: {e}"),
|
||||
Self::Store(e) => write!(f, "session store error: {e}"),
|
||||
Self::NoWorkers { all: true } => write!(
|
||||
f,
|
||||
"no workers found — start a fresh Worker with `yoi` and try again"
|
||||
),
|
||||
Self::NoWorkers { all: false } => write!(
|
||||
f,
|
||||
"no workers found in this workspace — use `yoi resume --all` to list all host/data-dir Workers"
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for PickerError {}
|
||||
|
||||
impl From<io::Error> for PickerError {
|
||||
fn from(e: io::Error) -> Self {
|
||||
Self::Io(e)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<session_store::StoreError> for PickerError {
|
||||
fn from(e: session_store::StoreError) -> Self {
|
||||
Self::Store(e)
|
||||
}
|
||||
}
|
||||
|
||||
pub enum PickerOutcome {
|
||||
/// User picked a Worker. `socket_override` is set for live rows when the
|
||||
/// runtime registry knows the exact socket path; stopped rows leave it
|
||||
/// empty so the caller restores by spawning the Worker runtime command.
|
||||
Picked {
|
||||
worker_name: String,
|
||||
socket_override: Option<PathBuf>,
|
||||
},
|
||||
Cancelled,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) struct PickerOptions {
|
||||
scope: PickerScope,
|
||||
include_stopped: bool,
|
||||
}
|
||||
|
||||
impl PickerOptions {
|
||||
pub(crate) fn workspace(workspace_root: PathBuf) -> Self {
|
||||
Self {
|
||||
scope: PickerScope::Workspace(workspace_root),
|
||||
include_stopped: true,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn all() -> Self {
|
||||
Self {
|
||||
scope: PickerScope::All,
|
||||
include_stopped: true,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn with_stopped(mut self, include_stopped: bool) -> Self {
|
||||
self.include_stopped = include_stopped;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
enum PickerScope {
|
||||
Workspace(PathBuf),
|
||||
All,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum WorkerRowState {
|
||||
Live,
|
||||
Stopped,
|
||||
Corrupt,
|
||||
}
|
||||
|
||||
impl WorkerRowState {
|
||||
fn label(self) -> &'static str {
|
||||
match self {
|
||||
Self::Live => "live",
|
||||
Self::Stopped => "stopped",
|
||||
Self::Corrupt => "corrupt",
|
||||
}
|
||||
}
|
||||
|
||||
fn style(self) -> Style {
|
||||
match self {
|
||||
Self::Live => Style::default()
|
||||
.fg(Color::Green)
|
||||
.add_modifier(Modifier::BOLD),
|
||||
Self::Stopped => Style::default().fg(Color::Yellow),
|
||||
Self::Corrupt => Style::default().fg(Color::Red).add_modifier(Modifier::BOLD),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn list_for_options(
|
||||
options: &PickerOptions,
|
||||
stored_workers: Vec<StoredWorkerInfo>,
|
||||
live_workers: Vec<LiveWorkerInfo>,
|
||||
) -> WorkerList {
|
||||
let stored_workers = if options.include_stopped {
|
||||
stored_workers
|
||||
} else {
|
||||
Vec::new()
|
||||
};
|
||||
match &options.scope {
|
||||
PickerScope::Workspace(workspace_root) => WorkerList::from_workspace_sources(
|
||||
WorkerVisibilitySource::ResumePicker,
|
||||
stored_workers,
|
||||
live_workers,
|
||||
None,
|
||||
MAX_ROWS,
|
||||
workspace_root,
|
||||
),
|
||||
PickerScope::All => WorkerList::from_sources(
|
||||
WorkerVisibilitySource::ResumePicker,
|
||||
stored_workers,
|
||||
live_workers,
|
||||
None,
|
||||
MAX_ROWS,
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn run(options: PickerOptions) -> Result<PickerOutcome, PickerError> {
|
||||
let store_dir = default_store_dir()?;
|
||||
let store = FsStore::new(&store_dir)?;
|
||||
let worker_metadata_store =
|
||||
FsWorkerStore::new(default_worker_metadata_dir()?).map_err(io::Error::other)?;
|
||||
let stored_workers = read_stored_worker_infos(&store, &worker_metadata_store)?;
|
||||
let live_workers = read_reachable_live_worker_infos(&store)
|
||||
.await
|
||||
.unwrap_or_default();
|
||||
let mut list = list_for_options(&options, stored_workers, live_workers);
|
||||
if list.entries.is_empty() {
|
||||
return Err(PickerError::NoWorkers {
|
||||
all: matches!(options.scope, PickerScope::All),
|
||||
});
|
||||
}
|
||||
|
||||
let mut terminal = make_inline_terminal()?;
|
||||
loop {
|
||||
terminal.draw(|f| draw(f, &list))?;
|
||||
match poll_event()? {
|
||||
None => continue,
|
||||
Some(Action::Up) => {
|
||||
let selected = list.selected_index().saturating_sub(1);
|
||||
list.select_index(selected);
|
||||
}
|
||||
Some(Action::Down) => {
|
||||
let selected = list.selected_index();
|
||||
if selected + 1 < list.entries.len() {
|
||||
list.select_index(selected + 1);
|
||||
}
|
||||
}
|
||||
Some(Action::Submit) => {
|
||||
close_viewport(&mut terminal)?;
|
||||
let entry = list.selected_entry().expect("non-empty worker list");
|
||||
return Ok(PickerOutcome::Picked {
|
||||
worker_name: entry.name.clone(),
|
||||
socket_override: entry.attach_socket_path().map(PathBuf::from),
|
||||
});
|
||||
}
|
||||
Some(Action::Cancel) => {
|
||||
close_viewport(&mut terminal)?;
|
||||
return Ok(PickerOutcome::Cancelled);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Park the cursor at the very bottom of the picker's inline viewport and emit
|
||||
/// one newline before dropping the terminal. This keeps any next inline viewport
|
||||
/// from drawing over the lower picker rows.
|
||||
fn close_viewport(terminal: &mut Terminal<CrosstermBackend<io::Stdout>>) -> io::Result<()> {
|
||||
let area = terminal.get_frame().area();
|
||||
let last_row = area.bottom().saturating_sub(1);
|
||||
terminal.set_cursor_position((0, last_row))?;
|
||||
use std::io::Write;
|
||||
let mut out = io::stdout();
|
||||
out.write_all(b"\r\n")?;
|
||||
out.flush()?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn default_store_dir() -> Result<PathBuf, PickerError> {
|
||||
manifest::paths::sessions_dir().ok_or_else(|| {
|
||||
PickerError::Io(io::Error::new(
|
||||
io::ErrorKind::NotFound,
|
||||
"could not resolve sessions directory \
|
||||
(set YOI_DATA_DIR, YOI_HOME, XDG_DATA_HOME, or HOME)",
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
fn default_worker_metadata_dir() -> Result<PathBuf, PickerError> {
|
||||
manifest::paths::data_dir()
|
||||
.map(|dir| dir.join("workers"))
|
||||
.ok_or_else(|| {
|
||||
PickerError::Io(io::Error::new(
|
||||
io::ErrorKind::NotFound,
|
||||
"could not resolve worker state directory \
|
||||
(set YOI_DATA_DIR, YOI_HOME, XDG_DATA_HOME, or HOME)",
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn live_socket_for_worker(worker_name: &str) -> Option<PathBuf> {
|
||||
worker_list_live_socket_for_worker(worker_name)
|
||||
}
|
||||
|
||||
fn make_inline_terminal() -> io::Result<Terminal<CrosstermBackend<io::Stdout>>> {
|
||||
let backend = CrosstermBackend::new(io::stdout());
|
||||
Terminal::with_options(
|
||||
backend,
|
||||
TerminalOptions {
|
||||
viewport: Viewport::Inline(VIEWPORT_LINES),
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
enum Action {
|
||||
Up,
|
||||
Down,
|
||||
Submit,
|
||||
Cancel,
|
||||
}
|
||||
|
||||
fn poll_event() -> io::Result<Option<Action>> {
|
||||
if !event::poll(Duration::from_millis(100))? {
|
||||
return Ok(None);
|
||||
}
|
||||
match event::read()? {
|
||||
TermEvent::Key(k) if k.kind != KeyEventKind::Release => {
|
||||
let ctrl = k.modifiers.contains(KeyModifiers::CONTROL);
|
||||
Ok(match k.code {
|
||||
KeyCode::Up => Some(Action::Up),
|
||||
KeyCode::Down => Some(Action::Down),
|
||||
KeyCode::Char('k') if !ctrl => Some(Action::Up),
|
||||
KeyCode::Char('j') if !ctrl => Some(Action::Down),
|
||||
KeyCode::Enter => Some(Action::Submit),
|
||||
KeyCode::Esc => Some(Action::Cancel),
|
||||
KeyCode::Char('c') if ctrl => Some(Action::Cancel),
|
||||
_ => None,
|
||||
})
|
||||
}
|
||||
_ => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
fn draw(f: &mut Frame<'_>, list: &WorkerList) {
|
||||
let area = f.area();
|
||||
let mut constraints: Vec<Constraint> = Vec::with_capacity(list.entries.len() + 3);
|
||||
constraints.push(Constraint::Length(1)); // title
|
||||
for _ in &list.entries {
|
||||
constraints.push(Constraint::Length(1));
|
||||
}
|
||||
constraints.push(Constraint::Length(1)); // hint
|
||||
constraints.push(Constraint::Length(1)); // spacer
|
||||
let layout = Layout::vertical(constraints).split(area);
|
||||
|
||||
f.render_widget(
|
||||
Paragraph::new(Line::from(vec![Span::styled(
|
||||
picker_title(),
|
||||
Style::default().add_modifier(Modifier::BOLD),
|
||||
)])),
|
||||
layout[0],
|
||||
);
|
||||
|
||||
let selected = list.selected_index();
|
||||
for (i, entry) in list.entries.iter().enumerate() {
|
||||
f.render_widget(
|
||||
Paragraph::new(row_line(entry, i == selected)),
|
||||
layout[i + 1],
|
||||
);
|
||||
}
|
||||
|
||||
f.render_widget(
|
||||
Paragraph::new(Line::from(vec![
|
||||
Span::raw(" "),
|
||||
Span::styled("[↑/↓]", Style::default().fg(Color::DarkGray)),
|
||||
Span::raw(" select "),
|
||||
Span::styled("[enter]", Style::default().fg(Color::Green)),
|
||||
Span::raw(" open/restore "),
|
||||
Span::styled("[esc]", Style::default().fg(Color::Yellow)),
|
||||
Span::raw(" cancel"),
|
||||
])),
|
||||
layout[list.entries.len() + 1],
|
||||
);
|
||||
}
|
||||
|
||||
fn picker_title() -> &'static str {
|
||||
"resume worker pick a worker"
|
||||
}
|
||||
|
||||
fn row_line(entry: &WorkerListEntry, selected: bool) -> Line<'_> {
|
||||
let marker = if selected { "▶ " } else { " " };
|
||||
let name_style = if selected {
|
||||
Style::default()
|
||||
.fg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD)
|
||||
} else {
|
||||
Style::default().fg(Color::Cyan)
|
||||
};
|
||||
let preview_style = if selected {
|
||||
Style::default().fg(Color::White)
|
||||
} else {
|
||||
Style::default().fg(Color::DarkGray)
|
||||
};
|
||||
let state = row_state(entry);
|
||||
let _visibility = entry.visibility;
|
||||
let _source_kinds = &entry.source_kinds;
|
||||
|
||||
let mut spans = vec![
|
||||
Span::raw(marker),
|
||||
Span::styled(entry.name.as_str(), name_style),
|
||||
Span::raw(" "),
|
||||
Span::styled(format!("[{}]", state.label()), state.style()),
|
||||
Span::raw(" "),
|
||||
Span::styled(
|
||||
format_updated_at(entry.summary.updated_at),
|
||||
Style::default().fg(Color::DarkGray),
|
||||
),
|
||||
Span::raw(" "),
|
||||
Span::styled(debug_ids(entry), Style::default().fg(Color::DarkGray)),
|
||||
];
|
||||
if let Some(preview) = entry.summary.preview.as_ref() {
|
||||
spans.push(Span::raw(" "));
|
||||
spans.push(Span::styled(preview.as_str(), preview_style));
|
||||
}
|
||||
Line::from(spans)
|
||||
}
|
||||
|
||||
fn row_state(entry: &WorkerListEntry) -> WorkerRowState {
|
||||
if entry.live.as_ref().is_some_and(|live| live.reachable) {
|
||||
return WorkerRowState::Live;
|
||||
}
|
||||
if entry
|
||||
.stored
|
||||
.as_ref()
|
||||
.is_some_and(|stored| matches!(stored.metadata_state, StoredMetadataState::Corrupt(_)))
|
||||
{
|
||||
return WorkerRowState::Corrupt;
|
||||
}
|
||||
WorkerRowState::Stopped
|
||||
}
|
||||
|
||||
fn format_updated_at(updated_at: u64) -> String {
|
||||
if updated_at == 0 {
|
||||
"updated: —".to_string()
|
||||
} else {
|
||||
format!("updated: {updated_at}")
|
||||
}
|
||||
}
|
||||
|
||||
fn debug_ids(entry: &WorkerListEntry) -> String {
|
||||
let session = entry
|
||||
.summary
|
||||
.active_session_id
|
||||
.map(short_id)
|
||||
.unwrap_or_else(|| "--------".to_string());
|
||||
let segment = entry
|
||||
.summary
|
||||
.active_segment_id
|
||||
.map(short_id)
|
||||
.unwrap_or_else(|| "--------".to_string());
|
||||
format!("s:{session} g:{segment}")
|
||||
}
|
||||
|
||||
fn short_id<T: ToString>(id: T) -> String {
|
||||
id.to_string().chars().take(8).collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn picker_title_names_pods_not_sessions() {
|
||||
assert_eq!(picker_title(), "resume worker pick a worker");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn picker_no_pods_message_mentions_all_for_workspace_scope() {
|
||||
let message = PickerError::NoWorkers { all: false }.to_string();
|
||||
assert!(message.contains("no workers found in this workspace"));
|
||||
assert!(message.contains("yoi resume --all"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn picker_no_pods_message_keeps_fresh_pod_hint_for_all_scope() {
|
||||
let message = PickerError::NoWorkers { all: true }.to_string();
|
||||
assert!(message.contains("start a fresh Worker with `yoi`"));
|
||||
assert!(!message.contains("yoi resume --all"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn picker_workspace_options_filter_by_workspace_metadata() {
|
||||
let list = list_for_options(
|
||||
&PickerOptions::workspace(PathBuf::from("/workspace/current")),
|
||||
vec![
|
||||
stored_pod("current", Some("/workspace/current"), 3),
|
||||
stored_pod("other", Some("/workspace/other"), 2),
|
||||
stored_pod("legacy", None, 1),
|
||||
],
|
||||
vec![],
|
||||
);
|
||||
|
||||
let names: Vec<_> = list
|
||||
.entries
|
||||
.iter()
|
||||
.map(|entry| entry.name.as_str())
|
||||
.collect();
|
||||
assert_eq!(names, vec!["current"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn picker_all_options_include_host_wide_and_legacy_pods() {
|
||||
let list = list_for_options(
|
||||
&PickerOptions::all(),
|
||||
vec![
|
||||
stored_pod("current", Some("/workspace/current"), 3),
|
||||
stored_pod("other", Some("/workspace/other"), 2),
|
||||
stored_pod("legacy", None, 1),
|
||||
],
|
||||
vec![],
|
||||
);
|
||||
|
||||
let names: Vec<_> = list
|
||||
.entries
|
||||
.iter()
|
||||
.map(|entry| entry.name.as_str())
|
||||
.collect();
|
||||
assert_eq!(names, vec!["current", "other", "legacy"]);
|
||||
}
|
||||
|
||||
fn stored_pod(name: &str, workspace_root: Option<&str>, updated_at: u64) -> StoredWorkerInfo {
|
||||
StoredWorkerInfo {
|
||||
worker_name: name.to_string(),
|
||||
metadata_state: StoredMetadataState::Present,
|
||||
active_session_id: None,
|
||||
active_segment_id: None,
|
||||
updated_at,
|
||||
workspace_root: workspace_root.map(PathBuf::from),
|
||||
preview: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn picker_row_shows_live_pending_preview_and_runtime_segment_id() {
|
||||
let segment_id = session_store::new_segment_id();
|
||||
let entry = WorkerList::from_sources(
|
||||
WorkerVisibilitySource::ResumePicker,
|
||||
vec![],
|
||||
vec![crate::worker_list::LiveWorkerInfo {
|
||||
worker_name: "pending".to_string(),
|
||||
socket_path: PathBuf::from("/tmp/pending.sock"),
|
||||
status: Some(protocol::WorkerStatus::Idle),
|
||||
reachable: true,
|
||||
segment_id: Some(segment_id),
|
||||
summary: crate::worker_list::WorkerEntrySummary::default(),
|
||||
}],
|
||||
None,
|
||||
10,
|
||||
)
|
||||
.entries
|
||||
.into_iter()
|
||||
.next()
|
||||
.unwrap();
|
||||
|
||||
let text = row_line(&entry, false)
|
||||
.spans
|
||||
.iter()
|
||||
.map(|span| span.content.as_ref())
|
||||
.collect::<String>();
|
||||
|
||||
assert!(text.contains("[live]"));
|
||||
assert!(text.contains("[live, pending segment]"));
|
||||
assert!(text.contains(&format!("g:{}", short_id(segment_id))));
|
||||
}
|
||||
}
|
||||
@@ -1,556 +0,0 @@
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::fs::{self, OpenOptions};
|
||||
use std::io;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::thread;
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
const REGISTRY_VERSION: u32 = 1;
|
||||
const REGISTRY_FILE: &str = "role-sessions.json";
|
||||
const REGISTRY_LOCK_FILE: &str = "role-sessions.lock";
|
||||
const CLAIMS_DIR: &str = "ticket-claims";
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) struct PanelRegistryStore {
|
||||
root: PathBuf,
|
||||
workspace_root: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub(crate) struct RoleSessionRegistry {
|
||||
pub version: u32,
|
||||
pub workspace_root: String,
|
||||
pub sessions: BTreeMap<String, RoleSessionRecord>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub(crate) struct RoleSessionRecord {
|
||||
pub role: String,
|
||||
pub worker_name: String,
|
||||
pub origin: RoleSessionOrigin,
|
||||
pub created_at: String,
|
||||
pub updated_at: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub session_id: Option<String>,
|
||||
#[serde(default)]
|
||||
pub related_tickets: Vec<RelatedTicketRef>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub(crate) enum RoleSessionOrigin {
|
||||
PreTicketIntake,
|
||||
TicketClaim,
|
||||
RoleLaunch,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, PartialOrd, Ord)]
|
||||
pub(crate) struct RelatedTicketRef {
|
||||
pub id: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub slug: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub(crate) struct TicketClaim {
|
||||
pub ticket_id: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub ticket_slug: Option<String>,
|
||||
pub worker_name: String,
|
||||
pub role: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(crate) struct PanelRegistrySnapshot {
|
||||
pub sessions: Vec<RoleSessionRecord>,
|
||||
pub claims: Vec<TicketClaim>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(crate) enum TicketClaimResult {
|
||||
Claimed,
|
||||
AlreadyOwned(TicketClaim),
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub(crate) enum PanelRegistryError {
|
||||
Io(io::Error),
|
||||
Json(serde_json::Error),
|
||||
TicketAlreadyClaimed(TicketClaim),
|
||||
}
|
||||
|
||||
impl std::fmt::Display for PanelRegistryError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::Io(error) => write!(f, "local role session registry I/O error: {error}"),
|
||||
Self::Json(error) => write!(f, "local role session registry JSON error: {error}"),
|
||||
Self::TicketAlreadyClaimed(claim) => write!(
|
||||
f,
|
||||
"Ticket {} is already claimed locally by {} ({})",
|
||||
claim.ticket_id, claim.worker_name, claim.role
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for PanelRegistryError {}
|
||||
|
||||
impl From<io::Error> for PanelRegistryError {
|
||||
fn from(error: io::Error) -> Self {
|
||||
Self::Io(error)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<serde_json::Error> for PanelRegistryError {
|
||||
fn from(error: serde_json::Error) -> Self {
|
||||
Self::Json(error)
|
||||
}
|
||||
}
|
||||
|
||||
impl PanelRegistryStore {
|
||||
pub(crate) fn default_for_workspace(workspace_root: &Path) -> Result<Self, PanelRegistryError> {
|
||||
let data_dir = manifest::paths::data_dir().ok_or_else(|| {
|
||||
PanelRegistryError::Io(io::Error::other("failed to resolve yoi data directory"))
|
||||
})?;
|
||||
Ok(Self::for_data_dir(data_dir, workspace_root))
|
||||
}
|
||||
|
||||
pub(crate) fn for_data_dir(data_dir: impl AsRef<Path>, workspace_root: &Path) -> Self {
|
||||
let workspace_root = normalized_workspace_key(workspace_root);
|
||||
let leaf = workspace_leaf(&workspace_root);
|
||||
let digest = fnv1a64_hex(workspace_root.as_bytes());
|
||||
Self {
|
||||
root: data_dir
|
||||
.as_ref()
|
||||
.join("panel")
|
||||
.join("workspaces")
|
||||
.join(format!("{leaf}-{digest}")),
|
||||
workspace_root: Some(workspace_root),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn from_root(root: impl Into<PathBuf>) -> Self {
|
||||
Self {
|
||||
root: root.into(),
|
||||
workspace_root: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn root(&self) -> &Path {
|
||||
&self.root
|
||||
}
|
||||
|
||||
pub(crate) fn snapshot(&self) -> Result<PanelRegistrySnapshot, PanelRegistryError> {
|
||||
let registry = self.load_registry()?;
|
||||
let claims = self.load_claims()?;
|
||||
Ok(PanelRegistrySnapshot {
|
||||
sessions: registry.sessions.into_values().collect(),
|
||||
claims,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn load_registry(&self) -> Result<RoleSessionRegistry, PanelRegistryError> {
|
||||
match fs::read(self.registry_path()) {
|
||||
Ok(bytes) => Ok(serde_json::from_slice(&bytes)?),
|
||||
Err(error) if error.kind() == io::ErrorKind::NotFound => Ok(RoleSessionRegistry {
|
||||
version: REGISTRY_VERSION,
|
||||
workspace_root: self.workspace_root.clone().unwrap_or_default(),
|
||||
sessions: BTreeMap::new(),
|
||||
}),
|
||||
Err(error) => Err(error.into()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn record_session(
|
||||
&self,
|
||||
worker_name: impl Into<String>,
|
||||
role: impl Into<String>,
|
||||
origin: RoleSessionOrigin,
|
||||
session_id: Option<String>,
|
||||
related_tickets: impl IntoIterator<Item = RelatedTicketRef>,
|
||||
) -> Result<(), PanelRegistryError> {
|
||||
let worker_name = worker_name.into();
|
||||
let role = role.into();
|
||||
let related_tickets: Vec<RelatedTicketRef> = related_tickets.into_iter().collect();
|
||||
self.update_registry(|registry| {
|
||||
let now = now_timestamp_string();
|
||||
let mut tickets: BTreeSet<RelatedTicketRef> = registry
|
||||
.sessions
|
||||
.get(&worker_name)
|
||||
.map(|record| record.related_tickets.iter().cloned().collect())
|
||||
.unwrap_or_default();
|
||||
tickets.extend(related_tickets);
|
||||
let created_at = registry
|
||||
.sessions
|
||||
.get(&worker_name)
|
||||
.map(|record| record.created_at.clone())
|
||||
.unwrap_or_else(|| now.clone());
|
||||
registry.sessions.insert(
|
||||
worker_name.clone(),
|
||||
RoleSessionRecord {
|
||||
role,
|
||||
worker_name,
|
||||
origin,
|
||||
created_at,
|
||||
updated_at: now,
|
||||
session_id,
|
||||
related_tickets: tickets.into_iter().collect(),
|
||||
},
|
||||
);
|
||||
Ok(())
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn claim_ticket(
|
||||
&self,
|
||||
ticket_id: &str,
|
||||
ticket_slug: Option<&str>,
|
||||
worker_name: &str,
|
||||
role: &str,
|
||||
) -> Result<TicketClaimResult, PanelRegistryError> {
|
||||
fs::create_dir_all(self.claims_dir())?;
|
||||
let claim_path = self.claim_path(ticket_id);
|
||||
let claim = TicketClaim {
|
||||
ticket_id: ticket_id.to_string(),
|
||||
ticket_slug: ticket_slug.map(ToOwned::to_owned),
|
||||
worker_name: worker_name.to_string(),
|
||||
role: role.to_string(),
|
||||
};
|
||||
match self.create_claim_file(&claim_path, &claim) {
|
||||
Ok(()) => {
|
||||
if let Err(error) = self.record_session(
|
||||
worker_name.to_string(),
|
||||
role.to_string(),
|
||||
RoleSessionOrigin::TicketClaim,
|
||||
None,
|
||||
[RelatedTicketRef {
|
||||
id: ticket_id.to_string(),
|
||||
slug: ticket_slug.map(ToOwned::to_owned),
|
||||
}],
|
||||
) {
|
||||
let _ = fs::remove_file(&claim_path);
|
||||
return Err(error);
|
||||
}
|
||||
Ok(TicketClaimResult::Claimed)
|
||||
}
|
||||
Err(error) if error.kind() == io::ErrorKind::AlreadyExists => {
|
||||
let existing = self.load_claim(ticket_id)?;
|
||||
if existing.worker_name == worker_name && existing.role == role {
|
||||
Ok(TicketClaimResult::AlreadyOwned(existing))
|
||||
} else {
|
||||
Err(PanelRegistryError::TicketAlreadyClaimed(existing))
|
||||
}
|
||||
}
|
||||
Err(error) => Err(error.into()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn load_claim(&self, ticket_id: &str) -> Result<TicketClaim, PanelRegistryError> {
|
||||
let bytes = fs::read(self.claim_path(ticket_id))?;
|
||||
Ok(serde_json::from_slice(&bytes)?)
|
||||
}
|
||||
|
||||
pub(crate) fn claim_for_ticket(
|
||||
&self,
|
||||
ticket_id: &str,
|
||||
) -> Result<Option<TicketClaim>, PanelRegistryError> {
|
||||
match self.load_claim(ticket_id) {
|
||||
Ok(claim) => Ok(Some(claim)),
|
||||
Err(PanelRegistryError::Io(error)) if error.kind() == io::ErrorKind::NotFound => {
|
||||
Ok(None)
|
||||
}
|
||||
Err(error) => Err(error),
|
||||
}
|
||||
}
|
||||
|
||||
fn update_registry(
|
||||
&self,
|
||||
update: impl FnOnce(&mut RoleSessionRegistry) -> Result<(), PanelRegistryError>,
|
||||
) -> Result<(), PanelRegistryError> {
|
||||
fs::create_dir_all(&self.root)?;
|
||||
let _lock = self.acquire_registry_lock()?;
|
||||
let mut registry = self.load_registry()?;
|
||||
registry.version = REGISTRY_VERSION;
|
||||
if let Some(workspace_root) = self.workspace_root.as_ref() {
|
||||
registry.workspace_root = workspace_root.clone();
|
||||
}
|
||||
update(&mut registry)?;
|
||||
self.save_registry(®istry)
|
||||
}
|
||||
|
||||
fn acquire_registry_lock(&self) -> Result<RegistryLockGuard, PanelRegistryError> {
|
||||
let lock_path = self.root.join(REGISTRY_LOCK_FILE);
|
||||
for _ in 0..50 {
|
||||
match OpenOptions::new()
|
||||
.write(true)
|
||||
.create_new(true)
|
||||
.open(&lock_path)
|
||||
{
|
||||
Ok(_) => return Ok(RegistryLockGuard { path: lock_path }),
|
||||
Err(error) if error.kind() == io::ErrorKind::AlreadyExists => {
|
||||
thread::sleep(Duration::from_millis(10));
|
||||
}
|
||||
Err(error) => return Err(error.into()),
|
||||
}
|
||||
}
|
||||
Err(PanelRegistryError::Io(io::Error::new(
|
||||
io::ErrorKind::WouldBlock,
|
||||
"timed out acquiring panel role session registry lock",
|
||||
)))
|
||||
}
|
||||
|
||||
fn save_registry(&self, registry: &RoleSessionRegistry) -> Result<(), PanelRegistryError> {
|
||||
let path = self.registry_path();
|
||||
let temp_path = path.with_extension(format!("json.{}.tmp", now_timestamp_string()));
|
||||
let bytes = serde_json::to_vec_pretty(registry)?;
|
||||
fs::write(&temp_path, [&bytes[..], b"\n"].concat())?;
|
||||
fs::rename(temp_path, path)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn create_claim_file(&self, claim_path: &Path, claim: &TicketClaim) -> io::Result<()> {
|
||||
let temp_path = self
|
||||
.claims_dir()
|
||||
.join(format!(".{}.tmp", now_timestamp_string()));
|
||||
let bytes = serde_json::to_vec_pretty(claim).map_err(io::Error::other)?;
|
||||
fs::write(&temp_path, [&bytes[..], b"\n"].concat())?;
|
||||
let link_result = fs::hard_link(&temp_path, claim_path);
|
||||
let remove_result = fs::remove_file(&temp_path);
|
||||
match (link_result, remove_result) {
|
||||
(Ok(()), Ok(())) | (Ok(()), Err(_)) => Ok(()),
|
||||
(Err(error), _) => Err(error),
|
||||
}
|
||||
}
|
||||
|
||||
fn load_claims(&self) -> Result<Vec<TicketClaim>, PanelRegistryError> {
|
||||
let mut claims: Vec<TicketClaim> = Vec::new();
|
||||
match fs::read_dir(self.claims_dir()) {
|
||||
Ok(entries) => {
|
||||
for entry in entries {
|
||||
let entry = entry?;
|
||||
if entry.file_type()?.is_file()
|
||||
&& entry
|
||||
.path()
|
||||
.extension()
|
||||
.is_some_and(|extension| extension == "json")
|
||||
{
|
||||
let bytes = fs::read(entry.path())?;
|
||||
claims.push(serde_json::from_slice(&bytes)?);
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(error) if error.kind() == io::ErrorKind::NotFound => {}
|
||||
Err(error) => return Err(error.into()),
|
||||
}
|
||||
claims.sort_by(|left, right| left.ticket_id.cmp(&right.ticket_id));
|
||||
Ok(claims)
|
||||
}
|
||||
|
||||
fn registry_path(&self) -> PathBuf {
|
||||
self.root.join(REGISTRY_FILE)
|
||||
}
|
||||
|
||||
fn claims_dir(&self) -> PathBuf {
|
||||
self.root.join(CLAIMS_DIR)
|
||||
}
|
||||
|
||||
fn claim_path(&self, ticket_id: &str) -> PathBuf {
|
||||
self.claims_dir()
|
||||
.join(format!("{}.json", encode_path_component(ticket_id)))
|
||||
}
|
||||
}
|
||||
|
||||
struct RegistryLockGuard {
|
||||
path: PathBuf,
|
||||
}
|
||||
|
||||
impl Drop for RegistryLockGuard {
|
||||
fn drop(&mut self) {
|
||||
let _ = fs::remove_file(&self.path);
|
||||
}
|
||||
}
|
||||
|
||||
impl PanelRegistrySnapshot {
|
||||
pub(crate) fn empty() -> Self {
|
||||
Self {
|
||||
sessions: Vec::new(),
|
||||
claims: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn claim_for_ticket(&self, ticket_id: &str) -> Option<&TicketClaim> {
|
||||
self.claims
|
||||
.iter()
|
||||
.find(|claim| claim.ticket_id == ticket_id)
|
||||
}
|
||||
}
|
||||
|
||||
fn normalized_workspace_key(path: &Path) -> String {
|
||||
path.to_string_lossy().replace('\\', "/")
|
||||
}
|
||||
|
||||
fn workspace_leaf(workspace_root: &str) -> String {
|
||||
let leaf = workspace_root
|
||||
.rsplit('/')
|
||||
.find(|part| !part.is_empty())
|
||||
.unwrap_or("workspace");
|
||||
let sanitized = leaf
|
||||
.chars()
|
||||
.map(|ch| {
|
||||
if ch.is_ascii_alphanumeric() || matches!(ch, '-' | '_') {
|
||||
ch
|
||||
} else {
|
||||
'-'
|
||||
}
|
||||
})
|
||||
.collect::<String>()
|
||||
.trim_matches('-')
|
||||
.to_string();
|
||||
if sanitized.is_empty() {
|
||||
"workspace".to_string()
|
||||
} else {
|
||||
sanitized
|
||||
}
|
||||
}
|
||||
|
||||
fn fnv1a64_hex(bytes: &[u8]) -> String {
|
||||
let mut hash = 0xcbf29ce484222325u64;
|
||||
for byte in bytes {
|
||||
hash ^= u64::from(*byte);
|
||||
hash = hash.wrapping_mul(0x100000001b3);
|
||||
}
|
||||
format!("{hash:016x}")
|
||||
}
|
||||
|
||||
fn encode_path_component(value: &str) -> String {
|
||||
let mut encoded = String::with_capacity(value.len());
|
||||
for byte in value.bytes() {
|
||||
match byte {
|
||||
b'a'..=b'z' | b'A'..=b'Z' | b'0'..=b'9' | b'-' | b'_' => encoded.push(byte as char),
|
||||
_ => encoded.push_str(&format!("%{byte:02X}")),
|
||||
}
|
||||
}
|
||||
encoded
|
||||
}
|
||||
|
||||
fn now_timestamp_string() -> String {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|duration| duration.as_nanos().to_string())
|
||||
.unwrap_or_else(|_| "0".to_string())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use tempfile::TempDir;
|
||||
|
||||
#[test]
|
||||
fn registry_path_is_workspace_scoped_under_data_dir() {
|
||||
let data_dir = TempDir::new().unwrap();
|
||||
let store = PanelRegistryStore::for_data_dir(data_dir.path(), Path::new("/repo/yoi"));
|
||||
let other = PanelRegistryStore::for_data_dir(data_dir.path(), Path::new("/repo/other"));
|
||||
|
||||
assert!(store.root().starts_with(data_dir.path()));
|
||||
let root = store.root().to_string_lossy();
|
||||
assert!(root.contains("panel/workspaces/yoi-"));
|
||||
assert_ne!(store.root(), other.root());
|
||||
|
||||
store
|
||||
.record_session(
|
||||
"ticket-intake-preticket",
|
||||
"intake",
|
||||
RoleSessionOrigin::PreTicketIntake,
|
||||
None,
|
||||
[],
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(store.load_registry().unwrap().workspace_root, "/repo/yoi");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claim_ticket_rejects_second_active_local_pod() {
|
||||
let temp = TempDir::new().unwrap();
|
||||
let store = PanelRegistryStore::from_root(temp.path().join("registry"));
|
||||
|
||||
assert!(matches!(
|
||||
store.claim_ticket("T-1", Some("ticket-one"), "ticket-one-intake", "intake"),
|
||||
Ok(TicketClaimResult::Claimed)
|
||||
));
|
||||
|
||||
let error = store
|
||||
.claim_ticket("T-1", Some("ticket-one"), "ticket-two-intake", "intake")
|
||||
.unwrap_err();
|
||||
assert!(matches!(error, PanelRegistryError::TicketAlreadyClaimed(_)));
|
||||
let claim = store.claim_for_ticket("T-1").unwrap().unwrap();
|
||||
assert_eq!(claim.worker_name, "ticket-one-intake");
|
||||
assert_eq!(claim.ticket_slug.as_deref(), Some("ticket-one"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn intake_session_relation_is_not_one_to_one_with_tickets() {
|
||||
let temp = TempDir::new().unwrap();
|
||||
let store = PanelRegistryStore::from_root(temp.path().join("registry"));
|
||||
|
||||
store
|
||||
.record_session(
|
||||
"ticket-intake-preticket",
|
||||
"intake",
|
||||
RoleSessionOrigin::PreTicketIntake,
|
||||
None,
|
||||
[],
|
||||
)
|
||||
.unwrap();
|
||||
store
|
||||
.record_session(
|
||||
"ticket-intake-shared",
|
||||
"intake",
|
||||
RoleSessionOrigin::RoleLaunch,
|
||||
None,
|
||||
[
|
||||
RelatedTicketRef {
|
||||
id: "T-1".to_string(),
|
||||
slug: Some("one".to_string()),
|
||||
},
|
||||
RelatedTicketRef {
|
||||
id: "T-2".to_string(),
|
||||
slug: Some("two".to_string()),
|
||||
},
|
||||
],
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let snapshot = store.snapshot().unwrap();
|
||||
let preticket = snapshot
|
||||
.sessions
|
||||
.iter()
|
||||
.find(|session| session.worker_name == "ticket-intake-preticket")
|
||||
.unwrap();
|
||||
let shared = snapshot
|
||||
.sessions
|
||||
.iter()
|
||||
.find(|session| session.worker_name == "ticket-intake-shared")
|
||||
.unwrap();
|
||||
|
||||
assert!(preticket.related_tickets.is_empty());
|
||||
assert_eq!(shared.role, "intake");
|
||||
assert_eq!(shared.origin, RoleSessionOrigin::RoleLaunch);
|
||||
assert!(!shared.created_at.is_empty());
|
||||
assert!(!shared.updated_at.is_empty());
|
||||
assert_eq!(
|
||||
shared.related_tickets,
|
||||
vec![
|
||||
RelatedTicketRef {
|
||||
id: "T-1".to_string(),
|
||||
slug: Some("one".to_string()),
|
||||
},
|
||||
RelatedTicketRef {
|
||||
id: "T-2".to_string(),
|
||||
slug: Some("two".to_string()),
|
||||
},
|
||||
]
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -228,7 +228,7 @@ worker_context_max_tokens = 100000
|
||||
enabled = true
|
||||
|
||||
[feature.memory]
|
||||
enabled = true
|
||||
enabled = false
|
||||
|
||||
[feature.web]
|
||||
enabled = true
|
||||
@@ -241,11 +241,6 @@ enabled = true
|
||||
authoring = true
|
||||
thread = true
|
||||
|
||||
[memory]
|
||||
extract_threshold = 50000
|
||||
consolidation_threshold_files = 5
|
||||
consolidation_threshold_bytes = 50000
|
||||
|
||||
[web]
|
||||
enabled = true
|
||||
|
||||
|
||||
@@ -1,756 +0,0 @@
|
||||
//! Inline-viewport "spawn Worker and attach" UX.
|
||||
//!
|
||||
//! Rendered at the user's current cursor position when `yoi` is invoked
|
||||
//! with no positional argument. Uses user-configured and bundled Profile
|
||||
//! choices plus bundled profiles, defaults to the builtin profile, prompts for
|
||||
//! the Worker's name, and on confirmation launches the Worker runtime command as an
|
||||
//! independent process. Once the process reports its socket via the
|
||||
//! `YOI-READY` stderr line, the dialog hands control back so main can
|
||||
//! switch the terminal to alternate-screen mode.
|
||||
//!
|
||||
//! The viewport's last frame stays in the terminal's scrollback so the
|
||||
//! user has a record of what was spawned (or why a spawn failed).
|
||||
|
||||
use std::io;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::time::Duration;
|
||||
|
||||
use client::{SpawnConfig, WorkerRuntimeCommand, spawn_worker};
|
||||
use crossterm::event::{self, Event as TermEvent, KeyCode, KeyEventKind, KeyModifiers};
|
||||
use manifest::ProfileDiscovery;
|
||||
use ratatui::Terminal;
|
||||
use ratatui::backend::CrosstermBackend;
|
||||
use ratatui::layout::{Constraint, Layout};
|
||||
use ratatui::style::{Color, Modifier, Style};
|
||||
use ratatui::text::{Line, Span};
|
||||
use ratatui::widgets::Paragraph;
|
||||
use ratatui::{Frame, TerminalOptions, Viewport};
|
||||
use session_store::SegmentId;
|
||||
|
||||
const VIEWPORT_LINES: u16 = 6;
|
||||
|
||||
pub struct SpawnReady {
|
||||
pub worker_name: String,
|
||||
pub socket_path: PathBuf,
|
||||
}
|
||||
|
||||
pub enum SpawnOutcome {
|
||||
Ready(SpawnReady),
|
||||
Cancelled,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum SpawnError {
|
||||
Io(io::Error),
|
||||
Spawn(client::SpawnError),
|
||||
}
|
||||
|
||||
impl std::fmt::Display for SpawnError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::Io(e) => write!(f, "io error: {e}"),
|
||||
Self::Spawn(e) => write!(f, "{e}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for SpawnError {}
|
||||
|
||||
impl From<io::Error> for SpawnError {
|
||||
fn from(e: io::Error) -> Self {
|
||||
Self::Io(e)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<client::SpawnError> for SpawnError {
|
||||
fn from(e: client::SpawnError) -> Self {
|
||||
Self::Spawn(e)
|
||||
}
|
||||
}
|
||||
|
||||
type InlineTerminal = Terminal<CrosstermBackend<io::Stdout>>;
|
||||
|
||||
/// Source session for a resume run. `None` = fresh spawn (current
|
||||
/// behaviour); `Some(id)` swaps the dialog into "Resume Worker" mode and
|
||||
/// passes `--session <id>` to the spawned Worker runtime child.
|
||||
pub async fn run(
|
||||
resume_from: Option<SegmentId>,
|
||||
worker_name: Option<String>,
|
||||
profile: Option<String>,
|
||||
runtime_command: WorkerRuntimeCommand,
|
||||
) -> Result<SpawnOutcome, SpawnError> {
|
||||
let defaults = load_spawn_defaults()?;
|
||||
let mut profile_choices = if resume_from.is_some() {
|
||||
Vec::new()
|
||||
} else {
|
||||
defaults.profile_choices
|
||||
};
|
||||
let profile_index = initial_profile_index(
|
||||
&mut profile_choices,
|
||||
profile.as_deref(),
|
||||
defaults.default_profile_index,
|
||||
);
|
||||
|
||||
let selected_name = worker_name.unwrap_or(defaults.default_name);
|
||||
let immediate = resume_from.is_some() || profile.is_some() && !selected_name.is_empty();
|
||||
let mut form = Form {
|
||||
cwd: defaults.cwd.clone(),
|
||||
scope_origin: defaults.scope_origin,
|
||||
name_cursor: selected_name.chars().count(),
|
||||
name: selected_name,
|
||||
message: None,
|
||||
editing: true,
|
||||
resume_from,
|
||||
profile_choices,
|
||||
profile_index,
|
||||
};
|
||||
|
||||
let mut terminal = make_inline_terminal()?;
|
||||
|
||||
// Phase 1: confirm / cancel.
|
||||
if !immediate {
|
||||
loop {
|
||||
terminal.draw(|f| draw_form(f, &form))?;
|
||||
match poll_event()? {
|
||||
None => continue,
|
||||
Some(Action::Submit) => {
|
||||
if form.name.trim().is_empty() {
|
||||
form.message = Some(("name is required".to_string(), MessageKind::Error));
|
||||
continue;
|
||||
}
|
||||
break;
|
||||
}
|
||||
Some(Action::Cancel) => {
|
||||
form.editing = false;
|
||||
form.message = Some(("cancelled".to_string(), MessageKind::Info));
|
||||
terminal.draw(|f| draw_form(f, &form))?;
|
||||
drop(terminal);
|
||||
return Ok(SpawnOutcome::Cancelled);
|
||||
}
|
||||
Some(Action::Char(c)) => form.insert_char(c),
|
||||
Some(Action::Backspace) => form.backspace(),
|
||||
Some(Action::Delete) => form.delete_forward(),
|
||||
Some(Action::Left) => form.move_left(),
|
||||
Some(Action::Right) => form.move_right(),
|
||||
Some(Action::Home) => form.name_cursor = 0,
|
||||
Some(Action::End) => form.name_cursor = form.name.chars().count(),
|
||||
Some(Action::ProfileNext) => form.cycle_profile_next(),
|
||||
Some(Action::ProfilePrev) => form.cycle_profile_prev(),
|
||||
}
|
||||
}
|
||||
} else if form.name.trim().is_empty() {
|
||||
return Err(SpawnError::Io(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"name is required",
|
||||
)));
|
||||
}
|
||||
|
||||
// Phase 2: launch worker and wait for ready line. Drop the cursor
|
||||
// out of the name field — subsequent frames are passive status
|
||||
// updates, not input — so the cursor doesn't end up parked there
|
||||
// when the inline terminal is finally dropped.
|
||||
form.editing = false;
|
||||
form.message = Some(("starting worker...".to_string(), MessageKind::Progress));
|
||||
terminal.draw(|f| draw_form(f, &form))?;
|
||||
|
||||
match wait_for_ready(&mut terminal, &mut form, &runtime_command).await {
|
||||
Ok(ready) => {
|
||||
form.message = Some((
|
||||
format!("ready: {} attaching...", ready.worker_name),
|
||||
MessageKind::Ok,
|
||||
));
|
||||
terminal.draw(|f| draw_form(f, &form))?;
|
||||
drop(terminal);
|
||||
Ok(SpawnOutcome::Ready(ready))
|
||||
}
|
||||
Err(e) => {
|
||||
form.message = Some((e.to_string(), MessageKind::Error));
|
||||
let _ = terminal.draw(|f| draw_form(f, &form));
|
||||
drop(terminal);
|
||||
Err(e)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Launch a Worker runtime command with `--worker <name>` without opening the name dialog. The child Worker
|
||||
/// resolves persisted Worker metadata if present, or creates a fresh same-name Worker
|
||||
/// from the default profile.
|
||||
pub async fn run_worker_name(
|
||||
worker_name: String,
|
||||
runtime_command: WorkerRuntimeCommand,
|
||||
) -> Result<SpawnOutcome, SpawnError> {
|
||||
let defaults = load_spawn_defaults()?;
|
||||
let mut form = form_for_worker_name(worker_name, defaults);
|
||||
let mut terminal = make_inline_terminal()?;
|
||||
terminal.draw(|f| draw_form(f, &form))?;
|
||||
|
||||
match wait_for_ready(&mut terminal, &mut form, &runtime_command).await {
|
||||
Ok(ready) => {
|
||||
form.message = Some((
|
||||
format!("ready: {} attaching...", ready.worker_name),
|
||||
MessageKind::Ok,
|
||||
));
|
||||
terminal.draw(|f| draw_form(f, &form))?;
|
||||
drop(terminal);
|
||||
Ok(SpawnOutcome::Ready(ready))
|
||||
}
|
||||
Err(e) => {
|
||||
form.message = Some((e.to_string(), MessageKind::Error));
|
||||
let _ = terminal.draw(|f| draw_form(f, &form));
|
||||
drop(terminal);
|
||||
Err(e)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct SpawnDefaults {
|
||||
cwd: PathBuf,
|
||||
scope_origin: ScopeOrigin,
|
||||
default_name: String,
|
||||
default_profile_index: usize,
|
||||
profile_choices: Vec<ProfileChoice>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
struct ProfileChoice {
|
||||
selector: Option<String>,
|
||||
label: String,
|
||||
is_default: bool,
|
||||
}
|
||||
|
||||
fn load_spawn_defaults() -> Result<SpawnDefaults, SpawnError> {
|
||||
let cwd = std::env::current_dir().map_err(SpawnError::Io)?;
|
||||
|
||||
let default_name = cwd
|
||||
.file_name()
|
||||
.and_then(|s| s.to_str())
|
||||
.map(sanitise_default_name)
|
||||
.filter(|s| !s.is_empty())
|
||||
.unwrap_or_else(|| "worker".to_string());
|
||||
|
||||
let (profile_choices, default_profile_index) = profile_choices_for_cwd(&cwd);
|
||||
|
||||
Ok(SpawnDefaults {
|
||||
cwd,
|
||||
scope_origin: ScopeOrigin::FromProfile,
|
||||
default_name,
|
||||
default_profile_index,
|
||||
profile_choices,
|
||||
})
|
||||
}
|
||||
|
||||
fn profile_choices_for_cwd(cwd: &Path) -> (Vec<ProfileChoice>, usize) {
|
||||
let Ok(registry) = ProfileDiscovery::for_cwd(cwd).discover() else {
|
||||
return (Vec::new(), 0);
|
||||
};
|
||||
|
||||
let mut choices = Vec::new();
|
||||
for entry in registry.entries() {
|
||||
let mut label = entry.qualified_name();
|
||||
if entry.is_default {
|
||||
label.push_str(" (default)");
|
||||
}
|
||||
if let Some(description) = entry.description.as_deref() {
|
||||
label.push_str(" — ");
|
||||
label.push_str(description);
|
||||
}
|
||||
choices.push(ProfileChoice {
|
||||
selector: Some(entry.qualified_name()),
|
||||
label,
|
||||
is_default: entry.is_default,
|
||||
});
|
||||
}
|
||||
|
||||
let default_index = choices
|
||||
.iter()
|
||||
.position(|choice| choice.is_default)
|
||||
.unwrap_or(0);
|
||||
(choices, default_index)
|
||||
}
|
||||
|
||||
fn initial_profile_index(
|
||||
choices: &mut Vec<ProfileChoice>,
|
||||
explicit_profile: Option<&str>,
|
||||
default_index: usize,
|
||||
) -> usize {
|
||||
let Some(selector) = explicit_profile else {
|
||||
return default_index.min(choices.len().saturating_sub(1));
|
||||
};
|
||||
if let Some(index) = choices
|
||||
.iter()
|
||||
.position(|choice| choice.selector.as_deref() == Some(selector))
|
||||
{
|
||||
return index;
|
||||
}
|
||||
choices.push(ProfileChoice {
|
||||
selector: Some(selector.to_string()),
|
||||
label: selector.to_string(),
|
||||
is_default: false,
|
||||
});
|
||||
choices.len() - 1
|
||||
}
|
||||
|
||||
fn form_for_worker_name(worker_name: String, defaults: SpawnDefaults) -> Form {
|
||||
Form {
|
||||
cwd: defaults.cwd,
|
||||
scope_origin: defaults.scope_origin,
|
||||
name_cursor: worker_name.chars().count(),
|
||||
name: worker_name,
|
||||
message: Some(("resuming worker...".to_string(), MessageKind::Progress)),
|
||||
editing: false,
|
||||
resume_from: None,
|
||||
profile_choices: Vec::new(),
|
||||
profile_index: 0,
|
||||
}
|
||||
}
|
||||
|
||||
fn make_inline_terminal() -> io::Result<InlineTerminal> {
|
||||
let backend = CrosstermBackend::new(io::stdout());
|
||||
Terminal::with_options(
|
||||
backend,
|
||||
TerminalOptions {
|
||||
viewport: Viewport::Inline(VIEWPORT_LINES),
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
enum Action {
|
||||
Submit,
|
||||
Cancel,
|
||||
Char(char),
|
||||
Backspace,
|
||||
Delete,
|
||||
Left,
|
||||
Right,
|
||||
Home,
|
||||
End,
|
||||
ProfileNext,
|
||||
ProfilePrev,
|
||||
}
|
||||
|
||||
fn poll_event() -> io::Result<Option<Action>> {
|
||||
if !event::poll(Duration::from_millis(100))? {
|
||||
return Ok(None);
|
||||
}
|
||||
match event::read()? {
|
||||
TermEvent::Key(k) if k.kind != KeyEventKind::Release => {
|
||||
let ctrl = k.modifiers.contains(KeyModifiers::CONTROL);
|
||||
Ok(match k.code {
|
||||
KeyCode::Enter => Some(Action::Submit),
|
||||
KeyCode::Esc => Some(Action::Cancel),
|
||||
KeyCode::Char('c') if ctrl => Some(Action::Cancel),
|
||||
KeyCode::Char('a') if ctrl => Some(Action::Home),
|
||||
KeyCode::Char('e') if ctrl => Some(Action::End),
|
||||
KeyCode::Char('u') if ctrl => Some(Action::Cancel),
|
||||
KeyCode::Backspace => Some(Action::Backspace),
|
||||
KeyCode::Delete => Some(Action::Delete),
|
||||
KeyCode::Left => Some(Action::Left),
|
||||
KeyCode::Right => Some(Action::Right),
|
||||
KeyCode::Up | KeyCode::BackTab => Some(Action::ProfilePrev),
|
||||
KeyCode::Down | KeyCode::Tab => Some(Action::ProfileNext),
|
||||
KeyCode::Home => Some(Action::Home),
|
||||
KeyCode::End => Some(Action::End),
|
||||
KeyCode::Char(c) if !ctrl && is_safe_name_char(c) => Some(Action::Char(c)),
|
||||
_ => None,
|
||||
})
|
||||
}
|
||||
_ => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
fn is_safe_name_char(c: char) -> bool {
|
||||
// Filesystem-safe; worker.name becomes a runtime-dir name.
|
||||
c.is_ascii_alphanumeric() || matches!(c, '-' | '_' | '.')
|
||||
}
|
||||
|
||||
fn sanitise_default_name(s: &str) -> String {
|
||||
s.chars()
|
||||
.map(|c| if is_safe_name_char(c) { c } else { '-' })
|
||||
.collect()
|
||||
}
|
||||
|
||||
async fn wait_for_ready(
|
||||
terminal: &mut InlineTerminal,
|
||||
form: &mut Form,
|
||||
runtime_command: &WorkerRuntimeCommand,
|
||||
) -> Result<SpawnReady, SpawnError> {
|
||||
let config = SpawnConfig {
|
||||
runtime_command: runtime_command.clone(),
|
||||
worker_name: form.name.clone(),
|
||||
profile: form.selected_profile_selector(),
|
||||
workspace_root: form.cwd.clone(),
|
||||
cwd: None,
|
||||
resume_from: form.resume_from,
|
||||
};
|
||||
let ready = spawn_worker(config, |line| {
|
||||
form.message = Some((line.to_string(), MessageKind::Progress));
|
||||
let _ = terminal.draw(|f| draw_form(f, form));
|
||||
})
|
||||
.await?;
|
||||
Ok(SpawnReady {
|
||||
worker_name: ready.worker_name,
|
||||
socket_path: ready.socket_path,
|
||||
})
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum MessageKind {
|
||||
Info,
|
||||
Ok,
|
||||
Error,
|
||||
Progress,
|
||||
}
|
||||
|
||||
enum ScopeOrigin {
|
||||
FromProfile,
|
||||
}
|
||||
|
||||
struct Form {
|
||||
cwd: PathBuf,
|
||||
/// Display label for the scope row in the dialog.
|
||||
scope_origin: ScopeOrigin,
|
||||
name: String,
|
||||
/// Cursor position counted in **chars**, not bytes — `name`
|
||||
/// currently only accepts ASCII so the two coincide, but we keep
|
||||
/// char-based bookkeeping in case we relax `is_safe_name_char`.
|
||||
name_cursor: usize,
|
||||
message: Option<(String, MessageKind)>,
|
||||
/// True while the dialog is accepting name input. Drives whether
|
||||
/// the rendered frame parks the terminal cursor inside the name
|
||||
/// field — when false (post-confirm / cancel / failure frames) the
|
||||
/// cursor stays out so it does not collide with the shell prompt
|
||||
/// after the inline terminal is dropped.
|
||||
editing: bool,
|
||||
/// `Some(id)` flips the dialog into "Resume Worker" mode: the title
|
||||
/// switches, the source session is shown to the user, and the
|
||||
/// child worker is launched with `--session <id>` so it restores
|
||||
/// from `id` and appends to the same session log.
|
||||
resume_from: Option<SegmentId>,
|
||||
/// Optional profile choices passed with `--profile` for
|
||||
/// fresh spawns. This is not used for resume/attach flows because those must
|
||||
/// restore Worker state rather than re-evaluate a profile source.
|
||||
profile_choices: Vec<ProfileChoice>,
|
||||
profile_index: usize,
|
||||
}
|
||||
|
||||
impl Form {
|
||||
fn insert_char(&mut self, c: char) {
|
||||
let byte = self.char_offset_to_byte(self.name_cursor);
|
||||
self.name.insert(byte, c);
|
||||
self.name_cursor += 1;
|
||||
}
|
||||
|
||||
fn backspace(&mut self) {
|
||||
if self.name_cursor == 0 {
|
||||
return;
|
||||
}
|
||||
let end = self.char_offset_to_byte(self.name_cursor);
|
||||
let start = self.char_offset_to_byte(self.name_cursor - 1);
|
||||
self.name.replace_range(start..end, "");
|
||||
self.name_cursor -= 1;
|
||||
}
|
||||
|
||||
fn delete_forward(&mut self) {
|
||||
let total = self.name.chars().count();
|
||||
if self.name_cursor >= total {
|
||||
return;
|
||||
}
|
||||
let start = self.char_offset_to_byte(self.name_cursor);
|
||||
let end = self.char_offset_to_byte(self.name_cursor + 1);
|
||||
self.name.replace_range(start..end, "");
|
||||
}
|
||||
|
||||
fn move_left(&mut self) {
|
||||
if self.name_cursor > 0 {
|
||||
self.name_cursor -= 1;
|
||||
}
|
||||
}
|
||||
|
||||
fn move_right(&mut self) {
|
||||
let total = self.name.chars().count();
|
||||
if self.name_cursor < total {
|
||||
self.name_cursor += 1;
|
||||
}
|
||||
}
|
||||
|
||||
fn selected_profile(&self) -> Option<&ProfileChoice> {
|
||||
self.profile_choices
|
||||
.get(self.profile_index)
|
||||
.filter(|choice| choice.selector.is_some())
|
||||
}
|
||||
|
||||
fn selected_profile_selector(&self) -> Option<String> {
|
||||
self.selected_profile()
|
||||
.and_then(|choice| choice.selector.clone())
|
||||
}
|
||||
|
||||
fn cycle_profile_next(&mut self) {
|
||||
if self.profile_choices.is_empty() {
|
||||
return;
|
||||
}
|
||||
self.profile_index = (self.profile_index + 1) % self.profile_choices.len();
|
||||
self.message = None;
|
||||
}
|
||||
|
||||
fn cycle_profile_prev(&mut self) {
|
||||
if self.profile_choices.is_empty() {
|
||||
return;
|
||||
}
|
||||
self.profile_index = if self.profile_index == 0 {
|
||||
self.profile_choices.len() - 1
|
||||
} else {
|
||||
self.profile_index - 1
|
||||
};
|
||||
self.message = None;
|
||||
}
|
||||
|
||||
fn char_offset_to_byte(&self, char_off: usize) -> usize {
|
||||
self.name
|
||||
.char_indices()
|
||||
.nth(char_off)
|
||||
.map(|(b, _)| b)
|
||||
.unwrap_or(self.name.len())
|
||||
}
|
||||
}
|
||||
|
||||
fn draw_form(f: &mut Frame<'_>, form: &Form) {
|
||||
let area = f.area();
|
||||
let layout = Layout::vertical([
|
||||
Constraint::Length(1), // title
|
||||
Constraint::Length(1), // name field
|
||||
Constraint::Length(1), // context (profile or scope default)
|
||||
Constraint::Length(1), // hint
|
||||
Constraint::Length(1), // message
|
||||
Constraint::Length(1), // spacer
|
||||
])
|
||||
.split(area);
|
||||
|
||||
let title_text = match form.resume_from {
|
||||
Some(id) => format!("resume worker session: {}", short_segment(id)),
|
||||
None => "spawn worker".to_string(),
|
||||
};
|
||||
let title = Paragraph::new(Line::from(vec![Span::styled(
|
||||
title_text,
|
||||
Style::default().add_modifier(Modifier::BOLD),
|
||||
)]));
|
||||
f.render_widget(title, layout[0]);
|
||||
|
||||
f.render_widget(Paragraph::new(name_line(form)), layout[1]);
|
||||
f.render_widget(Paragraph::new(context_line(form)), layout[2]);
|
||||
f.render_widget(Paragraph::new(hint_line()), layout[3]);
|
||||
f.render_widget(Paragraph::new(message_line(form)), layout[4]);
|
||||
|
||||
if form.editing {
|
||||
// Place the cursor inside the name field while the user is
|
||||
// editing. Skipped on post-confirm frames so the inline
|
||||
// viewport's drop leaves the cursor at the bottom of the
|
||||
// rendered area rather than parked on the name line, which
|
||||
// would let the shell prompt (or any later eprintln) clobber
|
||||
// the rendered name field after exit.
|
||||
let cursor_col = 2 + "name: ".len() + form.name_cursor;
|
||||
f.set_cursor_position((layout[1].x + cursor_col as u16, layout[1].y));
|
||||
}
|
||||
}
|
||||
|
||||
/// First 8 hex digits of a UUID — short enough to skim, long enough
|
||||
/// to disambiguate inside a 10-row picker.
|
||||
pub(crate) fn short_segment(id: SegmentId) -> String {
|
||||
let s = id.to_string();
|
||||
s.chars().take(8).collect()
|
||||
}
|
||||
|
||||
fn name_line(form: &Form) -> Line<'_> {
|
||||
Line::from(vec![
|
||||
Span::raw(" "),
|
||||
Span::styled("name: ", Style::default().fg(Color::DarkGray)),
|
||||
Span::styled(
|
||||
form.name.as_str(),
|
||||
Style::default()
|
||||
.fg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD),
|
||||
),
|
||||
])
|
||||
}
|
||||
|
||||
fn context_line(form: &Form) -> Line<'_> {
|
||||
if let Some(profile) = form.profile_choices.get(form.profile_index) {
|
||||
return Line::from(vec![
|
||||
Span::raw(" "),
|
||||
Span::styled("profile: ", Style::default().fg(Color::DarkGray)),
|
||||
Span::styled(profile.label.as_str(), Style::default().fg(Color::Green)),
|
||||
Span::styled(
|
||||
" (tab/down to change)",
|
||||
Style::default().fg(Color::DarkGray),
|
||||
),
|
||||
]);
|
||||
}
|
||||
|
||||
match form.scope_origin {
|
||||
ScopeOrigin::FromProfile => Line::from(vec![
|
||||
Span::raw(" "),
|
||||
Span::styled("scope: ", Style::default().fg(Color::DarkGray)),
|
||||
Span::styled("from selected profile", Style::default().fg(Color::Green)),
|
||||
]),
|
||||
}
|
||||
}
|
||||
|
||||
fn hint_line() -> Line<'static> {
|
||||
Line::from(vec![Span::styled(
|
||||
" enter spawn · tab/down next profile · shift-tab/up prev · esc cancel",
|
||||
Style::default().fg(Color::DarkGray),
|
||||
)])
|
||||
}
|
||||
|
||||
fn message_line(form: &Form) -> Line<'_> {
|
||||
let Some((text, kind)) = form.message.as_ref() else {
|
||||
return Line::from("");
|
||||
};
|
||||
let style = match kind {
|
||||
MessageKind::Info => Style::default().fg(Color::DarkGray),
|
||||
MessageKind::Ok => Style::default().fg(Color::Green),
|
||||
MessageKind::Error => Style::default().fg(Color::Red),
|
||||
MessageKind::Progress => Style::default().fg(Color::Yellow),
|
||||
};
|
||||
Line::from(vec![Span::raw(" "), Span::styled(text.as_str(), style)])
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn form(name: &str) -> Form {
|
||||
Form {
|
||||
cwd: PathBuf::from("/work/example"),
|
||||
scope_origin: ScopeOrigin::FromProfile,
|
||||
name: name.to_string(),
|
||||
name_cursor: name.chars().count(),
|
||||
message: None,
|
||||
editing: true,
|
||||
resume_from: None,
|
||||
profile_choices: Vec::new(),
|
||||
profile_index: 0,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn worker_name_form_restores_or_creates_by_worker_name() {
|
||||
let defaults = SpawnDefaults {
|
||||
cwd: PathBuf::from("/work/example"),
|
||||
scope_origin: ScopeOrigin::FromProfile,
|
||||
default_name: "ignored".to_string(),
|
||||
default_profile_index: 0,
|
||||
profile_choices: Vec::new(),
|
||||
};
|
||||
let f = form_for_worker_name("agent".to_string(), defaults);
|
||||
|
||||
assert_eq!(f.name, "agent");
|
||||
assert_eq!(f.name_cursor, "agent".chars().count());
|
||||
assert_eq!(f.resume_from, None);
|
||||
assert!(!f.editing);
|
||||
assert_eq!(
|
||||
f.message,
|
||||
Some(("resuming worker...".to_string(), MessageKind::Progress))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn profile_choices_ignore_repository_local_profile_registry() {
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
let project = temp.path().join("project");
|
||||
let yoi = project.join(".yoi");
|
||||
std::fs::create_dir_all(&yoi).unwrap();
|
||||
std::fs::write(
|
||||
yoi.join("profiles.toml"),
|
||||
"default = \"coder\"\n[profile]\ncoder = \"profiles/coder.toml\"\n",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let (choices, default_index) = profile_choices_for_cwd(&project);
|
||||
assert_eq!(default_index, 0);
|
||||
assert!(
|
||||
choices
|
||||
.iter()
|
||||
.all(|choice| { choice.selector.as_deref() != Some("project:coder") })
|
||||
);
|
||||
assert!(
|
||||
choices
|
||||
.iter()
|
||||
.any(|choice| { choice.selector.as_deref() == Some("builtin:companion") })
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn profile_cycle_selects_only_discovered_profiles() {
|
||||
let mut form = form("coder");
|
||||
form.profile_choices = vec![
|
||||
ProfileChoice {
|
||||
selector: Some("project:coder".to_string()),
|
||||
label: "project:coder (default)".to_string(),
|
||||
is_default: true,
|
||||
},
|
||||
ProfileChoice {
|
||||
selector: Some("user:reviewer".to_string()),
|
||||
label: "user:reviewer".to_string(),
|
||||
is_default: false,
|
||||
},
|
||||
];
|
||||
form.profile_index = 0;
|
||||
|
||||
assert_eq!(
|
||||
form.selected_profile_selector().as_deref(),
|
||||
Some("project:coder")
|
||||
);
|
||||
form.cycle_profile_next();
|
||||
assert_eq!(
|
||||
form.selected_profile_selector().as_deref(),
|
||||
Some("user:reviewer")
|
||||
);
|
||||
form.cycle_profile_next();
|
||||
assert_eq!(
|
||||
form.selected_profile_selector().as_deref(),
|
||||
Some("project:coder")
|
||||
);
|
||||
form.cycle_profile_prev();
|
||||
assert_eq!(
|
||||
form.selected_profile_selector().as_deref(),
|
||||
Some("user:reviewer")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn initial_profile_index_adds_explicit_selector_not_in_discovery_list() {
|
||||
let mut choices = Vec::new();
|
||||
let selected = initial_profile_index(&mut choices, Some("coder"), 0);
|
||||
assert_eq!(selected, 0);
|
||||
assert_eq!(choices[0].selector.as_deref(), Some("coder"));
|
||||
assert_eq!(choices[0].label, "coder");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn name_input_handles_insert_backspace_and_cursor() {
|
||||
let mut f = form("");
|
||||
for c in "abc".chars() {
|
||||
f.insert_char(c);
|
||||
}
|
||||
assert_eq!(f.name, "abc");
|
||||
assert_eq!(f.name_cursor, 3);
|
||||
|
||||
f.move_left();
|
||||
f.move_left();
|
||||
f.insert_char('X');
|
||||
assert_eq!(f.name, "aXbc");
|
||||
|
||||
f.backspace();
|
||||
assert_eq!(f.name, "abc");
|
||||
assert_eq!(f.name_cursor, 1);
|
||||
|
||||
f.delete_forward();
|
||||
assert_eq!(f.name, "ac");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sanitise_default_name_replaces_unsafe_chars() {
|
||||
assert_eq!(sanitise_default_name("my project!"), "my-project-");
|
||||
assert_eq!(sanitise_default_name("ok-name_2.0"), "ok-name_2.0");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,163 @@
|
||||
use std::io;
|
||||
use std::time::Duration;
|
||||
|
||||
use client::{StandaloneWorkerListIntent, StandaloneWorkerResumeIntent, Target};
|
||||
use crossterm::event::{self, Event as TermEvent, KeyCode, KeyEventKind, KeyModifiers};
|
||||
use ratatui::layout::{Constraint, Layout};
|
||||
use ratatui::prelude::{Color, Line, Modifier, Span, Style};
|
||||
use ratatui::widgets::Paragraph;
|
||||
use standalone::{StandaloneListScope, StandaloneWorkerRecord, StandaloneWorkerStore};
|
||||
use thiserror::Error;
|
||||
|
||||
use crate::inline_terminal::with_inline_terminal;
|
||||
|
||||
const LIMIT: usize = 100;
|
||||
|
||||
pub(crate) fn pick(
|
||||
target: &dyn Target,
|
||||
include_all: bool,
|
||||
) -> Result<Option<StandaloneWorkerResumeIntent>, StandalonePickerError> {
|
||||
let intent = target
|
||||
.standalone_worker_list(include_all)
|
||||
.map_err(StandalonePickerError::Target)?;
|
||||
let records = load_records(&intent)?;
|
||||
if records.is_empty() {
|
||||
return Err(StandalonePickerError::NoWorkers { include_all });
|
||||
}
|
||||
let selected = run_picker(records)?;
|
||||
selected
|
||||
.map(|record| {
|
||||
target
|
||||
.standalone_worker_resume(record.worker_id.to_string())
|
||||
.map_err(StandalonePickerError::Target)
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn load_records(
|
||||
intent: &StandaloneWorkerListIntent,
|
||||
) -> Result<Vec<StandaloneWorkerRecord>, StandalonePickerError> {
|
||||
let store = StandaloneWorkerStore::open(&intent.state_dir)
|
||||
.map_err(StandalonePickerError::StateStore)?;
|
||||
store
|
||||
.list(
|
||||
&intent.cwd,
|
||||
if intent.include_all {
|
||||
StandaloneListScope::All
|
||||
} else {
|
||||
StandaloneListScope::CurrentCwd
|
||||
},
|
||||
LIMIT,
|
||||
)
|
||||
.map_err(StandalonePickerError::StateStore)
|
||||
}
|
||||
|
||||
fn run_picker(
|
||||
records: Vec<StandaloneWorkerRecord>,
|
||||
) -> Result<Option<StandaloneWorkerRecord>, StandalonePickerError> {
|
||||
let height = u16::try_from(records.len().saturating_add(3).min(20)).unwrap_or(20);
|
||||
with_inline_terminal(height, |terminal| {
|
||||
let mut selected = 0usize;
|
||||
loop {
|
||||
terminal
|
||||
.draw(|frame| draw(frame, &records, selected))
|
||||
.map_err(StandalonePickerError::Io)?;
|
||||
if !event::poll(Duration::from_millis(100)).map_err(StandalonePickerError::Io)? {
|
||||
continue;
|
||||
}
|
||||
let TermEvent::Key(key) = event::read().map_err(StandalonePickerError::Io)? else {
|
||||
continue;
|
||||
};
|
||||
if key.kind == KeyEventKind::Release {
|
||||
continue;
|
||||
}
|
||||
let ctrl = key.modifiers.contains(KeyModifiers::CONTROL);
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') if !ctrl => {
|
||||
selected = selected.saturating_sub(1);
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if !ctrl => {
|
||||
selected = (selected + 1).min(records.len() - 1);
|
||||
}
|
||||
KeyCode::Enter => return Ok(Some(records[selected].clone())),
|
||||
KeyCode::Esc => return Ok(None),
|
||||
KeyCode::Char('c') if ctrl => return Ok(None),
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
fn draw(frame: &mut ratatui::Frame<'_>, records: &[StandaloneWorkerRecord], selected: usize) {
|
||||
let mut constraints = vec![Constraint::Length(1)];
|
||||
constraints.extend(records.iter().map(|_| Constraint::Length(1)));
|
||||
constraints.push(Constraint::Length(1));
|
||||
let rows = Layout::vertical(constraints).split(frame.area());
|
||||
frame.render_widget(
|
||||
Paragraph::new(Line::from(Span::styled(
|
||||
"resume standalone Worker",
|
||||
Style::default().add_modifier(Modifier::BOLD),
|
||||
))),
|
||||
rows[0],
|
||||
);
|
||||
for (index, record) in records.iter().enumerate() {
|
||||
let active = index == selected;
|
||||
let marker = if active { "▶ " } else { " " };
|
||||
let style = if active {
|
||||
Style::default()
|
||||
.fg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD)
|
||||
} else {
|
||||
Style::default().fg(Color::DarkGray)
|
||||
};
|
||||
let cwd = record.cwd.canonical_path.display();
|
||||
frame.render_widget(
|
||||
Paragraph::new(Line::from(vec![
|
||||
Span::raw(marker),
|
||||
Span::styled(
|
||||
format!("{} ({})", record.worker_name, record.worker_id.short()),
|
||||
style,
|
||||
),
|
||||
Span::raw(format!(
|
||||
" [{:?}] updated:{} {}",
|
||||
record.status, record.updated_at_unix_ms, cwd
|
||||
)),
|
||||
])),
|
||||
rows[index + 1],
|
||||
);
|
||||
}
|
||||
frame.render_widget(
|
||||
Paragraph::new(" [↑/↓] select [enter] restore [esc] cancel"),
|
||||
rows[records.len() + 1],
|
||||
);
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub(crate) enum StandalonePickerError {
|
||||
#[error("standalone target error: {0}")]
|
||||
Target(#[source] client::TargetError),
|
||||
#[error("standalone Worker state is unavailable: {0}")]
|
||||
StateStore(#[source] standalone::StandaloneStoreError),
|
||||
#[error(
|
||||
"no standalone Workers found for this cwd; use `yoi --local resume --all` to include all cwd identities"
|
||||
)]
|
||||
NoWorkers { include_all: bool },
|
||||
#[error("standalone Worker picker I/O failed: {0}")]
|
||||
Io(#[from] io::Error),
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use client::StandaloneTarget;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn empty_picker_keeps_current_cwd_as_default_scope() {
|
||||
let temp = tempfile::tempdir().expect("tempdir");
|
||||
let target = StandaloneTarget::new(temp.path());
|
||||
let error = pick(&target, false).expect_err("empty picker should fail explicitly");
|
||||
assert!(error.to_string().contains("this cwd"));
|
||||
assert!(error.to_string().contains("--all"));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,488 @@
|
||||
use std::io;
|
||||
use std::path::Path;
|
||||
use std::time::Duration;
|
||||
|
||||
use crossterm::event::{self, Event, KeyCode, KeyEvent, KeyEventKind, KeyModifiers};
|
||||
use manifest::ProfileDiscovery;
|
||||
use ratatui::layout::{Constraint, Direction, Layout};
|
||||
use ratatui::style::{Color, Modifier, Style};
|
||||
use ratatui::text::{Line, Span};
|
||||
use ratatui::widgets::Paragraph;
|
||||
use thiserror::Error;
|
||||
|
||||
use crate::inline_terminal::{InlineTerminal, with_inline_terminal};
|
||||
|
||||
const VIEWPORT_HEIGHT: u16 = 6;
|
||||
const FALLBACK_WORKER_NAME: &str = "worker";
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(crate) struct StandaloneSpawnSelection {
|
||||
pub worker_name: String,
|
||||
pub profile: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub(crate) enum StandaloneSpawnError {
|
||||
#[error("profile discovery failed: {0}")]
|
||||
ProfileDiscovery(#[from] manifest::ProfileError),
|
||||
#[error("no profiles are available")]
|
||||
NoProfiles,
|
||||
#[error("standalone spawn picker terminal error: {0}")]
|
||||
Terminal(#[from] io::Error),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
struct ProfileChoice {
|
||||
selector: String,
|
||||
label: String,
|
||||
is_default: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum StatusKind {
|
||||
Info,
|
||||
Progress,
|
||||
Error,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum SpawnAction {
|
||||
None,
|
||||
Submit,
|
||||
Cancel,
|
||||
}
|
||||
|
||||
struct SpawnForm {
|
||||
worker_name: String,
|
||||
cursor: usize,
|
||||
profile_choices: Vec<ProfileChoice>,
|
||||
selected_profile: usize,
|
||||
status: Option<(String, StatusKind)>,
|
||||
}
|
||||
|
||||
impl SpawnForm {
|
||||
fn new(
|
||||
worker_name: Option<String>,
|
||||
default_worker_name: String,
|
||||
profile_choices: Vec<ProfileChoice>,
|
||||
) -> Self {
|
||||
let worker_name = worker_name.unwrap_or(default_worker_name);
|
||||
let cursor = worker_name.chars().count();
|
||||
let selected_profile = profile_choices
|
||||
.iter()
|
||||
.position(|choice| choice.is_default)
|
||||
.unwrap_or(0);
|
||||
Self {
|
||||
worker_name,
|
||||
cursor,
|
||||
profile_choices,
|
||||
selected_profile,
|
||||
status: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn selected_profile(&self) -> &ProfileChoice {
|
||||
&self.profile_choices[self.selected_profile]
|
||||
}
|
||||
|
||||
fn apply_key(&mut self, key: KeyEvent) -> SpawnAction {
|
||||
if key.kind == KeyEventKind::Release {
|
||||
return SpawnAction::None;
|
||||
}
|
||||
|
||||
if key.modifiers.contains(KeyModifiers::CONTROL) {
|
||||
match key.code {
|
||||
KeyCode::Char('c') | KeyCode::Char('u') => return SpawnAction::Cancel,
|
||||
_ => return SpawnAction::None,
|
||||
}
|
||||
}
|
||||
|
||||
self.status = None;
|
||||
match key.code {
|
||||
KeyCode::Esc => SpawnAction::Cancel,
|
||||
KeyCode::Enter => {
|
||||
if self.worker_name.trim().is_empty() {
|
||||
self.status =
|
||||
Some(("worker name cannot be empty".to_owned(), StatusKind::Error));
|
||||
SpawnAction::None
|
||||
} else {
|
||||
SpawnAction::Submit
|
||||
}
|
||||
}
|
||||
KeyCode::Tab | KeyCode::Down => {
|
||||
self.selected_profile = (self.selected_profile + 1) % self.profile_choices.len();
|
||||
SpawnAction::None
|
||||
}
|
||||
KeyCode::BackTab | KeyCode::Up => {
|
||||
self.selected_profile = if self.selected_profile == 0 {
|
||||
self.profile_choices.len() - 1
|
||||
} else {
|
||||
self.selected_profile - 1
|
||||
};
|
||||
SpawnAction::None
|
||||
}
|
||||
KeyCode::Left => {
|
||||
self.cursor = self.cursor.saturating_sub(1);
|
||||
SpawnAction::None
|
||||
}
|
||||
KeyCode::Right => {
|
||||
self.cursor = (self.cursor + 1).min(self.worker_name.chars().count());
|
||||
SpawnAction::None
|
||||
}
|
||||
KeyCode::Home => {
|
||||
self.cursor = 0;
|
||||
SpawnAction::None
|
||||
}
|
||||
KeyCode::End => {
|
||||
self.cursor = self.worker_name.chars().count();
|
||||
SpawnAction::None
|
||||
}
|
||||
KeyCode::Backspace => {
|
||||
if self.cursor > 0 {
|
||||
let idx = byte_index(&self.worker_name, self.cursor - 1);
|
||||
self.worker_name.remove(idx);
|
||||
self.cursor -= 1;
|
||||
}
|
||||
SpawnAction::None
|
||||
}
|
||||
KeyCode::Delete => {
|
||||
if self.cursor < self.worker_name.chars().count() {
|
||||
let idx = byte_index(&self.worker_name, self.cursor);
|
||||
self.worker_name.remove(idx);
|
||||
}
|
||||
SpawnAction::None
|
||||
}
|
||||
KeyCode::Char(ch) if is_safe_worker_char(ch) => {
|
||||
let idx = byte_index(&self.worker_name, self.cursor);
|
||||
self.worker_name.insert(idx, ch);
|
||||
self.cursor += 1;
|
||||
SpawnAction::None
|
||||
}
|
||||
_ => SpawnAction::None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn select(
|
||||
workspace_root: &Path,
|
||||
worker_name: Option<String>,
|
||||
profile: Option<String>,
|
||||
) -> Result<Option<StandaloneSpawnSelection>, StandaloneSpawnError> {
|
||||
let default_worker_name = default_worker_name(workspace_root);
|
||||
if let Some(profile) = profile {
|
||||
return Ok(Some(StandaloneSpawnSelection {
|
||||
worker_name: worker_name.unwrap_or(default_worker_name),
|
||||
profile,
|
||||
}));
|
||||
}
|
||||
|
||||
let registry = ProfileDiscovery::user_settings().discover()?;
|
||||
let choices = profile_choices(®istry);
|
||||
if choices.is_empty() {
|
||||
return Err(StandaloneSpawnError::NoProfiles);
|
||||
}
|
||||
|
||||
with_inline_terminal(VIEWPORT_HEIGHT, |terminal| {
|
||||
run_picker(
|
||||
terminal,
|
||||
SpawnForm::new(worker_name, default_worker_name, choices),
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn run_picker(
|
||||
terminal: &mut InlineTerminal,
|
||||
mut form: SpawnForm,
|
||||
) -> Result<Option<StandaloneSpawnSelection>, StandaloneSpawnError> {
|
||||
loop {
|
||||
terminal.draw(|frame| draw_form(frame, &form))?;
|
||||
if !event::poll(Duration::from_millis(100))? {
|
||||
continue;
|
||||
}
|
||||
let Event::Key(key) = event::read()? else {
|
||||
continue;
|
||||
};
|
||||
|
||||
match form.apply_key(key) {
|
||||
SpawnAction::None => {}
|
||||
SpawnAction::Cancel => {
|
||||
form.status = Some(("cancelled".to_owned(), StatusKind::Info));
|
||||
terminal.draw(|frame| draw_form(frame, &form))?;
|
||||
return Ok(None);
|
||||
}
|
||||
SpawnAction::Submit => {
|
||||
let selection = StandaloneSpawnSelection {
|
||||
worker_name: form.worker_name.trim().to_owned(),
|
||||
profile: form.selected_profile().selector.clone(),
|
||||
};
|
||||
form.status = Some(("starting worker...".to_owned(), StatusKind::Progress));
|
||||
terminal.draw(|frame| draw_form(frame, &form))?;
|
||||
return Ok(Some(selection));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn profile_choices(registry: &manifest::ProfileRegistry) -> Vec<ProfileChoice> {
|
||||
registry
|
||||
.entries()
|
||||
.iter()
|
||||
.map(|entry| {
|
||||
let selector = entry.qualified_name();
|
||||
let default_marker = if entry.is_default { " (default)" } else { "" };
|
||||
let mut label = format!("{selector}{default_marker}");
|
||||
if let Some(description) = &entry.description {
|
||||
label.push_str(" — ");
|
||||
label.push_str(description);
|
||||
}
|
||||
ProfileChoice {
|
||||
selector,
|
||||
label,
|
||||
is_default: entry.is_default,
|
||||
}
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn draw_form(frame: &mut ratatui::Frame<'_>, form: &SpawnForm) {
|
||||
let chunks = Layout::default()
|
||||
.direction(Direction::Vertical)
|
||||
.constraints([
|
||||
Constraint::Length(1),
|
||||
Constraint::Length(1),
|
||||
Constraint::Length(1),
|
||||
Constraint::Length(1),
|
||||
Constraint::Length(1),
|
||||
Constraint::Min(0),
|
||||
])
|
||||
.split(frame.area());
|
||||
|
||||
frame.render_widget(
|
||||
Paragraph::new(Line::from(vec![
|
||||
Span::raw(" "),
|
||||
Span::styled(
|
||||
"spawn worker",
|
||||
Style::default().add_modifier(Modifier::BOLD),
|
||||
),
|
||||
])),
|
||||
chunks[0],
|
||||
);
|
||||
|
||||
frame.render_widget(
|
||||
Paragraph::new(Line::from(vec![
|
||||
Span::raw(" "),
|
||||
Span::styled("name: ", Style::default().fg(Color::DarkGray)),
|
||||
Span::styled(
|
||||
&form.worker_name,
|
||||
Style::default()
|
||||
.fg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD),
|
||||
),
|
||||
])),
|
||||
chunks[1],
|
||||
);
|
||||
|
||||
frame.render_widget(
|
||||
Paragraph::new(Line::from(vec![
|
||||
Span::raw(" "),
|
||||
Span::styled("profile: ", Style::default().fg(Color::DarkGray)),
|
||||
Span::styled(
|
||||
&form.selected_profile().label,
|
||||
Style::default().fg(Color::Green),
|
||||
),
|
||||
Span::styled(
|
||||
" (tab/down to change)",
|
||||
Style::default().fg(Color::DarkGray),
|
||||
),
|
||||
])),
|
||||
chunks[2],
|
||||
);
|
||||
|
||||
frame.render_widget(
|
||||
Paragraph::new(Line::from(Span::styled(
|
||||
" enter spawn · left/right edit · esc cancel",
|
||||
Style::default().fg(Color::DarkGray),
|
||||
))),
|
||||
chunks[3],
|
||||
);
|
||||
|
||||
let (message, color) = form
|
||||
.status
|
||||
.as_ref()
|
||||
.map(|(message, kind)| {
|
||||
let color = match kind {
|
||||
StatusKind::Info => Color::DarkGray,
|
||||
StatusKind::Progress => Color::Yellow,
|
||||
StatusKind::Error => Color::Red,
|
||||
};
|
||||
(message.as_str(), color)
|
||||
})
|
||||
.unwrap_or(("", Color::Reset));
|
||||
frame.render_widget(
|
||||
Paragraph::new(Line::from(vec![
|
||||
Span::raw(" "),
|
||||
Span::styled(message, Style::default().fg(color)),
|
||||
])),
|
||||
chunks[4],
|
||||
);
|
||||
|
||||
let prefix_width = " name: ".chars().count() as u16;
|
||||
let x = chunks[1]
|
||||
.x
|
||||
.saturating_add(prefix_width)
|
||||
.saturating_add(form.cursor as u16)
|
||||
.min(chunks[1].right().saturating_sub(1));
|
||||
frame.set_cursor_position((x, chunks[1].y));
|
||||
}
|
||||
|
||||
fn default_worker_name(workspace_root: &Path) -> String {
|
||||
workspace_root
|
||||
.file_name()
|
||||
.and_then(|name| name.to_str())
|
||||
.map(sanitise_default_name)
|
||||
.filter(|name| !name.is_empty())
|
||||
.unwrap_or_else(|| FALLBACK_WORKER_NAME.to_owned())
|
||||
}
|
||||
|
||||
fn sanitise_default_name(name: &str) -> String {
|
||||
name.chars()
|
||||
.map(|ch| if is_safe_worker_char(ch) { ch } else { '-' })
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn is_safe_worker_char(ch: char) -> bool {
|
||||
ch.is_ascii_alphanumeric() || matches!(ch, '-' | '_' | '.')
|
||||
}
|
||||
|
||||
fn byte_index(input: &str, char_index: usize) -> usize {
|
||||
input
|
||||
.char_indices()
|
||||
.nth(char_index)
|
||||
.map_or(input.len(), |(idx, _)| idx)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crossterm::event::{KeyEvent, KeyModifiers};
|
||||
|
||||
use super::*;
|
||||
|
||||
fn choices() -> Vec<ProfileChoice> {
|
||||
vec![
|
||||
ProfileChoice {
|
||||
selector: "builtin:default".to_owned(),
|
||||
label: "builtin:default (default) — Default".to_owned(),
|
||||
is_default: true,
|
||||
},
|
||||
ProfileChoice {
|
||||
selector: "builtin:coder".to_owned(),
|
||||
label: "builtin:coder — Coder".to_owned(),
|
||||
is_default: false,
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_form_preserves_old_spawn_layout_defaults() {
|
||||
let form = SpawnForm::new(None, "yoi".to_owned(), choices());
|
||||
assert_eq!(form.worker_name, "yoi");
|
||||
assert_eq!(form.selected_profile().selector, "builtin:default");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tab_and_arrows_cycle_profiles() {
|
||||
let mut form = SpawnForm::new(None, "yoi".to_owned(), choices());
|
||||
assert_eq!(
|
||||
form.apply_key(KeyEvent::new(KeyCode::Tab, KeyModifiers::NONE)),
|
||||
SpawnAction::None
|
||||
);
|
||||
assert_eq!(form.selected_profile().selector, "builtin:coder");
|
||||
form.apply_key(KeyEvent::new(KeyCode::Down, KeyModifiers::NONE));
|
||||
assert_eq!(form.selected_profile().selector, "builtin:default");
|
||||
form.apply_key(KeyEvent::new(KeyCode::Up, KeyModifiers::NONE));
|
||||
assert_eq!(form.selected_profile().selector, "builtin:coder");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn name_input_uses_old_safe_character_policy() {
|
||||
let mut form = SpawnForm::new(Some("worker".to_owned()), "yoi".to_owned(), choices());
|
||||
form.apply_key(KeyEvent::new(KeyCode::Char('-'), KeyModifiers::NONE));
|
||||
form.apply_key(KeyEvent::new(KeyCode::Char('1'), KeyModifiers::NONE));
|
||||
form.apply_key(KeyEvent::new(KeyCode::Char('/'), KeyModifiers::NONE));
|
||||
assert_eq!(form.worker_name, "worker-1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn enter_rejects_empty_name_and_escape_cancels() {
|
||||
let mut form = SpawnForm::new(Some(String::new()), "yoi".to_owned(), choices());
|
||||
assert_eq!(
|
||||
form.apply_key(KeyEvent::new(KeyCode::Enter, KeyModifiers::NONE)),
|
||||
SpawnAction::None
|
||||
);
|
||||
assert_eq!(
|
||||
form.status.as_ref().map(|(message, _)| message.as_str()),
|
||||
Some("worker name cannot be empty")
|
||||
);
|
||||
assert_eq!(
|
||||
form.apply_key(KeyEvent::new(KeyCode::Esc, KeyModifiers::NONE)),
|
||||
SpawnAction::Cancel
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn renderer_preserves_legacy_inline_spawn_form() {
|
||||
let backend = ratatui::backend::TestBackend::new(100, VIEWPORT_HEIGHT);
|
||||
let mut terminal = ratatui::Terminal::new(backend).unwrap();
|
||||
let form = SpawnForm::new(None, "yoi".to_owned(), choices());
|
||||
terminal.draw(|frame| draw_form(frame, &form)).unwrap();
|
||||
let buffer = terminal.backend().buffer();
|
||||
let rendered = buffer
|
||||
.content
|
||||
.chunks(buffer.area.width as usize)
|
||||
.map(|row| row.iter().map(|cell| cell.symbol()).collect::<String>())
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n");
|
||||
|
||||
assert!(rendered.contains("spawn worker"));
|
||||
assert!(rendered.contains("name: yoi"));
|
||||
assert!(rendered.contains("profile: builtin:default (default) — Default"));
|
||||
assert!(rendered.contains("enter spawn · left/right edit · esc cancel"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builtin_discovery_produces_a_default_profile_choice() {
|
||||
let registry = ProfileDiscovery::with_sources(None, None)
|
||||
.discover()
|
||||
.unwrap();
|
||||
let choices = profile_choices(®istry);
|
||||
let default = choices.iter().find(|choice| choice.is_default).unwrap();
|
||||
assert_eq!(default.selector, "builtin:default");
|
||||
assert!(default.label.contains("(default)"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_worker_name_comes_from_sanitised_directory_basename() {
|
||||
assert_eq!(
|
||||
default_worker_name(Path::new("/home/hare/Project/yoi")),
|
||||
"yoi"
|
||||
);
|
||||
assert_eq!(
|
||||
default_worker_name(Path::new("/home/hare/Project/my project")),
|
||||
"my-project"
|
||||
);
|
||||
assert_eq!(default_worker_name(Path::new("/")), "worker");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn explicit_profile_bypasses_discovery_and_uses_directory_name() {
|
||||
let selection = select(
|
||||
Path::new("/home/hare/Project/yoi"),
|
||||
None,
|
||||
Some("builtin:coder".to_owned()),
|
||||
)
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(selection.worker_name, "yoi");
|
||||
assert_eq!(selection.profile, "builtin:coder");
|
||||
}
|
||||
}
|
||||
+153
-45
@@ -36,6 +36,9 @@ use crate::task::{TaskCounts, TaskEntry, TaskStatus, TaskStore};
|
||||
use crate::text_selection::{HistoryViewport, SelectionRow};
|
||||
use crate::view_mode::Mode;
|
||||
|
||||
const RUN_SPINNER_FRAMES: [&str; 8] = ["⣷", "⣯", "⣟", "⡿", "⢿", "⣻", "⣽", "⣾"];
|
||||
const RUN_SPINNER_FRAME_MS: u128 = 80;
|
||||
|
||||
pub fn draw(frame: &mut Frame, app: &mut App) {
|
||||
let area = frame.area();
|
||||
// Input content starts after the prompt (`> ` or `: `), so the width
|
||||
@@ -57,19 +60,27 @@ pub fn draw(frame: &mut Frame, app: &mut App) {
|
||||
let tabs = app.worker_view_tabs();
|
||||
let show_tabs = tabs.len() > 1;
|
||||
let mini_view_h = task_mini_view_height(&app.selected_worker_view().task_store, show_tabs);
|
||||
// One blank row separates the history tail from the mini-view so
|
||||
// the latest message doesn't visually crash into the task summary.
|
||||
// Folds away with the mini-view when there are no tasks.
|
||||
let mini_view_gap = if mini_view_h > 0 { 1 } else { 0 };
|
||||
let run_status_h = u16::from(app.running);
|
||||
let run_status_gap = run_status_h;
|
||||
// One blank row separates the history tail from the run/task mini-view so
|
||||
// the latest message doesn't visually crash into operational status.
|
||||
// Folds away when neither run status nor tasks are visible.
|
||||
let mini_view_gap = if mini_view_h > 0 || run_status_h > 0 {
|
||||
1
|
||||
} else {
|
||||
0
|
||||
};
|
||||
|
||||
let chunks = Layout::vertical([
|
||||
Constraint::Min(0), // history view
|
||||
Constraint::Length(mini_view_gap), // gap above mini-view
|
||||
Constraint::Length(mini_view_h), // task mini-view (0 when empty)
|
||||
Constraint::Length(1), // separator
|
||||
Constraint::Length(1), // status
|
||||
Constraint::Length(input_height), // input area
|
||||
Constraint::Length(1), // actionbar
|
||||
Constraint::Min(0), // history view
|
||||
Constraint::Length(mini_view_gap), // gap above run/task mini-view
|
||||
Constraint::Length(run_status_h), // active run status
|
||||
Constraint::Length(run_status_gap), // gap below active run status
|
||||
Constraint::Length(mini_view_h), // task mini-view (0 when empty)
|
||||
Constraint::Length(1), // separator
|
||||
Constraint::Length(1), // status
|
||||
Constraint::Length(input_height), // input area
|
||||
Constraint::Length(1), // actionbar
|
||||
])
|
||||
.split(area);
|
||||
|
||||
@@ -82,24 +93,27 @@ pub fn draw(frame: &mut Frame, app: &mut App) {
|
||||
} else {
|
||||
draw_history(frame, app, chunks[0]);
|
||||
}
|
||||
if run_status_h > 0 {
|
||||
draw_run_status(frame, app, chunks[2]);
|
||||
}
|
||||
if mini_view_h > 0 {
|
||||
draw_task_mini_view(
|
||||
frame,
|
||||
&app.selected_worker_view().task_store,
|
||||
&tabs,
|
||||
chunks[2],
|
||||
chunks[4],
|
||||
);
|
||||
}
|
||||
draw_separator(frame, chunks[3]);
|
||||
draw_separator(frame, chunks[5]);
|
||||
// Status/composer/control surfaces remain parent-owned. View selection changes
|
||||
// only transcript/task presentation and never implies SubWorker control.
|
||||
draw_status(frame, app, chunks[4]);
|
||||
draw_input(frame, app, &input_render, chunks[5]);
|
||||
draw_actionbar(frame, app, chunks[6]);
|
||||
draw_status(frame, app, chunks[6]);
|
||||
draw_input(frame, app, &input_render, chunks[7]);
|
||||
draw_actionbar(frame, app, chunks[8]);
|
||||
if app.is_command_mode() {
|
||||
draw_command_popup(frame, app, chunks[5]);
|
||||
draw_command_popup(frame, app, chunks[7]);
|
||||
} else if let Some(state) = app.completion.as_ref().filter(|c| c.is_active()) {
|
||||
draw_completion_popup(frame, state, chunks[5]);
|
||||
draw_completion_popup(frame, state, chunks[7]);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -120,6 +134,65 @@ fn task_mini_view_height(store: &TaskStore, show_tabs: bool) -> u16 {
|
||||
(active_shown as u16).saturating_add(1)
|
||||
}
|
||||
|
||||
fn draw_run_status(frame: &mut Frame, app: &App, area: Rect) {
|
||||
frame.render_widget(Paragraph::new(run_status_line(app, Instant::now())), area);
|
||||
}
|
||||
|
||||
fn run_status_line(app: &App, now: Instant) -> Line<'static> {
|
||||
let elapsed = app
|
||||
.run_started_at
|
||||
.and_then(|started_at| now.checked_duration_since(started_at))
|
||||
.unwrap_or_default();
|
||||
let spinner_index =
|
||||
((elapsed.as_millis() / RUN_SPINNER_FRAME_MS) as usize) % RUN_SPINNER_FRAMES.len();
|
||||
let request_label = if app.run_requests == 1 {
|
||||
"1 req".to_owned()
|
||||
} else {
|
||||
format!("{} reqs", app.run_requests)
|
||||
};
|
||||
|
||||
Line::from(vec![
|
||||
Span::styled(
|
||||
RUN_SPINNER_FRAMES[spinner_index],
|
||||
Style::default()
|
||||
.fg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD),
|
||||
),
|
||||
Span::raw(" "),
|
||||
Span::styled(
|
||||
fmt_run_elapsed(elapsed.as_secs()),
|
||||
Style::default().fg(Color::Gray),
|
||||
),
|
||||
Span::styled(" ・ ", Style::default().fg(Color::DarkGray)),
|
||||
Span::styled(request_label, Style::default().fg(Color::Gray)),
|
||||
Span::styled(" | ", Style::default().fg(Color::DarkGray)),
|
||||
Span::styled("↑", Style::default().fg(Color::Green)),
|
||||
Span::styled(
|
||||
fmt_tokens(app.run_upload_tokens),
|
||||
Style::default().fg(Color::Green),
|
||||
),
|
||||
Span::styled("/", Style::default().fg(Color::DarkGray)),
|
||||
Span::styled("↓", Style::default().fg(Color::Yellow)),
|
||||
Span::styled(
|
||||
fmt_tokens(app.run_output_tokens),
|
||||
Style::default().fg(Color::Yellow),
|
||||
),
|
||||
])
|
||||
}
|
||||
|
||||
fn fmt_run_elapsed(secs: u64) -> String {
|
||||
let hours = secs / 3600;
|
||||
let minutes = (secs % 3600) / 60;
|
||||
let seconds = secs % 60;
|
||||
if hours > 0 {
|
||||
format!("{hours}h {minutes}m {seconds:02}s")
|
||||
} else if minutes > 0 {
|
||||
format!("{minutes}m {seconds:02}s")
|
||||
} else {
|
||||
format!("{seconds}s")
|
||||
}
|
||||
}
|
||||
|
||||
fn draw_task_mini_view(frame: &mut Frame, store: &TaskStore, tabs: &[WorkerViewTab], area: Rect) {
|
||||
if area.height == 0 || area.width == 0 {
|
||||
return;
|
||||
@@ -1223,6 +1296,28 @@ fn chip_span_for(seg: &Segment, fallback: Style) -> (Style, String) {
|
||||
Style::default().fg(Color::Magenta),
|
||||
format!("[Clipboard #{id} | {chars} chars, {line_count} lines]"),
|
||||
),
|
||||
Segment::PasteArtifact { artifact } => (
|
||||
Style::default().fg(Color::Magenta),
|
||||
format!(
|
||||
"[Paste artifact {} | {} chars, {} lines, {}, {}, created {} ms]",
|
||||
artifact.artifact_id,
|
||||
artifact.char_count,
|
||||
artifact.line_count,
|
||||
artifact.media_type.as_str(),
|
||||
artifact.availability.as_str(),
|
||||
artifact.created_at_ms
|
||||
),
|
||||
),
|
||||
Segment::UploadedFile { file } => (
|
||||
Style::default().fg(Color::Cyan),
|
||||
format!(
|
||||
"[Attached {} | {} bytes, {}, {}]",
|
||||
file.file_name,
|
||||
file.byte_len,
|
||||
file.media_type,
|
||||
file.availability.as_str()
|
||||
),
|
||||
),
|
||||
Segment::FileRef { path } => (Style::default().fg(Color::Cyan), format!("@{path}")),
|
||||
Segment::Flow { selector } => (
|
||||
Style::default().fg(Color::Yellow),
|
||||
@@ -1241,6 +1336,22 @@ fn segment_display_text(seg: &Segment) -> String {
|
||||
Segment::Paste {
|
||||
id, chars, lines, ..
|
||||
} => format!("[Clipboard #{id} | {chars} chars, {lines} lines]"),
|
||||
Segment::PasteArtifact { artifact } => format!(
|
||||
"[Paste artifact {} | {} chars, {} lines, {}, {}, created {} ms]",
|
||||
artifact.artifact_id,
|
||||
artifact.char_count,
|
||||
artifact.line_count,
|
||||
artifact.media_type.as_str(),
|
||||
artifact.availability.as_str(),
|
||||
artifact.created_at_ms
|
||||
),
|
||||
Segment::UploadedFile { file } => format!(
|
||||
"[Attached {} | {} bytes, {}, {}]",
|
||||
file.file_name,
|
||||
file.byte_len,
|
||||
file.media_type,
|
||||
file.availability.as_str()
|
||||
),
|
||||
Segment::FileRef { path } => format!("@{path}"),
|
||||
Segment::Flow { selector } => format!("[Flow: {selector}]"),
|
||||
Segment::Unknown => "[unknown segment]".to_owned(),
|
||||
@@ -1726,32 +1837,7 @@ fn draw_status(frame: &mut Frame, app: &App, area: Rect) {
|
||||
),
|
||||
];
|
||||
|
||||
if app.running {
|
||||
let status = if let Some(wait_event) = &app.latest_llm_wait_event {
|
||||
format!(
|
||||
"request: {} | ↑{}/↓{} | {wait_event}",
|
||||
app.run_requests,
|
||||
fmt_tokens(app.run_upload_tokens),
|
||||
fmt_tokens(app.run_output_tokens),
|
||||
)
|
||||
} else if let Some(tool) = &app.current_tool {
|
||||
format!(
|
||||
"request: {} | ↑{}/↓{} | tool: {tool}",
|
||||
app.run_requests,
|
||||
fmt_tokens(app.run_upload_tokens),
|
||||
fmt_tokens(app.run_output_tokens),
|
||||
)
|
||||
} else {
|
||||
format!(
|
||||
"request: {} | ↑{}/↓{}",
|
||||
app.run_requests,
|
||||
fmt_tokens(app.run_upload_tokens),
|
||||
fmt_tokens(app.run_output_tokens),
|
||||
)
|
||||
};
|
||||
spans.push(Span::raw(" | "));
|
||||
spans.push(Span::styled(status, Style::default().fg(Color::Yellow)));
|
||||
} else if app.paused {
|
||||
if app.paused {
|
||||
spans.push(Span::raw(" | "));
|
||||
spans.push(Span::styled(
|
||||
"paused",
|
||||
@@ -1763,7 +1849,7 @@ fn draw_status(frame: &mut Frame, app: &App, area: Rect) {
|
||||
" — Enter to resume, Ctrl-X to cancel, type to start new turn",
|
||||
Style::default().fg(Color::DarkGray),
|
||||
));
|
||||
} else {
|
||||
} else if !app.running {
|
||||
spans.push(Span::styled(" idle", Style::default().fg(Color::DarkGray)));
|
||||
}
|
||||
|
||||
@@ -2053,6 +2139,28 @@ mod tests {
|
||||
use protocol::WorkerStatus;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
#[test]
|
||||
fn run_status_line_matches_console_metrics_and_spinner_frame() {
|
||||
let now = Instant::now();
|
||||
let mut app = App::new("worker".into());
|
||||
app.run_started_at = now.checked_sub(Duration::from_millis(160));
|
||||
app.run_requests = 1;
|
||||
app.run_upload_tokens = 1_200;
|
||||
app.run_output_tokens = 45;
|
||||
|
||||
assert_eq!(
|
||||
line_text(&run_status_line(&app, now)),
|
||||
"⣟ 0s ・ 1 req | ↑1.2k/↓45"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn run_elapsed_uses_console_style_units() {
|
||||
assert_eq!(fmt_run_elapsed(9), "9s");
|
||||
assert_eq!(fmt_run_elapsed(65), "1m 05s");
|
||||
assert_eq!(fmt_run_elapsed(3_726), "1h 2m 06s");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn task_summary_right_aligns_worker_tabs_and_highlights_selection() {
|
||||
let tabs = vec![
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user