Compare commits
277
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
016dbd7cb1 | ||
|
|
7210d3c202 | ||
|
|
572204b49a | ||
|
|
86dd67a24c | ||
|
|
beeba1fdfc | ||
|
|
41b7b289d0 | ||
|
|
24237249d1 | ||
|
|
e448073b98 | ||
|
|
c08152d069 | ||
|
|
3995561220 | ||
|
|
aea51caeb4 | ||
|
|
3eca380bd8 | ||
|
|
8a3e06bc81 | ||
|
|
c4274c42cb | ||
|
|
a61ad15767 | ||
|
|
7f1e374fd7 | ||
|
|
e3f5445a02 | ||
|
|
d2cb50d081 | ||
|
|
6c609808c9 | ||
|
|
2d4c7b383a | ||
|
|
c21ed7dff2 | ||
|
|
448e392a0e | ||
|
|
d97c40d6af | ||
|
|
2d512b6be6 | ||
|
|
f061a95b48 | ||
|
|
eefdef1bef | ||
|
|
7675f81999 | ||
|
|
8fb592071f | ||
|
|
2528312142 | ||
|
|
08d7965ea8 | ||
|
|
e0badad91f | ||
|
|
f0a91ce2d8 | ||
|
|
24cab83f48 | ||
|
|
33a2b5d702 | ||
|
|
7f807004ad | ||
|
|
5564425488 | ||
|
|
4a89c04732 | ||
|
|
ec5a403ec6 | ||
|
|
f6ce1df766 | ||
|
|
9d7ddcc04a | ||
|
|
3df611636b | ||
|
|
d0999326bd | ||
|
|
6fbc65476c | ||
|
|
fcc7d79d80 | ||
|
|
18fd6a1f5e | ||
|
|
a072562034 | ||
|
|
2b4a2bc688 | ||
|
|
3344d9f8b2 | ||
|
|
fae36d220d | ||
|
|
7b6a84a550 | ||
|
|
f29c343879 | ||
|
|
f5e9f49a13 | ||
|
|
73a35599d2 | ||
|
|
5080d7860e | ||
|
|
7fb1d4056c | ||
|
|
243a081874 | ||
|
|
04924cf796 | ||
|
|
8f0917b8bc | ||
|
|
d4ad46127a | ||
|
|
fba5ecf54c | ||
|
|
e035df9e7b | ||
|
|
3baf0b6358 | ||
|
|
4de04e42b5 | ||
|
|
ebec98a14c | ||
|
|
f966470d33 | ||
|
|
c76ede2ab4 | ||
|
|
2fd043b634 | ||
|
|
fb6bbe9145 | ||
|
|
d1585d7483 | ||
|
|
13d853217d | ||
|
|
cc27d57e4a | ||
|
|
f5c5ea5a0b | ||
|
|
f1baea1705 | ||
|
|
31f7d39647 | ||
|
|
c5a834bfd2 | ||
|
|
1eef9b75ee | ||
|
|
ba9c885f52 | ||
|
|
88683a8d8f | ||
|
|
0c48c5dee3 | ||
|
|
101a0acb6b | ||
|
|
2d1956b653 | ||
|
|
b7bba8b53a | ||
|
|
7b25b767f8 | ||
|
|
4631b95144 | ||
|
|
38627c498b | ||
|
|
282a8d31b5 | ||
|
|
ab4fb4c1ee | ||
|
|
e3e9e83bc1 | ||
|
|
4269ebec04 | ||
|
|
e8b9adcde4 | ||
|
|
668a9062b3 | ||
|
|
5fd2ccf084 | ||
|
|
5686bbc9fd | ||
|
|
2cd57a32b2 | ||
|
|
89f4f99622 | ||
|
|
78d571ed14 | ||
|
|
e5332f4a7f | ||
|
|
3ed1545c3c | ||
|
|
9da20d15da | ||
|
|
052d60bd7d | ||
|
|
2456d6fda5 | ||
|
|
e7803d1aba | ||
|
|
ca5fddf89b | ||
|
|
88fad3893e | ||
|
|
82f9b0e48c | ||
|
|
ddb4c1454d | ||
|
|
7b1cf854f2 | ||
|
|
51c6d7f835 | ||
|
|
9d55ce0a87 | ||
|
|
0f8d61188a | ||
|
|
7363dffb9d | ||
|
|
8e4b7deaa4 | ||
|
|
ec845cbc25 | ||
|
|
e7079e223f | ||
|
|
dea5bd581d | ||
|
|
75b85b46d1 | ||
|
|
68f00bc948 | ||
|
|
5e9f7a7dc3 | ||
|
|
b038f022d3 | ||
|
|
cf7515fb35 | ||
|
|
bb4c1dfe4f | ||
|
|
72b56964c3 | ||
|
|
5b0a6691f8 | ||
|
|
1239c638a5 | ||
|
|
d2fa0787d8 | ||
|
|
a7056702e8 | ||
|
|
bb56283063 | ||
|
|
130ef1f0fe | ||
|
|
724205b1df | ||
|
|
69824ea45d | ||
|
|
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 |
@@ -1,21 +1,19 @@
|
||||
すでにシステムのドッグフーディングに成功しているが、一旦安定した旧バージョンで、ブラウザ/TUI Client/backend/runtimeの分離とチームスペースとしてのworkspaceを作るObjectiveを進めている。
|
||||
すでにシステムのドッグフーディングに成功しており、ブラウザ/TUI Client/backend/runtimeの分離とチームスペースとしてのworkspaceの実装を進めている。
|
||||
|
||||
## このシステムに置ける設計要旨
|
||||
|
||||
- プロンプトはすべて resources/promptsに集約している。管理効率の向上と同時に、ユーザーがオーバーライドする形式でもある。
|
||||
- プロンプトはすべて`resources/prompts`に集約している。管理効率の向上のためであると同時に、ユーザーがオーバーライドする形式でもある。
|
||||
- 変更量を最小にするために設計を歪めたり、設計問題に対して不必要な後方互換性を作らない。長期的なメンテナンスと型安全性を追求すること。
|
||||
|
||||
### LLM コンテキストの加工原則
|
||||
|
||||
LLM に投げる context への割り込みは、大きく2種類に分かれる。**前者は許されるが、後者は禁止**。
|
||||
LLM に投げる context はappend-onlyが基本であり、またその永続化形式からAPIコールの形式を純粋に再現可能である必要が有る。
|
||||
|
||||
Workerの状態から純粋に再現可能で、且つ揮発性の無い操作であることが望ましい。(pruning、tool result の content 切り詰め、prompt cache anchor の付与等)。
|
||||
原則として、コンテキストは積み重ねるものであり、一時的にメッセージを差し込むことや、過去のメッセージを改ざんすることはKVキャッシュのヒット率を下げる。
|
||||
一時的にメッセージを差し込む等の、揮発性の有るコンテキストの改変や、過去のメッセージを改ざんすることは基本的に禁止されている。
|
||||
これを行うと、 LLM はそのコンテキストに基づいて生成を行う一方、次以降のターンでhistoryに残らないため、「自分がなぜその発言/tool call をしたか」の根拠が消えるうえ、prompt cache のヒット率も低下させることになる。
|
||||
|
||||
**禁止**: ターンを跨ぐことができない情報に基づいて、history に記録せずに context だけにコンテンツを差し込むこと。これをやると LLM はそれに反応して生成を行う一方、次以降のターンでhistoryに残らないため、「自分がなぜその発言/tool call をしたか」の根拠が消えるうえ、prompt cache のヒット率も低下させることになる。
|
||||
|
||||
新しい input を context に乗せたいなら、必ず先に `worker.history` に append して commit すること。`history.json` への永続化はそこから自動的についてくる。Notify / WorkerEvent / typed `SystemItem` reminder はこの原則で扱う。
|
||||
また、キャッシュを破壊するタイミングは正確にコントロールされる必要があり、キャッシュ破壊とトークン消費のトレードオフに基づいて慎重に設計されるべきである。
|
||||
過去のコンテキストの圧縮は、キャッシュ破壊とトークン消費のトレードオフであり、必要であれば行っている。
|
||||
しかし、キャッシュを破壊するタイミングと頻度は正確にコントロールされる必要があり、実際のセッションデータの解析に基づいて慎重に設計されるべきである。
|
||||
|
||||
---
|
||||
|
||||
|
||||
Generated
+75
-656
File diff suppressed because it is too large
Load Diff
@@ -68,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" }
|
||||
@@ -126,6 +132,7 @@ tokio-tungstenite = "0.29"
|
||||
tower = "0.5"
|
||||
toml = "1.1"
|
||||
tracing = "0.1"
|
||||
tracing-subscriber = { version = "0.3", features = ["env-filter", "json"] }
|
||||
url = "2.5"
|
||||
uuid = "1.23"
|
||||
zeroize = "1"
|
||||
|
||||
@@ -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));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -18,7 +18,6 @@ tokio = { workspace = true, features = ["rt", "macros", "net", "io-util", "sync"
|
||||
tokio-tungstenite = { workspace = true }
|
||||
uuid = { workspace = true }
|
||||
workspace-api.workspace = true
|
||||
workdir = { workspace = true }
|
||||
|
||||
[dev-dependencies]
|
||||
tempfile = { workspace = true }
|
||||
|
||||
@@ -192,6 +192,32 @@ impl BackendApiClient {
|
||||
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 {
|
||||
@@ -235,6 +261,18 @@ fn redirect_policy(origin: BackendOrigin) -> redirect::Policy {
|
||||
})
|
||||
}
|
||||
|
||||
#[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),
|
||||
@@ -266,6 +304,11 @@ pub enum BackendApiClientError {
|
||||
origin: BackendOrigin,
|
||||
status: u16,
|
||||
},
|
||||
BackendResponse {
|
||||
origin: BackendOrigin,
|
||||
status: u16,
|
||||
detail: Option<String>,
|
||||
},
|
||||
Io {
|
||||
path: PathBuf,
|
||||
source: std::io::Error,
|
||||
@@ -312,6 +355,17 @@ impl fmt::Display for BackendApiClientError {
|
||||
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())
|
||||
}
|
||||
@@ -584,6 +638,23 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[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 [
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
use crate::BackendOrigin;
|
||||
use serde::{Deserialize, Serialize};
|
||||
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,
|
||||
@@ -28,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),
|
||||
@@ -74,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>,
|
||||
@@ -91,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
|
||||
@@ -104,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,
|
||||
@@ -119,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 {
|
||||
@@ -162,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,19 +1,30 @@
|
||||
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_tungstenite::tungstenite::client::IntoClientRequest;
|
||||
use tokio_tungstenite::tungstenite::http::HeaderValue;
|
||||
use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION;
|
||||
pub use workdir::workspace::WorkingDirectorySummary as BackendWorkingDirectorySummary;
|
||||
pub use workspace_api::{
|
||||
Diagnostic as BackendDiagnostic, DiagnosticSeverity as BackendDiagnosticSeverity,
|
||||
ListResponse as BackendRuntimeListResponse, RuntimeSummary as BackendRuntimeSummary,
|
||||
BrowserCreateWorkerResponse as BackendCreateWorkerResponse,
|
||||
CreateWorkspaceWorkerRequest as BackendCreateWorkerRequest, Diagnostic as BackendDiagnostic,
|
||||
DiagnosticSeverity as BackendDiagnosticSeverity, ListResponse as BackendRuntimeListResponse,
|
||||
RuntimeSummary as BackendRuntimeSummary,
|
||||
WorkerCapabilitySummary as BackendWorkerCapabilitySummary,
|
||||
WorkerImplementationSummary as BackendWorkerImplementationSummary,
|
||||
WorkerLaunchOptionsResponse as BackendWorkerLaunchOptions,
|
||||
WorkerLaunchProfileCandidate as BackendWorkerLaunchProfileCandidate,
|
||||
WorkerLaunchRuntimeOption as BackendWorkerLaunchRuntimeOption,
|
||||
WorkerOperationState as BackendWorkerOperationState,
|
||||
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)]
|
||||
@@ -47,6 +58,164 @@ 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)]
|
||||
pub struct BackendWorkerLaunchTarget {
|
||||
pub base_url: String,
|
||||
pub workspace_id: Option<String>,
|
||||
}
|
||||
|
||||
impl BackendWorkerLaunchTarget {
|
||||
pub fn new(base_url: impl Into<String>, workspace_id: Option<String>) -> Self {
|
||||
Self {
|
||||
base_url: base_url.into(),
|
||||
workspace_id,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn select_workspace(&mut self, workspace_id: impl Into<String>) {
|
||||
self.workspace_id = Some(workspace_id.into());
|
||||
}
|
||||
|
||||
pub fn workspace_id(&self) -> Option<&str> {
|
||||
self.workspace_id.as_deref()
|
||||
}
|
||||
|
||||
pub fn runtime_target(
|
||||
&self,
|
||||
runtime_id: impl Into<String>,
|
||||
worker_id: impl Into<String>,
|
||||
) -> Result<BackendRuntimeTarget, BackendRuntimeClientError> {
|
||||
let workspace_id = self.workspace_id.clone().ok_or_else(|| {
|
||||
BackendRuntimeClientError::InvalidTarget(
|
||||
"workspace_id is required before creating a Backend worker".to_string(),
|
||||
)
|
||||
})?;
|
||||
Ok(BackendRuntimeTarget::new(
|
||||
self.base_url.clone(),
|
||||
workspace_id,
|
||||
runtime_id,
|
||||
worker_id,
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
@@ -133,6 +302,58 @@ impl From<reqwest::Error> for BackendRuntimeClientError {
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn get_backend_worker_launch_options(
|
||||
target: &BackendWorkerLaunchTarget,
|
||||
) -> Result<BackendWorkerLaunchOptions, BackendRuntimeClientError> {
|
||||
validate_launch_target(target)?;
|
||||
let api = BackendApiClient::from_stored_token(&target.base_url)?;
|
||||
get_backend_worker_launch_options_with_client(target, &api).await
|
||||
}
|
||||
|
||||
async fn get_backend_worker_launch_options_with_client(
|
||||
target: &BackendWorkerLaunchTarget,
|
||||
api: &BackendApiClient,
|
||||
) -> Result<BackendWorkerLaunchOptions, BackendRuntimeClientError> {
|
||||
let path = backend_workspace_workers_launch_options_path(
|
||||
target
|
||||
.workspace_id
|
||||
.as_deref()
|
||||
.expect("validated Backend Workspace scope"),
|
||||
);
|
||||
let response = api.request(HttpMethod::GET, &path)?.send().await?;
|
||||
let response = api.require_success(response).await?;
|
||||
Ok(response.json::<BackendWorkerLaunchOptions>().await?)
|
||||
}
|
||||
|
||||
pub async fn create_backend_worker(
|
||||
target: &BackendWorkerLaunchTarget,
|
||||
request: &BackendCreateWorkerRequest,
|
||||
) -> Result<BackendCreateWorkerResponse, BackendRuntimeClientError> {
|
||||
validate_launch_target(target)?;
|
||||
let api = BackendApiClient::from_stored_token(&target.base_url)?;
|
||||
create_backend_worker_with_client(target, request, &api).await
|
||||
}
|
||||
|
||||
async fn create_backend_worker_with_client(
|
||||
target: &BackendWorkerLaunchTarget,
|
||||
request: &BackendCreateWorkerRequest,
|
||||
api: &BackendApiClient,
|
||||
) -> Result<BackendCreateWorkerResponse, BackendRuntimeClientError> {
|
||||
let path = backend_workspace_workers_path(
|
||||
target
|
||||
.workspace_id
|
||||
.as_deref()
|
||||
.expect("validated Backend Workspace scope"),
|
||||
);
|
||||
let response = api
|
||||
.request(HttpMethod::POST, &path)?
|
||||
.json(request)
|
||||
.send()
|
||||
.await?;
|
||||
let response = api.require_success(response).await?;
|
||||
Ok(response.json::<BackendCreateWorkerResponse>().await?)
|
||||
}
|
||||
|
||||
pub async fn list_backend_workers(
|
||||
target: &BackendRuntimeListTarget,
|
||||
) -> Result<BackendRuntimeListResponse<BackendWorkerSummary>, BackendRuntimeClientError> {
|
||||
@@ -265,7 +486,7 @@ pub async fn restore_backend_worker(
|
||||
.json(&serde_json::json!({}))
|
||||
.send()
|
||||
.await?;
|
||||
api.check_status(response.status())?;
|
||||
let response = api.require_success(response).await?;
|
||||
Ok(response.json::<BackendWorkerRestoreResponse>().await?)
|
||||
}
|
||||
|
||||
@@ -340,6 +561,30 @@ fn validate_target(target: &BackendRuntimeTarget) -> Result<(), BackendRuntimeCl
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_launch_target(
|
||||
target: &BackendWorkerLaunchTarget,
|
||||
) -> Result<(), BackendRuntimeClientError> {
|
||||
if target.base_url.trim().is_empty() {
|
||||
return Err(BackendRuntimeClientError::InvalidTarget(
|
||||
"Backend API base URL is required".to_string(),
|
||||
));
|
||||
}
|
||||
if !(target.base_url.starts_with("http://") || target.base_url.starts_with("https://")) {
|
||||
return Err(BackendRuntimeClientError::InvalidTarget(
|
||||
"Backend API base URL must start with http:// or https://".to_string(),
|
||||
));
|
||||
}
|
||||
match target.workspace_id.as_deref() {
|
||||
Some("") => Err(BackendRuntimeClientError::InvalidTarget(
|
||||
"workspace_id must not be empty".to_string(),
|
||||
)),
|
||||
None => Err(BackendRuntimeClientError::InvalidTarget(
|
||||
"workspace selection is required before creating a Backend worker".to_string(),
|
||||
)),
|
||||
Some(_) => Ok(()),
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_list_target(
|
||||
target: &BackendRuntimeListTarget,
|
||||
) -> Result<(), BackendRuntimeClientError> {
|
||||
@@ -374,6 +619,17 @@ fn validate_list_target(
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn backend_workspace_workers_path(workspace_id: &str) -> String {
|
||||
format!("/api/w/{}/workers", path_segment_encode(workspace_id))
|
||||
}
|
||||
|
||||
fn backend_workspace_workers_launch_options_path(workspace_id: &str) -> String {
|
||||
format!(
|
||||
"{}/launch-options",
|
||||
backend_workspace_workers_path(workspace_id)
|
||||
)
|
||||
}
|
||||
|
||||
fn backend_runtimes_path(workspace_id: &str) -> String {
|
||||
format!("/api/w/{}/runtimes", path_segment_encode(workspace_id))
|
||||
}
|
||||
@@ -458,6 +714,155 @@ fn percent_encode(input: &str, keep: impl Fn(u8) -> bool) -> String {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::TcpListener;
|
||||
|
||||
async fn serve_json_once(body: serde_json::Value) -> (String, tokio::task::JoinHandle<String>) {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let base_url = format!("http://{}", listener.local_addr().unwrap());
|
||||
let task = tokio::spawn(async move {
|
||||
let (mut socket, _) = listener.accept().await.unwrap();
|
||||
let mut request = Vec::new();
|
||||
let header_end = loop {
|
||||
let mut buffer = [0_u8; 4096];
|
||||
let read = socket.read(&mut buffer).await.unwrap();
|
||||
assert!(read > 0, "client closed before sending HTTP headers");
|
||||
request.extend_from_slice(&buffer[..read]);
|
||||
if let Some(position) = request.windows(4).position(|part| part == b"\r\n\r\n") {
|
||||
break position + 4;
|
||||
}
|
||||
};
|
||||
let headers = String::from_utf8_lossy(&request[..header_end]);
|
||||
let content_length = headers
|
||||
.lines()
|
||||
.find_map(|line| {
|
||||
let (name, value) = line.split_once(':')?;
|
||||
name.eq_ignore_ascii_case("content-length")
|
||||
.then(|| value.trim().parse::<usize>().unwrap())
|
||||
})
|
||||
.unwrap_or(0);
|
||||
while request.len() < header_end + content_length {
|
||||
let mut buffer = [0_u8; 4096];
|
||||
let read = socket.read(&mut buffer).await.unwrap();
|
||||
assert!(read > 0, "client closed before sending HTTP body");
|
||||
request.extend_from_slice(&buffer[..read]);
|
||||
}
|
||||
|
||||
let body = serde_json::to_vec(&body).unwrap();
|
||||
let response = format!(
|
||||
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
|
||||
body.len()
|
||||
);
|
||||
socket.write_all(response.as_bytes()).await.unwrap();
|
||||
socket.write_all(&body).await.unwrap();
|
||||
String::from_utf8(request).unwrap()
|
||||
});
|
||||
(base_url, task)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn launch_options_request_uses_workspace_path_and_bearer_auth() {
|
||||
let (base_url, server) = serve_json_once(serde_json::json!({
|
||||
"workspace_id": "team main",
|
||||
"runtimes": [{
|
||||
"runtime_id": "embedded",
|
||||
"display_name": "Embedded",
|
||||
"built_in": true,
|
||||
"worker_creation_available": true,
|
||||
"working_directory_required": false,
|
||||
"status": "online",
|
||||
"diagnostics": []
|
||||
}],
|
||||
"default_profile": "builtin:default",
|
||||
"profiles": [{
|
||||
"id": "builtin:default",
|
||||
"label": "Default",
|
||||
"description": ""
|
||||
}],
|
||||
"repositories": [],
|
||||
"working_directories": [],
|
||||
"diagnostics": []
|
||||
}))
|
||||
.await;
|
||||
let target = BackendWorkerLaunchTarget::new(&base_url, Some("team main".to_string()));
|
||||
let api = BackendApiClient::from_access_token_for_test(&base_url, "launch-secret").unwrap();
|
||||
|
||||
let response = get_backend_worker_launch_options_with_client(&target, &api)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.runtimes[0].runtime_id, "embedded");
|
||||
let request = server.await.unwrap();
|
||||
assert!(request.starts_with("GET /api/w/team%20main/workers/launch-options HTTP/1.1\r\n"));
|
||||
assert!(
|
||||
request
|
||||
.to_ascii_lowercase()
|
||||
.contains("authorization: bearer launch-secret\r\n")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn create_worker_posts_frontend_contract_to_workspace_path() {
|
||||
let (base_url, server) = serve_json_once(serde_json::json!({
|
||||
"workspace_id": "workspace-1",
|
||||
"runtime_id": "embedded",
|
||||
"worker_id": "worker-1",
|
||||
"console_href": "/w/workspace-1/workers/embedded/worker-1",
|
||||
"worker": {
|
||||
"runtime_id": "embedded",
|
||||
"worker_id": "worker-1",
|
||||
"host_id": "host-1",
|
||||
"display_name": "Coder one",
|
||||
"label": "Coder one",
|
||||
"profile": "builtin:coder",
|
||||
"singleton_key": null,
|
||||
"tags": [],
|
||||
"workspace": {
|
||||
"visibility": "workspace",
|
||||
"identity": "workspace",
|
||||
"workspace_id": "workspace-1"
|
||||
},
|
||||
"state": "idle",
|
||||
"last_seen_at": null,
|
||||
"pinned": false,
|
||||
"retention_state": "resident",
|
||||
"implementation": {"kind": "embedded", "display_hint": "Embedded"},
|
||||
"capabilities": {"can_stop": true, "can_spawn_followup": false},
|
||||
"diagnostics": []
|
||||
},
|
||||
"diagnostics": []
|
||||
}))
|
||||
.await;
|
||||
let target = BackendWorkerLaunchTarget::new(&base_url, Some("workspace-1".to_string()));
|
||||
let api = BackendApiClient::from_access_token_for_test(&base_url, "create-secret").unwrap();
|
||||
let create = BackendCreateWorkerRequest {
|
||||
runtime_id: "embedded".to_string(),
|
||||
display_name: "Coder one".to_string(),
|
||||
profile: Some("builtin:coder".to_string()),
|
||||
ticket_assignment: None,
|
||||
initial_submit: Vec::new(),
|
||||
working_directory: None,
|
||||
control_operation_id: None,
|
||||
};
|
||||
|
||||
let response = create_backend_worker_with_client(&target, &create, &api)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.worker_id, "worker-1");
|
||||
let request = server.await.unwrap();
|
||||
assert!(request.starts_with("POST /api/w/workspace-1/workers HTTP/1.1\r\n"));
|
||||
assert!(
|
||||
request
|
||||
.to_ascii_lowercase()
|
||||
.contains("authorization: bearer create-secret\r\n")
|
||||
);
|
||||
let body = request.split_once("\r\n\r\n").unwrap().1;
|
||||
let body: serde_json::Value = serde_json::from_str(body).unwrap();
|
||||
assert_eq!(body["runtime_id"], "embedded");
|
||||
assert_eq!(body["display_name"], "Coder one");
|
||||
assert_eq!(body["profile"], "builtin:coder");
|
||||
assert_eq!(body["initial_submit"], serde_json::json!([]));
|
||||
assert_eq!(body["working_directory"], serde_json::Value::Null);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn protocol_url_uses_backend_runtime_worker_identity() {
|
||||
@@ -508,8 +913,8 @@ mod tests {
|
||||
"capabilities": {"can_stop": true, "can_spawn_followup": false},
|
||||
"working_directory": {
|
||||
"working_directory_id": "wd-1",
|
||||
"repository_id": "main",
|
||||
"materializer_kind": "local_git_worktree",
|
||||
"repository_key": "main",
|
||||
"materializer_kind": "runtime_git_clone",
|
||||
"status": "active",
|
||||
"occupied_by": {
|
||||
"runtime_id": "arcadia",
|
||||
@@ -521,13 +926,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,60 +1,18 @@
|
||||
use crate::{BackendApiClient, BackendApiClientError};
|
||||
use reqwest::Method;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::fmt;
|
||||
use workspace_api::{RepositoryObservedStatus, RepositorySource};
|
||||
use workspace_api::{
|
||||
InitialRepositoryIntent, RepositoryListResponse, RepositorySummary,
|
||||
WorkspaceCatalogListResponse, WorkspaceCreateRequest, WorkspaceCreateResponse,
|
||||
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,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, PartialEq, Eq)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct CreateBackendWorkspaceRequest {
|
||||
pub operation_key: String,
|
||||
pub display_name: String,
|
||||
pub repository: CreateBackendWorkspaceRepository,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct CreateBackendWorkspaceRepository {
|
||||
pub uri: String,
|
||||
pub display_name: Option<String>,
|
||||
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>,
|
||||
}
|
||||
pub type BackendWorkspace = WorkspaceSummary;
|
||||
pub type CreateBackendWorkspaceResponse = WorkspaceCreateResponse;
|
||||
pub type CreateBackendWorkspaceRequest = WorkspaceCreateRequest;
|
||||
pub type CreateBackendWorkspaceRepository = InitialRepositoryIntent;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct BackendWorkspaceCatalogTarget {
|
||||
@@ -100,6 +58,48 @@ impl From<reqwest::Error> for BackendWorkspaceClientError {
|
||||
}
|
||||
}
|
||||
|
||||
pub fn list_backend_workspaces_blocking(
|
||||
target: &BackendWorkspaceCatalogTarget,
|
||||
) -> Result<Vec<BackendWorkspace>, BackendWorkspaceClientError> {
|
||||
let client = BackendApiClient::from_stored_token(&target.base_url)?;
|
||||
let response = client
|
||||
.blocking_request(
|
||||
Method::GET,
|
||||
&format!("/api/workspaces?limit={DEFAULT_WORKSPACE_LIMIT}"),
|
||||
)?
|
||||
.send()?;
|
||||
client.check_status(response.status())?;
|
||||
Ok(response.json::<WorkspaceCatalogListResponse>()?.0)
|
||||
}
|
||||
|
||||
pub fn list_backend_workspace_repositories_blocking(
|
||||
target: &BackendWorkspaceCatalogTarget,
|
||||
workspace_id: &str,
|
||||
) -> Result<Vec<RepositorySummary>, BackendWorkspaceClientError> {
|
||||
if workspace_id.is_empty()
|
||||
|| workspace_id.len() > 200
|
||||
|| !workspace_id
|
||||
.bytes()
|
||||
.all(|byte| byte.is_ascii_alphanumeric() || byte == b'-')
|
||||
{
|
||||
return Err(BackendWorkspaceClientError::InvalidTarget(
|
||||
"Workspace id returned by Backend is invalid".to_string(),
|
||||
));
|
||||
}
|
||||
let client = BackendApiClient::from_stored_token(&target.base_url)?;
|
||||
let response = client
|
||||
.blocking_request(Method::GET, &format!("/api/w/{workspace_id}/repositories"))?
|
||||
.send()?;
|
||||
client.check_status(response.status())?;
|
||||
let response = response.json::<RepositoryListResponse>()?;
|
||||
if response.workspace_id != workspace_id {
|
||||
return Err(BackendWorkspaceClientError::InvalidTarget(
|
||||
"Repository catalog response does not match the requested Workspace".to_string(),
|
||||
));
|
||||
}
|
||||
Ok(response.items)
|
||||
}
|
||||
|
||||
pub async fn list_backend_workspaces(
|
||||
target: &BackendWorkspaceCatalogTarget,
|
||||
) -> Result<Vec<BackendWorkspace>, BackendWorkspaceClientError> {
|
||||
@@ -118,7 +118,7 @@ async fn list_backend_workspaces_with_client(
|
||||
.send()
|
||||
.await?;
|
||||
client.check_status(response.status())?;
|
||||
Ok(response.json::<Vec<BackendWorkspace>>().await?)
|
||||
Ok(response.json::<WorkspaceCatalogListResponse>().await?.0)
|
||||
}
|
||||
|
||||
pub async fn create_backend_workspace(
|
||||
@@ -176,8 +176,8 @@ mod tests {
|
||||
operation_key: "workspace-create-1".to_string(),
|
||||
display_name: "Alpha".to_string(),
|
||||
repository: CreateBackendWorkspaceRepository {
|
||||
repository_key: "main".to_string(),
|
||||
uri: "/srv/repos/alpha".to_string(),
|
||||
display_name: Some("Main".to_string()),
|
||||
default_ref: Some("develop".to_string()),
|
||||
},
|
||||
};
|
||||
|
||||
@@ -112,26 +112,27 @@ mod tests {
|
||||
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,
|
||||
encode_event(&Event::WorkerState {
|
||||
snapshot: WorkerStatus::Idle.into(),
|
||||
})
|
||||
.expect("encode event"),
|
||||
);
|
||||
let mut client = Client::new(socket);
|
||||
|
||||
client
|
||||
.send(&Method::run_text("hello"))
|
||||
.send(&Method::submit_text(
|
||||
protocol::new_submission_request_id(),
|
||||
"hello",
|
||||
))
|
||||
.await
|
||||
.expect("send method");
|
||||
assert!(matches!(
|
||||
decode_method(&client.socket.sent[0]),
|
||||
Ok(Method::Run { .. })
|
||||
Ok(Method::Submit { .. })
|
||||
));
|
||||
assert!(matches!(
|
||||
client.next_event().await,
|
||||
Ok(Some(Event::Status {
|
||||
status: WorkerStatus::Idle
|
||||
}))
|
||||
Ok(Some(Event::WorkerState { .. }))
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
+22
-11
@@ -21,23 +21,34 @@ pub use backend_auth::{
|
||||
poll_device_login, start_device_login, wait_for_device_login,
|
||||
};
|
||||
pub use backend_runtime::{
|
||||
BackendDiagnostic, BackendDiagnosticSeverity, BackendRuntimeClientError,
|
||||
BackendRuntimeListResponse, BackendRuntimeListTarget, BackendRuntimeSummary,
|
||||
BackendRuntimeTarget, BackendWorkerCapabilitySummary, BackendWorkerImplementationSummary,
|
||||
BackendWorkerRestoreResponse, BackendWorkerRestoreResult, BackendWorkerSummary,
|
||||
BackendWorkerWorkspaceSummary, BackendWorkingDirectorySummary, connect_backend_runtime,
|
||||
list_backend_stopped_workers, list_backend_workers, restore_backend_worker,
|
||||
BackendCreateWorkerRequest, BackendCreateWorkerResponse, BackendDiagnostic,
|
||||
BackendDiagnosticSeverity, BackendRuntimeClientError, BackendRuntimeListResponse,
|
||||
BackendRuntimeListTarget, BackendRuntimeSummary, BackendRuntimeTarget,
|
||||
BackendWorkerCapabilitySummary, BackendWorkerImplementationSummary, BackendWorkerLaunchOptions,
|
||||
BackendWorkerLaunchProfileCandidate, BackendWorkerLaunchRuntimeOption,
|
||||
BackendWorkerLaunchTarget, BackendWorkerOperationState, BackendWorkerRestoreResponse,
|
||||
BackendWorkerRestoreResult, BackendWorkerSummary, BackendWorkerWorkspaceSummary,
|
||||
BackendWorkingDirectorySummary, connect_backend_runtime, create_backend_worker,
|
||||
get_backend_worker_launch_options, 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,
|
||||
CreateBackendWorkspaceResponse, create_backend_workspace,
|
||||
list_backend_workspace_repositories_blocking, list_backend_workspaces,
|
||||
list_backend_workspaces_blocking,
|
||||
};
|
||||
pub use client::{Client, ClientError};
|
||||
pub use target::{
|
||||
BackendTarget, Dashboard, ResolvedTarget, StandaloneTarget, StandaloneWorkerListIntent,
|
||||
StandaloneWorkerResumeIntent, Target, TargetError, TargetKind, WorkerConnection,
|
||||
WorkerConnectionSelector, WorkerList, WorkerListRequest, WorkerSpawn,
|
||||
BackendTarget, BackendWorkerLaunch, Dashboard, ResolvedTarget, StandaloneTarget,
|
||||
StandaloneWorkerListIntent, StandaloneWorkerResumeIntent, Target, TargetError, TargetKind,
|
||||
WorkerConnection, WorkerConnectionSelector, WorkerList, WorkerListRequest, WorkerSpawn,
|
||||
};
|
||||
pub use workspace_api::{
|
||||
CompanionCancelRequest, CompanionLifecycleState, CompanionMessageDisposition,
|
||||
CompanionMessageRequest, CompanionMessageResponse, CompanionStatusResponse,
|
||||
CompanionTranscriptItem, CompanionTranscriptProjection, CompanionTranscriptRole,
|
||||
CompanionTransportSummary, ObjectiveDetail, ObjectiveSummary,
|
||||
};
|
||||
pub use workspace_api::{ObjectiveDetail, ObjectiveSummary};
|
||||
pub use workspace_product::BackendWorkspaceProductClient;
|
||||
|
||||
@@ -2,7 +2,7 @@ use std::{fmt, path::PathBuf};
|
||||
|
||||
use crate::{
|
||||
BackendApiClient, BackendApiClientError, BackendOrigin, BackendRuntimeListTarget,
|
||||
BackendRuntimeTarget,
|
||||
BackendRuntimeTarget, BackendWorkerLaunchTarget,
|
||||
};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
@@ -123,6 +123,11 @@ pub struct Dashboard {
|
||||
pub workspace_id: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct BackendWorkerLaunch {
|
||||
pub target: BackendWorkerLaunchTarget,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct WorkerList {
|
||||
pub backend_target: BackendRuntimeListTarget,
|
||||
@@ -199,6 +204,13 @@ pub trait Target: fmt::Debug + Send + Sync {
|
||||
Err(TargetError::unsupported("Worker dashboard", self.kind()))
|
||||
}
|
||||
|
||||
fn launch_backend_worker(&self) -> Result<BackendWorkerLaunch, TargetError> {
|
||||
Err(TargetError::unsupported(
|
||||
"Backend Worker launch",
|
||||
self.kind(),
|
||||
))
|
||||
}
|
||||
|
||||
fn list_workers(&self, _request: WorkerListRequest) -> Result<WorkerList, TargetError> {
|
||||
Err(TargetError::unsupported("Worker listing", self.kind()))
|
||||
}
|
||||
@@ -299,6 +311,15 @@ impl Target for BackendTarget {
|
||||
})
|
||||
}
|
||||
|
||||
fn launch_backend_worker(&self) -> Result<BackendWorkerLaunch, TargetError> {
|
||||
Ok(BackendWorkerLaunch {
|
||||
target: BackendWorkerLaunchTarget::new(
|
||||
self.base_url.clone(),
|
||||
self.workspace_id.clone(),
|
||||
),
|
||||
})
|
||||
}
|
||||
|
||||
fn list_workers(&self, request: WorkerListRequest) -> Result<WorkerList, TargetError> {
|
||||
Ok(WorkerList {
|
||||
backend_target: BackendRuntimeListTarget::new(
|
||||
|
||||
@@ -89,17 +89,20 @@ mod tests {
|
||||
let mut client = Client::new(socket);
|
||||
|
||||
client
|
||||
.send(&Method::run_text("hello"))
|
||||
.send(&Method::submit_text(
|
||||
protocol::new_submission_request_id(),
|
||||
"hello",
|
||||
))
|
||||
.await
|
||||
.expect("send method");
|
||||
assert!(matches!(
|
||||
peer.next().await.as_deref().map(decode_method),
|
||||
Some(Ok(Method::Run { .. }))
|
||||
Some(Ok(Method::Submit { .. }))
|
||||
));
|
||||
|
||||
peer.send(
|
||||
encode_event(&Event::Status {
|
||||
status: WorkerStatus::Idle,
|
||||
encode_event(&Event::WorkerState {
|
||||
snapshot: WorkerStatus::Idle.into(),
|
||||
})
|
||||
.expect("encode event"),
|
||||
)
|
||||
@@ -107,9 +110,7 @@ mod tests {
|
||||
.expect("send event");
|
||||
assert!(matches!(
|
||||
client.next_event().await,
|
||||
Ok(Some(Event::Status {
|
||||
status: WorkerStatus::Idle
|
||||
}))
|
||||
Ok(Some(Event::WorkerState { .. }))
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -113,8 +113,8 @@ mod tests {
|
||||
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,
|
||||
let event = encode_event(&Event::WorkerState {
|
||||
snapshot: WorkerStatus::Idle.into(),
|
||||
})
|
||||
.unwrap();
|
||||
stream.write_all(event.as_bytes()).await.unwrap();
|
||||
@@ -126,12 +126,7 @@ mod tests {
|
||||
.await
|
||||
.expect("client should receive event while alive")
|
||||
.expect("transport should succeed");
|
||||
assert!(matches!(
|
||||
event,
|
||||
Some(Event::Status {
|
||||
status: WorkerStatus::Idle
|
||||
})
|
||||
));
|
||||
assert!(matches!(event, Some(Event::WorkerState { .. })));
|
||||
server.await.unwrap();
|
||||
}
|
||||
|
||||
@@ -147,12 +142,18 @@ mod tests {
|
||||
|
||||
let mut client = Client::new(Socket::connect(&socket_path).await.unwrap());
|
||||
client
|
||||
.send(&Method::run_text("hello"))
|
||||
.send(&Method::submit_text(
|
||||
protocol::new_submission_request_id(),
|
||||
"hello",
|
||||
))
|
||||
.await
|
||||
.expect("send method");
|
||||
|
||||
let received = server.await.unwrap().expect("method message");
|
||||
assert!(matches!(decode_method(&received), Ok(Method::Run { .. })));
|
||||
assert!(matches!(
|
||||
decode_method(&received),
|
||||
Ok(Method::Submit { .. })
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
@@ -114,10 +114,10 @@ mod tests {
|
||||
assert!(matches!(
|
||||
message,
|
||||
Message::Text(ref text)
|
||||
if matches!(decode_method(text), Ok(Method::Run { .. }))
|
||||
if matches!(decode_method(text), Ok(Method::Submit { .. }))
|
||||
));
|
||||
let event = encode_event(&Event::Status {
|
||||
status: WorkerStatus::Idle,
|
||||
let event = encode_event(&Event::WorkerState {
|
||||
snapshot: WorkerStatus::Idle.into(),
|
||||
})
|
||||
.unwrap();
|
||||
socket.send(Message::Text(event.into())).await.unwrap();
|
||||
@@ -126,14 +126,15 @@ mod tests {
|
||||
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"))
|
||||
.send(&Method::submit_text(
|
||||
protocol::new_submission_request_id(),
|
||||
"hello",
|
||||
))
|
||||
.await
|
||||
.expect("send method");
|
||||
assert!(matches!(
|
||||
client.next_event().await,
|
||||
Ok(Some(Event::Status {
|
||||
status: WorkerStatus::Idle
|
||||
}))
|
||||
Ok(Some(Event::WorkerState { .. }))
|
||||
));
|
||||
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,39 +9,19 @@ 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, RevokeRuntimeTrustKeyRequest,
|
||||
RuntimeTrustKeyRevealResponse, TICKET_ORCHESTRATION_PLANS_QUERY_PATH,
|
||||
TICKET_RELATIONS_QUERY_PATH, WorkerLaunchOptionsResponse, WorkspaceRuntimeDetail,
|
||||
WorkspaceRuntimeResource,
|
||||
};
|
||||
|
||||
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.
|
||||
@@ -263,11 +243,57 @@ impl BackendWorkspaceProductClient {
|
||||
)
|
||||
}
|
||||
|
||||
pub fn list_runtimes(
|
||||
&self,
|
||||
) -> Result<ListResponse<WorkspaceRuntimeResource>, BackendWorkspaceClientError> {
|
||||
self.get_json("/runtimes")
|
||||
}
|
||||
|
||||
pub fn runtime_detail(
|
||||
&self,
|
||||
runtime_id: &str,
|
||||
) -> Result<WorkspaceRuntimeDetail, BackendWorkspaceClientError> {
|
||||
self.get_json(&format!("/runtimes/{}", encode_path_segment(runtime_id)))
|
||||
}
|
||||
|
||||
pub fn reveal_runtime_trust_key(
|
||||
&self,
|
||||
runtime_id: &str,
|
||||
) -> Result<RuntimeTrustKeyRevealResponse, BackendWorkspaceClientError> {
|
||||
self.get_json(&format!(
|
||||
"/runtimes/{}/trust-key",
|
||||
encode_path_segment(runtime_id)
|
||||
))
|
||||
}
|
||||
|
||||
pub fn revoke_runtime_trust_key(
|
||||
&self,
|
||||
runtime_id: &str,
|
||||
request: &RevokeRuntimeTrustKeyRequest,
|
||||
) -> Result<WorkspaceRuntimeDetail, BackendWorkspaceClientError> {
|
||||
self.send_json(
|
||||
Method::DELETE,
|
||||
&format!("/runtimes/{}/trust-key", encode_path_segment(runtime_id)),
|
||||
Some(request),
|
||||
)
|
||||
}
|
||||
|
||||
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()
|
||||
@@ -278,19 +304,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
|
||||
@@ -298,7 +324,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(
|
||||
@@ -690,6 +716,82 @@ mod tests {
|
||||
(format!("http://{address}"), receiver, handle)
|
||||
}
|
||||
|
||||
#[test]
|
||||
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_with_access_token(
|
||||
base_url,
|
||||
"workspace-a",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let response = client.memory_document().unwrap();
|
||||
|
||||
assert_eq!(response.record_source, "workspace-sqlite");
|
||||
assert!(
|
||||
request
|
||||
.recv()
|
||||
.unwrap()
|
||||
.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":[]}"#;
|
||||
@@ -792,11 +894,11 @@ 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_with_access_token(
|
||||
@@ -824,7 +926,7 @@ 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_with_access_token(
|
||||
base_url,
|
||||
|
||||
@@ -9,14 +9,21 @@ fn workspace_creation_request_preserves_operation_key_for_retry() {
|
||||
operation_key: "workspace-create-1".to_string(),
|
||||
display_name: "Alpha".to_string(),
|
||||
repository: CreateBackendWorkspaceRepository {
|
||||
repository_key: "main".to_string(),
|
||||
uri: "/srv/repos/alpha".to_string(),
|
||||
display_name: Some("Main".to_string()),
|
||||
default_ref: Some("develop".to_string()),
|
||||
},
|
||||
};
|
||||
|
||||
assert_eq!(request.clone(), request);
|
||||
assert_eq!(request.operation_key, "workspace-create-1");
|
||||
let json = serde_json::to_value(&request).unwrap();
|
||||
assert_eq!(json["operation_key"], "workspace-create-1");
|
||||
assert_eq!(json["repository"]["repository_key"], "main");
|
||||
assert_eq!(json["repository"]["uri"], "/srv/repos/alpha");
|
||||
assert!(json.get("operation_id").is_none());
|
||||
assert!(json["repository"].get("display_name").is_none());
|
||||
assert!(json["repository"].get("source").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -3,7 +3,7 @@ use std::path::{Path, PathBuf};
|
||||
use globset::Glob;
|
||||
use ignore::WalkBuilder;
|
||||
|
||||
use crate::{FsAccessPolicy, FsError, FsPath, GlobRequest, GlobResult, direct_symlink};
|
||||
use crate::{FsAccessPolicy, FsError, FsPath, GlobRequest, GlobResult, resolve_access_path};
|
||||
|
||||
/// Execute a bounded glob entirely inside the provider process.
|
||||
pub fn run_glob(
|
||||
@@ -15,26 +15,24 @@ pub fn run_glob(
|
||||
if !root.is_absolute() {
|
||||
return Err(FsError::RelativePath(root.to_path_buf()));
|
||||
}
|
||||
if !access.is_readable(base) {
|
||||
let base_resolved = resolve_access_path(base).map_err(|error| FsError::Io {
|
||||
path: PathBuf::from(request.path.as_str()),
|
||||
source: error,
|
||||
})?;
|
||||
if !access.is_readable_paths(base, &base_resolved) {
|
||||
return Err(FsError::OutOfScope(PathBuf::from(request.path.as_str())));
|
||||
}
|
||||
if let Some(info) = direct_symlink(base)
|
||||
&& info.target_exists
|
||||
&& info.resolved_path.is_dir()
|
||||
{
|
||||
return Err(FsError::SymlinkDirectoryNotTraversed {
|
||||
tool: "Glob",
|
||||
path: PathBuf::from(request.path.as_str()),
|
||||
target: PathBuf::from("<provider-internal target>"),
|
||||
});
|
||||
}
|
||||
let matcher = Glob::new(&request.pattern)
|
||||
.map_err(|error| FsError::InvalidGlob(error.to_string()))?
|
||||
.compile_matcher();
|
||||
let mut matches = Vec::new();
|
||||
for entry in WalkBuilder::new(base).hidden(false).build().flatten() {
|
||||
let mut walker = WalkBuilder::new(base);
|
||||
walker.hidden(false).follow_links(false);
|
||||
for entry in walker.build().flatten() {
|
||||
let path = entry.path();
|
||||
if !path.is_file() || !access.is_readable(path) {
|
||||
let readable = resolve_access_path(path)
|
||||
.is_ok_and(|resolved| access.is_readable_paths(path, &resolved));
|
||||
if !path.is_file() || !readable {
|
||||
continue;
|
||||
}
|
||||
let relative = path.strip_prefix(base).unwrap_or(path);
|
||||
|
||||
@@ -14,7 +14,7 @@ use std::path::{Path, PathBuf};
|
||||
use thiserror::Error;
|
||||
|
||||
pub use glob::run_glob;
|
||||
pub use local::{run_edit, run_list, run_read, run_stat, run_write};
|
||||
pub use local::{resolve_access_path, run_edit, run_list, run_read, run_stat, run_write};
|
||||
pub use operation::*;
|
||||
pub use search::run_grep;
|
||||
|
||||
@@ -22,6 +22,19 @@ pub use search::run_grep;
|
||||
pub trait FsAccessPolicy: Send + Sync {
|
||||
fn is_readable(&self, path: &Path) -> bool;
|
||||
fn is_writable(&self, path: &Path) -> bool;
|
||||
|
||||
/// Authorize both the Workdir-visible path and its provider-resolved
|
||||
/// target. Implementations that do not distinguish symbolic-link identity
|
||||
/// retain resolved-target semantics through the defaults.
|
||||
fn is_readable_paths(&self, logical: &Path, resolved: &Path) -> bool {
|
||||
let _ = logical;
|
||||
self.is_readable(resolved)
|
||||
}
|
||||
|
||||
fn is_writable_paths(&self, logical: &Path, resolved: &Path) -> bool {
|
||||
let _ = logical;
|
||||
self.is_writable(resolved)
|
||||
}
|
||||
}
|
||||
|
||||
/// First symlink encountered while resolving a provider path.
|
||||
@@ -477,13 +490,14 @@ mod tests {
|
||||
|
||||
#[cfg(unix)]
|
||||
#[test]
|
||||
fn grep_keeps_direct_symlink_directory_and_broken_path_guards() {
|
||||
fn grep_traverses_a_direct_symlink_directory_and_rejects_a_broken_path() {
|
||||
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-dir/nested.rs"), "needle nested\n").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();
|
||||
@@ -501,18 +515,35 @@ mod tests {
|
||||
assert_eq!(file_result.match_count, 1);
|
||||
assert!(file_result.output.starts_with("file-link.rs\n"));
|
||||
|
||||
let directory_error = run_grep(
|
||||
let directory_result = 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")
|
||||
));
|
||||
.unwrap();
|
||||
assert_eq!(directory_result.match_count, 1);
|
||||
assert!(
|
||||
directory_result
|
||||
.output
|
||||
.starts_with("directory-link/nested.rs\n")
|
||||
);
|
||||
|
||||
let glob_result = run_glob(
|
||||
&root,
|
||||
&root.join("directory-link"),
|
||||
GlobRequest {
|
||||
pattern: "**/*.rs".to_string(),
|
||||
path: FsPath::new("directory-link").unwrap(),
|
||||
limit: 10,
|
||||
},
|
||||
&readable,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
glob_result.paths,
|
||||
vec![FsPath::new("directory-link/nested.rs").unwrap()]
|
||||
);
|
||||
|
||||
let broken_error = run_grep(
|
||||
&root,
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
use std::ffi::OsString;
|
||||
use std::fs;
|
||||
use std::io::Write;
|
||||
use std::path::{Path, PathBuf};
|
||||
@@ -18,7 +19,8 @@ pub fn run_stat(
|
||||
) -> Result<StatResult, FsError> {
|
||||
let logical = request.path;
|
||||
let path = resolve(root, &logical)?;
|
||||
if !access.is_readable(&path) {
|
||||
let resolved = resolve_access_path(&path).map_err(|error| map_io(&logical, error))?;
|
||||
if !access.is_readable_paths(&path, &resolved) {
|
||||
return Err(FsError::OutOfScope(PathBuf::from(logical.as_str())));
|
||||
}
|
||||
let metadata = fs::symlink_metadata(&path).map_err(|error| map_io(&logical, error))?;
|
||||
@@ -45,7 +47,7 @@ pub fn run_read(
|
||||
) -> Result<ReadResult, FsError> {
|
||||
let logical = request.path;
|
||||
let path = resolve(root, &logical)?;
|
||||
let path = require_access(&path, &logical, access, false)?;
|
||||
let path = require_access(&path, &logical, access, false, false)?;
|
||||
let metadata = fs::metadata(&path).map_err(|error| map_io(&logical, error))?;
|
||||
if metadata.is_dir() {
|
||||
return Err(FsError::IsDirectory(PathBuf::from(logical.as_str())));
|
||||
@@ -99,7 +101,7 @@ pub fn run_write(
|
||||
let path = resolve(root, &logical)?;
|
||||
let created = !path.exists();
|
||||
if path.exists() {
|
||||
let target = require_access(&path, &logical, access, true)?;
|
||||
let target = require_access(&path, &logical, access, true, false)?;
|
||||
let metadata = fs::metadata(&target).map_err(|error| map_io(&logical, error))?;
|
||||
if metadata.is_dir() {
|
||||
return Err(FsError::IsDirectory(PathBuf::from(logical.as_str())));
|
||||
@@ -113,12 +115,8 @@ pub fn run_write(
|
||||
if request.expected_hash.is_some() {
|
||||
return Err(FsError::Conflict(logical.as_str().to_string()));
|
||||
}
|
||||
let parent = path.parent().ok_or_else(|| {
|
||||
FsError::InvalidArgument(format!("{} has no parent", logical.as_str()))
|
||||
})?;
|
||||
let parent_logical = logical_parent(&logical);
|
||||
require_access(parent, &parent_logical, access, true)?;
|
||||
atomic_write(&path, &request.content, &logical)?;
|
||||
let target = require_access(&path, &logical, access, true, true)?;
|
||||
atomic_write(&target, &request.content, &logical)?;
|
||||
}
|
||||
Ok(WriteResult {
|
||||
bytes_written: request.content.len(),
|
||||
@@ -133,7 +131,7 @@ pub fn run_edit(
|
||||
) -> Result<EditResult, FsError> {
|
||||
let logical = request.path;
|
||||
let path = resolve(root, &logical)?;
|
||||
let target = require_access(&path, &logical, access, true)?;
|
||||
let target = require_access(&path, &logical, access, true, false)?;
|
||||
let bytes = fs::read(&target).map_err(|error| map_io(&logical, error))?;
|
||||
let actual_hash = hash_bytes(&bytes);
|
||||
if actual_hash != request.expected_hash {
|
||||
@@ -173,7 +171,8 @@ pub fn run_list(
|
||||
) -> Result<ListResult, FsError> {
|
||||
let logical = request.path;
|
||||
let path = resolve(root, &logical)?;
|
||||
let path = require_access(&path, &logical, access, false)?;
|
||||
let logical_base = path.clone();
|
||||
let path = require_access(&path, &logical, access, false, true)?;
|
||||
let metadata = fs::metadata(&path).map_err(|error| map_io(&logical, error))?;
|
||||
if !metadata.is_dir() {
|
||||
return Err(FsError::NotDirectory(PathBuf::from(logical.as_str())));
|
||||
@@ -183,7 +182,15 @@ pub fn run_list(
|
||||
for entry in read_dir {
|
||||
let entry = entry.map_err(|error| map_io(&logical, error))?;
|
||||
let absolute = entry.path();
|
||||
if !access.is_readable(&absolute) {
|
||||
let relative_to_base = absolute.strip_prefix(&path).map_err(|_| {
|
||||
FsError::InvalidArgument("provider returned a path outside its list base".to_string())
|
||||
})?;
|
||||
let logical_absolute = logical_base.join(relative_to_base);
|
||||
let resolved = match resolve_access_path(&absolute) {
|
||||
Ok(resolved) => resolved,
|
||||
Err(_) => continue,
|
||||
};
|
||||
if !access.is_readable_paths(&logical_absolute, &resolved) {
|
||||
continue;
|
||||
}
|
||||
let link_metadata =
|
||||
@@ -203,7 +210,7 @@ pub fn run_list(
|
||||
} else {
|
||||
EntryKind::Other
|
||||
};
|
||||
let relative = absolute.strip_prefix(root).map_err(|_| {
|
||||
let relative = logical_absolute.strip_prefix(root).map_err(|_| {
|
||||
FsError::InvalidArgument("provider returned a path outside its root".to_string())
|
||||
})?;
|
||||
entries.push(ListEntry {
|
||||
@@ -247,19 +254,24 @@ fn require_access(
|
||||
logical: &FsPath,
|
||||
access: &dyn FsAccessPolicy,
|
||||
write: bool,
|
||||
allow_symlink_directory: bool,
|
||||
) -> Result<PathBuf, FsError> {
|
||||
if let Some(info) = direct_symlink(path) {
|
||||
if !info.target_exists {
|
||||
return Err(FsError::BrokenSymlink {
|
||||
path: PathBuf::from(logical.as_str()),
|
||||
link: PathBuf::from(logical.as_str()),
|
||||
target: PathBuf::from("<provider-internal target>"),
|
||||
});
|
||||
}
|
||||
let symlink = direct_symlink(path);
|
||||
if let Some(info) = symlink.as_ref()
|
||||
&& !info.target_exists
|
||||
{
|
||||
return Err(FsError::BrokenSymlink {
|
||||
path: PathBuf::from(logical.as_str()),
|
||||
link: PathBuf::from(logical.as_str()),
|
||||
target: PathBuf::from("<provider-internal target>"),
|
||||
});
|
||||
}
|
||||
let resolved = resolve_access_path(path).map_err(|error| map_io(logical, error))?;
|
||||
if let Some(info) = symlink {
|
||||
let allowed = if write {
|
||||
access.is_writable(&info.resolved_path)
|
||||
access.is_writable_paths(path, &resolved)
|
||||
} else {
|
||||
access.is_readable(&info.resolved_path)
|
||||
access.is_readable_paths(path, &resolved)
|
||||
};
|
||||
if !allowed {
|
||||
return Err(FsError::SymlinkOutOfScope {
|
||||
@@ -268,21 +280,21 @@ fn require_access(
|
||||
required_permission: if write { "write" } else { "read" },
|
||||
});
|
||||
}
|
||||
if write && info.resolved_path.is_dir() {
|
||||
if !allow_symlink_directory && info.resolved_path.is_dir() {
|
||||
return Err(FsError::SymlinkTargetIsDirectory {
|
||||
path: PathBuf::from(logical.as_str()),
|
||||
target: PathBuf::from("<provider-internal target>"),
|
||||
});
|
||||
}
|
||||
return Ok(info.resolved_path);
|
||||
return Ok(resolved);
|
||||
}
|
||||
let allowed = if write {
|
||||
access.is_writable(path)
|
||||
access.is_writable_paths(path, &resolved)
|
||||
} else {
|
||||
access.is_readable(path)
|
||||
access.is_readable_paths(path, &resolved)
|
||||
};
|
||||
if allowed {
|
||||
Ok(path.to_path_buf())
|
||||
Ok(resolved)
|
||||
} else if write {
|
||||
Err(FsError::ReadOnly(PathBuf::from(logical.as_str())))
|
||||
} else {
|
||||
@@ -290,12 +302,38 @@ fn require_access(
|
||||
}
|
||||
}
|
||||
|
||||
fn logical_parent(path: &FsPath) -> FsPath {
|
||||
let parent = Path::new(path.as_str())
|
||||
.parent()
|
||||
.unwrap_or_else(|| Path::new(""))
|
||||
.to_string_lossy();
|
||||
FsPath::new(parent).unwrap_or_else(|_| FsPath::root())
|
||||
/// Resolve every existing component of an absolute provider path while
|
||||
/// retaining a missing final tail for create operations. Dangling symlinks are
|
||||
/// rejected because no resolved authority identity can be established.
|
||||
pub fn resolve_access_path(path: &Path) -> std::io::Result<PathBuf> {
|
||||
let mut cursor = path;
|
||||
let mut missing = Vec::<OsString>::new();
|
||||
loop {
|
||||
match fs::canonicalize(cursor) {
|
||||
Ok(mut resolved) => {
|
||||
for component in missing.iter().rev() {
|
||||
resolved.push(component);
|
||||
}
|
||||
return Ok(resolved);
|
||||
}
|
||||
Err(error) if error.kind() == std::io::ErrorKind::NotFound => {
|
||||
if fs::symlink_metadata(cursor)
|
||||
.is_ok_and(|metadata| metadata.file_type().is_symlink())
|
||||
{
|
||||
return Err(error);
|
||||
}
|
||||
let name = cursor.file_name().ok_or(error)?;
|
||||
missing.push(name.to_os_string());
|
||||
cursor = cursor.parent().ok_or_else(|| {
|
||||
std::io::Error::new(
|
||||
std::io::ErrorKind::NotFound,
|
||||
"path has no existing ancestor",
|
||||
)
|
||||
})?;
|
||||
}
|
||||
Err(error) => return Err(error),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn atomic_write(path: &Path, content: &[u8], logical: &FsPath) -> Result<(), FsError> {
|
||||
|
||||
@@ -10,7 +10,9 @@ use ignore::WalkBuilder;
|
||||
use ignore::overrides::{Override, OverrideBuilder};
|
||||
use ignore::types::{Types, TypesBuilder};
|
||||
|
||||
use crate::{FsError, GrepOutputMode, GrepRequest, GrepResult, direct_symlink};
|
||||
use crate::{
|
||||
FsError, GrepOutputMode, GrepRequest, GrepResult, direct_symlink, resolve_access_path,
|
||||
};
|
||||
|
||||
struct ContentLine {
|
||||
path: PathBuf,
|
||||
@@ -220,14 +222,28 @@ pub fn run_grep(
|
||||
return Err(FsError::RelativePath(base));
|
||||
}
|
||||
let symlink = direct_symlink(&base);
|
||||
if !access.is_readable(&base) {
|
||||
if let Some(info) = symlink.as_ref()
|
||||
&& !info.target_exists
|
||||
{
|
||||
return Err(FsError::BrokenSymlink {
|
||||
path: base.clone(),
|
||||
link: info.link_path.clone(),
|
||||
target: info.resolved_path.clone(),
|
||||
});
|
||||
}
|
||||
let resolved_base = resolve_access_path(&base).map_err(|error| FsError::io(&base, error))?;
|
||||
if !access.is_readable_paths(&base, &resolved_base) {
|
||||
return Err(if let Some(info) = symlink.as_ref() {
|
||||
let link_parent_readable = info
|
||||
.link_path
|
||||
.parent()
|
||||
.map(|parent| access.is_readable(parent))
|
||||
.and_then(|parent| {
|
||||
resolve_access_path(parent)
|
||||
.ok()
|
||||
.map(|resolved| access.is_readable_paths(parent, &resolved))
|
||||
})
|
||||
.unwrap_or(false);
|
||||
if info.target_exists && link_parent_readable {
|
||||
if link_parent_readable {
|
||||
FsError::SymlinkOutOfScope {
|
||||
path: base.clone(),
|
||||
target: info.resolved_path.clone(),
|
||||
@@ -240,15 +256,6 @@ pub fn run_grep(
|
||||
FsError::OutOfScope(base.clone())
|
||||
});
|
||||
}
|
||||
if let Some(info) = symlink.as_ref() {
|
||||
if !info.target_exists {
|
||||
return Err(FsError::BrokenSymlink {
|
||||
path: base.clone(),
|
||||
link: info.link_path.clone(),
|
||||
target: info.target_path.clone(),
|
||||
});
|
||||
}
|
||||
}
|
||||
let base_meta = std::fs::metadata(&base).map_err(|e| match e.kind() {
|
||||
std::io::ErrorKind::NotFound => FsError::NotFound(base.clone()),
|
||||
_ => FsError::io(&base, e),
|
||||
@@ -259,16 +266,6 @@ pub fn run_grep(
|
||||
base.display()
|
||||
)));
|
||||
}
|
||||
if base_meta.is_dir()
|
||||
&& let Some(info) = symlink.as_ref()
|
||||
{
|
||||
return Err(FsError::SymlinkDirectoryNotTraversed {
|
||||
tool: "Grep",
|
||||
path: base.clone(),
|
||||
target: info.resolved_path.clone(),
|
||||
});
|
||||
}
|
||||
|
||||
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())?;
|
||||
@@ -331,7 +328,9 @@ pub fn run_grep(
|
||||
continue;
|
||||
}
|
||||
let path = entry.path();
|
||||
if !access.is_readable(path) {
|
||||
let readable = resolve_access_path(path)
|
||||
.is_ok_and(|resolved| access.is_readable_paths(path, &resolved));
|
||||
if !readable {
|
||||
continue;
|
||||
}
|
||||
if scan_path(
|
||||
|
||||
+187
-101
@@ -15,13 +15,13 @@ use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::defaults;
|
||||
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
|
||||
@@ -54,10 +54,6 @@ pub struct WorkerManifestConfig {
|
||||
/// disabled after cascade merge.
|
||||
#[serde(default)]
|
||||
pub feature: FeatureConfigPartial,
|
||||
/// Explicit plugin package enablement entries. Discovery/resolution is a
|
||||
/// separate step and does not run during config merge.
|
||||
#[serde(default)]
|
||||
pub plugins: PluginConfig,
|
||||
/// Explicit Model Context Protocol provider declarations. Config parsing
|
||||
/// never starts a local MCP subprocess.
|
||||
#[serde(default)]
|
||||
@@ -67,15 +63,13 @@ 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>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct FeatureConfigPartial {
|
||||
#[serde(default)]
|
||||
pub task: Option<FeatureFlagConfigPartial>,
|
||||
@@ -103,8 +97,6 @@ pub struct FeatureConfigPartial {
|
||||
pub merge_request: Option<MergeRequestFeatureConfigPartial>,
|
||||
#[serde(default)]
|
||||
pub orchestration: Option<FeatureFlagConfigPartial>,
|
||||
#[serde(default)]
|
||||
pub plugins: Option<FeatureFlagConfigPartial>,
|
||||
}
|
||||
|
||||
impl FeatureConfigPartial {
|
||||
@@ -147,7 +139,6 @@ impl FeatureConfigPartial {
|
||||
other.orchestration,
|
||||
FeatureFlagConfigPartial::merge,
|
||||
),
|
||||
plugins: merge_option(self.plugins, other.plugins, FeatureFlagConfigPartial::merge),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -193,18 +184,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),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -253,13 +312,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(),
|
||||
@@ -296,10 +363,6 @@ impl From<FeatureConfigPartial> for FeatureConfig {
|
||||
.orchestration
|
||||
.map(FeatureFlagConfig::from)
|
||||
.unwrap_or_default(),
|
||||
plugins: value
|
||||
.plugins
|
||||
.map(FeatureFlagConfig::from)
|
||||
.unwrap_or_default(),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -329,20 +392,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),
|
||||
}),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -411,7 +506,6 @@ impl From<FeatureConfig> for FeatureConfigPartial {
|
||||
ticket: Some(value.ticket.into()),
|
||||
merge_request: Some(value.merge_request.into()),
|
||||
orchestration: Some(value.orchestration.into()),
|
||||
plugins: Some(value.plugins.into()),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -543,13 +637,23 @@ pub(crate) fn reject_removed_manifest_fields(s: &str) -> Result<(), toml::de::Er
|
||||
(removed; use compaction.prune_protected_tokens)",
|
||||
));
|
||||
}
|
||||
if value.get("memory").is_some() {
|
||||
return Err(toml::de::Error::custom(
|
||||
"unknown field in manifest: memory (removed; configure feature.memory)",
|
||||
));
|
||||
}
|
||||
if value.get("plugins").is_some() {
|
||||
return Err(toml::de::Error::custom(
|
||||
"unknown field in manifest: plugins (dynamic Plugins are not supported)",
|
||||
));
|
||||
}
|
||||
if value
|
||||
.get("memory")
|
||||
.get("feature")
|
||||
.and_then(toml::Value::as_table)
|
||||
.is_some_and(|table| table.contains_key("extract_worker_max_input_tokens"))
|
||||
.is_some_and(|table| table.contains_key("plugins"))
|
||||
{
|
||||
return Err(toml::de::Error::custom(
|
||||
"unknown field in manifest: memory.extract_worker_max_input_tokens (removed)",
|
||||
"unknown field in manifest: feature.plugins (dynamic Plugins are not supported)",
|
||||
));
|
||||
}
|
||||
if value
|
||||
@@ -633,11 +737,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
|
||||
{
|
||||
@@ -674,7 +773,6 @@ impl WorkerManifestConfig {
|
||||
PermissionConfigPartial::merge,
|
||||
),
|
||||
feature: self.feature.merge(upper.feature),
|
||||
plugins: merge_plugin_config(self.plugins, upper.plugins),
|
||||
mcp: merge_mcp_config(self.mcp, upper.mcp),
|
||||
compaction: merge_option(
|
||||
self.compaction,
|
||||
@@ -682,7 +780,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),
|
||||
}
|
||||
}
|
||||
@@ -695,16 +792,6 @@ impl SkillsConfig {
|
||||
}
|
||||
}
|
||||
|
||||
fn merge_plugin_config(mut base: PluginConfig, upper: PluginConfig) -> PluginConfig {
|
||||
let upper_has_resolved_plan = upper.has_resolved_plan();
|
||||
base.enabled.extend(upper.enabled);
|
||||
if upper_has_resolved_plan {
|
||||
base.resolved = upper.resolved;
|
||||
base.diagnostics = upper.diagnostics;
|
||||
}
|
||||
base
|
||||
}
|
||||
|
||||
fn merge_mcp_config(mut base: McpConfig, upper: McpConfig) -> McpConfig {
|
||||
base.stdio_servers.extend(upper.stdio_servers);
|
||||
base
|
||||
@@ -754,32 +841,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 {
|
||||
@@ -1219,11 +1280,9 @@ impl TryFrom<WorkerManifestConfig> for WorkerManifest {
|
||||
session,
|
||||
permissions,
|
||||
feature: FeatureConfig::from(cfg.feature),
|
||||
plugins: cfg.plugins,
|
||||
mcp: cfg.mcp,
|
||||
compaction,
|
||||
web: cfg.web,
|
||||
memory: cfg.memory,
|
||||
skills: cfg.skills,
|
||||
profile: None,
|
||||
})
|
||||
@@ -1260,18 +1319,17 @@ mod tests {
|
||||
target: abs("/worker"),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}],
|
||||
deny: Vec::new(),
|
||||
},
|
||||
delegation_scope: ScopeConfig::default(),
|
||||
permissions: None,
|
||||
feature: FeatureConfigPartial::default(),
|
||||
plugins: PluginConfig::default(),
|
||||
mcp: McpConfig::default(),
|
||||
session: None,
|
||||
compaction: None,
|
||||
web: None,
|
||||
memory: None,
|
||||
skills: None,
|
||||
}
|
||||
}
|
||||
@@ -1507,6 +1565,7 @@ mod tests {
|
||||
target: PathBuf::from("secrets"),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
});
|
||||
let resolved = cfg.resolve_paths(Path::new("/workspace/proj"));
|
||||
assert_eq!(resolved.scope.allow[0].target, Path::new("/workspace/proj"));
|
||||
@@ -1644,6 +1703,7 @@ mod tests {
|
||||
target: abs("/a"),
|
||||
permission: Permission::Read,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}],
|
||||
deny: Vec::new(),
|
||||
},
|
||||
@@ -1655,11 +1715,13 @@ mod tests {
|
||||
target: abs("/b"),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}],
|
||||
deny: vec![ScopeRule {
|
||||
target: abs("/a/secret"),
|
||||
permission: Permission::Read,
|
||||
recursive: false,
|
||||
symlink_policy: Default::default(),
|
||||
}],
|
||||
},
|
||||
..Default::default()
|
||||
@@ -1846,29 +1908,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]
|
||||
@@ -1948,7 +2031,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);
|
||||
@@ -2002,6 +2085,7 @@ enabled = false
|
||||
target: abs("/worker"),
|
||||
permission: Permission::Read,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}],
|
||||
deny: Vec::new(),
|
||||
},
|
||||
@@ -2025,8 +2109,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);
|
||||
}
|
||||
|
||||
@@ -2074,7 +2158,7 @@ readiness_check = true
|
||||
enabled = true
|
||||
|
||||
[feature.memory]
|
||||
staging = true
|
||||
staging_tools = true
|
||||
|
||||
[feature.manage_workdir]
|
||||
enabled = true
|
||||
@@ -2104,6 +2188,7 @@ enabled = true
|
||||
target: abs("/worker"),
|
||||
permission: Permission::Read,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}],
|
||||
deny: Vec::new(),
|
||||
},
|
||||
@@ -2111,8 +2196,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);
|
||||
@@ -2180,6 +2265,7 @@ permission = "write"
|
||||
target: abs("/worker"),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}],
|
||||
deny: Vec::new(),
|
||||
},
|
||||
|
||||
@@ -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);
|
||||
|
||||
+666
-177
@@ -29,7 +29,7 @@ pub use profile::{
|
||||
WorkspaceAuthorityRequirement, resolve_profile_artifact, resolve_profile_artifact_value,
|
||||
validate_profile_execution_target,
|
||||
};
|
||||
pub use protocol::{Permission, ScopeRule};
|
||||
pub use protocol::{Permission, ScopeRule, SymlinkPolicy};
|
||||
pub use scope::{DelegationScope, Scope, ScopeError, SharedScope};
|
||||
|
||||
use std::collections::{BTreeMap, HashMap};
|
||||
@@ -47,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,
|
||||
@@ -68,10 +69,6 @@ pub struct WorkerManifest {
|
||||
/// resolve disabled so Profile authors choose the exposed built-in surfaces.
|
||||
#[serde(default)]
|
||||
pub feature: FeatureConfig,
|
||||
/// Explicit plugin package enablement. Discovery remains read-only; only
|
||||
/// source-qualified entries listed here may resolve to active plugin metadata.
|
||||
#[serde(default)]
|
||||
pub plugins: plugin::PluginConfig,
|
||||
/// Explicit external Model Context Protocol provider configuration. This
|
||||
/// is config data only: declaring a server never starts a subprocess or
|
||||
/// grants OS sandboxing. Runtime MCP lifecycle/registration is a separate
|
||||
@@ -80,11 +77,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`.
|
||||
@@ -109,12 +101,13 @@ 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)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct FeatureConfig {
|
||||
#[serde(default)]
|
||||
pub task: FeatureFlagConfig,
|
||||
#[serde(default)]
|
||||
pub memory: MemoryFeatureConfig,
|
||||
pub memory: ResolvedMemoryFeatureConfig,
|
||||
#[serde(default)]
|
||||
pub web: FeatureFlagConfig,
|
||||
#[serde(default)]
|
||||
@@ -139,15 +132,13 @@ pub struct FeatureConfig {
|
||||
pub merge_request: MergeRequestFeatureConfig,
|
||||
#[serde(default)]
|
||||
pub orchestration: FeatureFlagConfig,
|
||||
#[serde(default)]
|
||||
pub plugins: FeatureFlagConfig,
|
||||
}
|
||||
|
||||
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(),
|
||||
@@ -159,7 +150,6 @@ impl Default for FeatureConfig {
|
||||
ticket: TicketFeatureConfig::default(),
|
||||
merge_request: MergeRequestFeatureConfig::default(),
|
||||
orchestration: FeatureFlagConfig::disabled(),
|
||||
plugins: FeatureFlagConfig::disabled(),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -222,34 +212,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(())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -484,98 +579,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 {
|
||||
@@ -931,6 +934,10 @@ impl Default for CompactionConfig {
|
||||
}
|
||||
|
||||
impl WorkerManifest {
|
||||
pub fn requires_persisted_execution_snapshot(&self) -> bool {
|
||||
self.profile.is_some() || 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)?;
|
||||
@@ -941,6 +948,267 @@ 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 = 3;
|
||||
const PREVIOUS_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
|
||||
&& version != PREVIOUS_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 mut 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",
|
||||
)));
|
||||
}
|
||||
if version == PREVIOUS_RESOLVED_MANIFEST_SNAPSHOT_SCHEMA_VERSION {
|
||||
migrate_legacy_manifest_authority(&mut manifest)?;
|
||||
}
|
||||
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_manifest_authority(
|
||||
manifest: &mut serde_json::Value,
|
||||
) -> Result<(), serde_json::Error> {
|
||||
let root = manifest.as_object_mut().ok_or_else(|| {
|
||||
serde_json::Error::io(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
"resolved Worker manifest must be an object",
|
||||
))
|
||||
})?;
|
||||
root.remove("plugins");
|
||||
if let Some(feature) = root.get_mut("feature") {
|
||||
let feature = feature.as_object_mut().ok_or_else(|| {
|
||||
serde_json::Error::io(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
"resolved Worker manifest feature must be an object",
|
||||
))
|
||||
})?;
|
||||
feature.remove("plugins");
|
||||
feature.remove("ticket_orchestration");
|
||||
if let Some(workers) = feature.remove("workers") {
|
||||
feature
|
||||
.entry("sub_worker".to_string())
|
||||
.or_insert_with(|| workers.clone());
|
||||
feature.entry("worker".to_string()).or_insert(workers);
|
||||
}
|
||||
if let Some(ticket) = feature
|
||||
.get_mut("ticket")
|
||||
.and_then(serde_json::Value::as_object_mut)
|
||||
&& let Some(access) = ticket.remove("access")
|
||||
&& ticket
|
||||
.get("enabled")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
&& access.as_str() == Some("lifecycle")
|
||||
{
|
||||
ticket.insert("authoring".to_string(), serde_json::Value::Bool(true));
|
||||
ticket.insert("thread".to_string(), serde_json::Value::Bool(true));
|
||||
ticket.insert("workflow".to_string(), serde_json::Value::Bool(true));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
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 requested_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 !requested_enabled {
|
||||
workspace_settings = None;
|
||||
}
|
||||
// Legacy standalone manifests could enable process-local Memory without a
|
||||
// Workspace-owned settings snapshot. That authority no longer exists, so
|
||||
// migration safely disables Memory instead of treating the whole Worker
|
||||
// snapshot as corrupt.
|
||||
let enabled = requested_enabled && workspace_settings.is_some();
|
||||
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);
|
||||
migrate_legacy_manifest_authority(&mut snapshot)?;
|
||||
validate_persisted_worker_manifest(serde_json::from_value(snapshot)?)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
@@ -1101,33 +1369,61 @@ model_id = "claude-sonnet-4-20250514"
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_plugin_enablement_config() {
|
||||
fn dynamic_plugin_manifest_config_is_rejected() {
|
||||
let toml = format!(
|
||||
"{MINIMAL_REQUIRED}\n\
|
||||
[[plugins.enabled]]\n\
|
||||
id = \"project:example\"\n\
|
||||
version = \"0.1.0\"\n\
|
||||
digest = \"sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa\"\n\
|
||||
surfaces = [\"hook\"]\n\n\
|
||||
[plugins.enabled.config]\n\
|
||||
greeting = \"hello\"\n"
|
||||
id = \"project:example\"\n"
|
||||
);
|
||||
let manifest = WorkerManifest::from_toml(&toml).unwrap();
|
||||
assert_eq!(manifest.plugins.enabled.len(), 1);
|
||||
let enabled = &manifest.plugins.enabled[0];
|
||||
assert_eq!(enabled.id, "project:example");
|
||||
assert_eq!(
|
||||
enabled.version.as_ref().map(|version| version.0.as_str()),
|
||||
Some("0.1.0")
|
||||
let error = WorkerManifest::from_toml(&toml).unwrap_err();
|
||||
assert!(
|
||||
error
|
||||
.to_string()
|
||||
.contains("dynamic Plugins are not supported"),
|
||||
"unexpected error: {error}"
|
||||
);
|
||||
assert_eq!(enabled.surfaces, vec![plugin::PluginSurface::Hook]);
|
||||
assert_eq!(
|
||||
enabled
|
||||
.config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("greeting"))
|
||||
.and_then(|value| value.as_str()),
|
||||
Some("hello")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn persisted_manifest_with_dynamic_plugin_plan_is_rejected() {
|
||||
let base =
|
||||
serde_json::to_value(WorkerManifest::from_toml(MINIMAL_REQUIRED).unwrap()).unwrap();
|
||||
|
||||
let mut top_level = base.clone();
|
||||
top_level.as_object_mut().unwrap().insert(
|
||||
"plugins".to_string(),
|
||||
serde_json::json!({
|
||||
"resolved": [{
|
||||
"package_path": "/tmp/ambient.yoi-plugin"
|
||||
}]
|
||||
}),
|
||||
);
|
||||
let error = serde_json::from_value::<WorkerManifest>(top_level).unwrap_err();
|
||||
assert!(error.to_string().contains("unknown field `plugins`"));
|
||||
|
||||
let mut nested = base;
|
||||
nested
|
||||
.get_mut("feature")
|
||||
.unwrap()
|
||||
.as_object_mut()
|
||||
.unwrap()
|
||||
.insert(
|
||||
"plugins".to_string(),
|
||||
serde_json::json!({ "enabled": true }),
|
||||
);
|
||||
let error = serde_json::from_value::<WorkerManifest>(nested).unwrap_err();
|
||||
assert!(error.to_string().contains("unknown field `plugins`"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dynamic_plugin_feature_flag_is_rejected() {
|
||||
let toml = format!("{MINIMAL_REQUIRED}\n[feature.plugins]\nenabled = true\n");
|
||||
let error = WorkerManifest::from_toml(&toml).unwrap_err();
|
||||
assert!(
|
||||
error
|
||||
.to_string()
|
||||
.contains("dynamic Plugins are not supported"),
|
||||
"unexpected error: {error}"
|
||||
);
|
||||
}
|
||||
|
||||
@@ -1246,36 +1542,237 @@ 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"], 3);
|
||||
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_drops_removed_plugin_authority() {
|
||||
let manifest = WorkerManifest::from_toml(MINIMAL_REQUIRED).unwrap();
|
||||
let mut versioned = write_persisted_worker_manifest_snapshot(&manifest).unwrap();
|
||||
versioned["schema_version"] = serde_json::json!(2);
|
||||
versioned["manifest"]["feature"]["plugins"] = serde_json::json!({ "enabled": true });
|
||||
versioned["manifest"]["feature"]
|
||||
.as_object_mut()
|
||||
.unwrap()
|
||||
.remove("sub_worker");
|
||||
versioned["manifest"]["feature"]
|
||||
.as_object_mut()
|
||||
.unwrap()
|
||||
.remove("worker");
|
||||
versioned["manifest"]["feature"]["workers"] = serde_json::json!({ "enabled": true });
|
||||
versioned["manifest"]["feature"]["ticket"] =
|
||||
serde_json::json!({ "enabled": true, "access": "lifecycle" });
|
||||
versioned["manifest"]["feature"]["ticket_orchestration"] =
|
||||
serde_json::json!({ "enabled": false });
|
||||
versioned["manifest"]["plugins"] = serde_json::json!({
|
||||
"enabled": ["legacy-plugin"],
|
||||
"config": { "legacy-plugin": { "legacy": true } }
|
||||
});
|
||||
|
||||
let restored = read_persisted_worker_manifest_snapshot(versioned).unwrap();
|
||||
let current = write_persisted_worker_manifest_snapshot(&restored).unwrap();
|
||||
assert_eq!(current["schema_version"], 3);
|
||||
assert!(current["manifest"].get("plugins").is_none());
|
||||
assert!(current["manifest"]["feature"].get("plugins").is_none());
|
||||
assert!(current["manifest"]["feature"].get("workers").is_none());
|
||||
assert_eq!(
|
||||
current["manifest"]["feature"]["sub_worker"]["enabled"],
|
||||
true
|
||||
);
|
||||
assert_eq!(current["manifest"]["feature"]["worker"]["enabled"], true);
|
||||
assert_eq!(current["manifest"]["feature"]["ticket"]["authoring"], true);
|
||||
assert_eq!(current["manifest"]["feature"]["ticket"]["thread"], true);
|
||||
assert_eq!(current["manifest"]["feature"]["ticket"]["workflow"], true);
|
||||
|
||||
let mut legacy = serde_json::to_value(manifest).unwrap();
|
||||
legacy.as_object_mut().unwrap().remove("memory");
|
||||
legacy["feature"]["memory"] = serde_json::json!({
|
||||
"enabled": true,
|
||||
"staging": false
|
||||
});
|
||||
legacy["feature"]["plugins"] = serde_json::json!({ "enabled": false });
|
||||
legacy["plugins"] = serde_json::json!({ "enabled": [] });
|
||||
let legacy = read_persisted_worker_manifest_snapshot(legacy).unwrap();
|
||||
let current = write_persisted_worker_manifest_snapshot(&legacy).unwrap();
|
||||
assert_eq!(
|
||||
current["manifest"]["feature"]["memory"]["profile"]["enabled"],
|
||||
false
|
||||
);
|
||||
}
|
||||
|
||||
#[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": 4,
|
||||
"manifest": manifest,
|
||||
}))
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
|
||||
@@ -1291,14 +1788,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 =
|
||||
|
||||
+100
-1874
File diff suppressed because it is too large
Load Diff
@@ -18,11 +18,10 @@ 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";
|
||||
@@ -148,7 +147,6 @@ pub enum WorkspaceAuthorityRequirement {
|
||||
MergeRequest,
|
||||
Objective,
|
||||
Orchestration,
|
||||
Plugins,
|
||||
Ticket,
|
||||
Worker,
|
||||
}
|
||||
@@ -162,7 +160,6 @@ impl fmt::Display for WorkspaceAuthorityRequirement {
|
||||
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"),
|
||||
}
|
||||
@@ -185,7 +182,7 @@ pub fn validate_profile_execution_target(
|
||||
if feature.manage_workdir.enabled {
|
||||
requirements.insert(WorkspaceAuthorityRequirement::ManageWorkdir);
|
||||
}
|
||||
if feature.memory.enabled || feature.memory.staging {
|
||||
if feature.memory.profile.enabled || feature.memory.profile.staging_tools {
|
||||
requirements.insert(WorkspaceAuthorityRequirement::Memory);
|
||||
}
|
||||
if feature.merge_request.show
|
||||
@@ -202,9 +199,6 @@ pub fn validate_profile_execution_target(
|
||||
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
|
||||
@@ -638,11 +632,9 @@ fn resolve_profile_value(
|
||||
session: profile.session,
|
||||
permissions: profile.permissions,
|
||||
feature: profile.feature,
|
||||
plugins: profile.plugins,
|
||||
mcp: profile.mcp,
|
||||
compaction,
|
||||
web: profile.web,
|
||||
memory: profile.memory.map(Into::into),
|
||||
skills: profile.skills,
|
||||
};
|
||||
let config =
|
||||
@@ -663,51 +655,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 {
|
||||
@@ -730,16 +677,12 @@ struct ProfileConfig {
|
||||
#[serde(default)]
|
||||
feature: FeatureConfigPartial,
|
||||
#[serde(default)]
|
||||
plugins: PluginConfig,
|
||||
#[serde(default)]
|
||||
mcp: McpConfig,
|
||||
#[serde(default)]
|
||||
compaction: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
web: Option<WebConfig>,
|
||||
#[serde(default)]
|
||||
memory: Option<ProfileMemoryConfig>,
|
||||
#[serde(default)]
|
||||
skills: Option<SkillsConfig>,
|
||||
}
|
||||
|
||||
@@ -940,12 +883,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() {
|
||||
@@ -1024,6 +961,7 @@ fn profile_scope_intent_to_config(
|
||||
target: workspace_base.join(path),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
});
|
||||
}
|
||||
Ok(ScopeConfig {
|
||||
@@ -1031,6 +969,7 @@ fn profile_scope_intent_to_config(
|
||||
target: workspace_base.to_path_buf(),
|
||||
permission,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}],
|
||||
deny,
|
||||
})
|
||||
@@ -1299,7 +1238,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 {
|
||||
@@ -1319,6 +1260,51 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ambient_plugin_directories_do_not_affect_builtin_profile_resolution() {
|
||||
let tmp = TempDir::new().unwrap();
|
||||
let workspace = tmp.path().join("workspace/nested");
|
||||
std::fs::create_dir_all(&workspace).unwrap();
|
||||
for root in [tmp.path(), tmp.path().join("workspace").as_path()] {
|
||||
let package = root.join(".yoi/plugins/broken.yoi-plugin");
|
||||
std::fs::create_dir_all(package.parent().unwrap()).unwrap();
|
||||
std::fs::write(package, b"malformed ambient package").unwrap();
|
||||
}
|
||||
|
||||
let resolved = ProfileResolver::new()
|
||||
.with_workspace_base(&workspace)
|
||||
.resolve_for_target(
|
||||
&ProfileSelector::source_named(ProfileRegistrySource::Builtin, "default"),
|
||||
ProfileResolveOptions::with_worker_name("standalone-worker"),
|
||||
ProfileExecutionTarget::Standalone,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(resolved.manifest.worker.name, "standalone-worker");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn profile_rejects_dynamic_plugin_configuration() {
|
||||
let tmp = TempDir::new().unwrap();
|
||||
for body in [
|
||||
"[feature.plugins]\nenabled = true\n",
|
||||
"[[plugins.enabled]]\nid = \"explicit:example\"\n",
|
||||
] {
|
||||
let profile = write_profile(tmp.path(), "plugin.toml", body);
|
||||
let error = ProfileResolver::new()
|
||||
.with_workspace_base(tmp.path())
|
||||
.resolve(
|
||||
&ProfileSelector::path(profile),
|
||||
ProfileResolveOptions::with_worker_name("runtime-worker"),
|
||||
)
|
||||
.unwrap_err();
|
||||
assert!(
|
||||
error.to_string().contains("unknown field"),
|
||||
"unexpected error: {error}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builtin_default_resolves_as_a_standalone_local_capability_profile() {
|
||||
let tmp = TempDir::new().unwrap();
|
||||
@@ -1351,14 +1337,12 @@ mod tests {
|
||||
assert!(resolved.manifest.delegation_scope.allow.iter().any(|rule| {
|
||||
rule.permission == protocol::Permission::Write && rule.target == tmp.path()
|
||||
}));
|
||||
assert!(!resolved.manifest.feature.memory.enabled);
|
||||
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]
|
||||
@@ -1438,6 +1422,28 @@ mod tests {
|
||||
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]
|
||||
fn profile_resolution_requires_runtime_worker_name() {
|
||||
let tmp = TempDir::new().unwrap();
|
||||
@@ -1608,7 +1614,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);
|
||||
|
||||
+271
-69
@@ -3,16 +3,17 @@
|
||||
//! Built from [`crate::ScopeConfig`] via [`Scope::from_config`]. Every
|
||||
//! rule `target` must already be an absolute path — per-layer path
|
||||
//! resolution runs earlier, inside [`crate::WorkerManifestConfig::resolve_paths`].
|
||||
//! All rule `target` paths inside the [`Scope`] are canonicalised (where
|
||||
//! possible) so access checks are pure path comparisons.
|
||||
//! All rule targets retain both their lexically normalized logical identity and
|
||||
//! their provider-resolved identity. Allow rules select one identity explicitly;
|
||||
//! deny rules always inspect both so aliases cannot bypass a restriction.
|
||||
|
||||
use std::ffi::OsString;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::path::{Component, Path, PathBuf};
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use arc_swap::{ArcSwap, Guard};
|
||||
|
||||
use crate::{Permission, ScopeConfig, ScopeRule};
|
||||
use crate::{Permission, ScopeConfig, ScopeRule, SymlinkPolicy};
|
||||
|
||||
/// Parsed, pwd-resolved set of allow/deny rules for a Worker.
|
||||
///
|
||||
@@ -26,10 +27,13 @@ pub struct Scope {
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
struct ResolvedRule {
|
||||
/// Absolute, canonicalized-or-normalized target directory/file.
|
||||
target: PathBuf,
|
||||
/// Absolute, lexically normalized target as presented through the Workdir.
|
||||
logical_target: PathBuf,
|
||||
/// Absolute target after provider-side symbolic-link resolution.
|
||||
resolved_target: PathBuf,
|
||||
permission: Permission,
|
||||
recursive: bool,
|
||||
symlink_policy: SymlinkPolicy,
|
||||
}
|
||||
|
||||
/// Parsed filesystem authority this Worker may pass to spawned children.
|
||||
@@ -98,18 +102,46 @@ fn permission_denies_requested(denied: Permission, requested: Permission) -> boo
|
||||
|
||||
fn rule_covers(available: &ResolvedRule, requested: &ResolvedRule) -> bool {
|
||||
permission_covers(available.permission, requested.permission)
|
||||
&& rule_path_set_contains(available, requested)
|
||||
&& available.symlink_policy >= requested.symlink_policy
|
||||
&& rule_path_set_contains(
|
||||
available,
|
||||
requested,
|
||||
match available.symlink_policy {
|
||||
SymlinkPolicy::Resolved => RuleIdentity::Resolved,
|
||||
SymlinkPolicy::Logical => RuleIdentity::Logical,
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
fn denial_overlaps_requested(deny: &ResolvedRule, requested: &ResolvedRule) -> bool {
|
||||
permission_denies_requested(deny.permission, requested.permission)
|
||||
&& rule_path_sets_overlap(deny, requested)
|
||||
&& (rule_path_sets_overlap(deny, requested, RuleIdentity::Logical)
|
||||
|| rule_path_sets_overlap(deny, requested, RuleIdentity::Resolved))
|
||||
}
|
||||
|
||||
fn rule_path_set_contains(available: &ResolvedRule, requested: &ResolvedRule) -> bool {
|
||||
#[derive(Clone, Copy)]
|
||||
enum RuleIdentity {
|
||||
Logical,
|
||||
Resolved,
|
||||
}
|
||||
|
||||
fn rule_target(rule: &ResolvedRule, identity: RuleIdentity) -> &Path {
|
||||
match identity {
|
||||
RuleIdentity::Logical => &rule.logical_target,
|
||||
RuleIdentity::Resolved => &rule.resolved_target,
|
||||
}
|
||||
}
|
||||
|
||||
fn rule_path_set_contains(
|
||||
available: &ResolvedRule,
|
||||
requested: &ResolvedRule,
|
||||
identity: RuleIdentity,
|
||||
) -> bool {
|
||||
let available_target = rule_target(available, identity);
|
||||
let requested_target = rule_target(requested, identity);
|
||||
match (available.recursive, requested.recursive) {
|
||||
// A recursive grant contains every possible requested path below its target.
|
||||
(true, _) => requested.target.starts_with(&available.target),
|
||||
(true, _) => requested_target.starts_with(available_target),
|
||||
// A non-recursive grant contains only the target and its direct children;
|
||||
// a recursive request always includes descendants beyond that finite-depth
|
||||
// set.
|
||||
@@ -117,36 +149,42 @@ fn rule_path_set_contains(available: &ResolvedRule, requested: &ResolvedRule) ->
|
||||
// Two non-recursive rules have the same finite-depth set only when their
|
||||
// target is identical. A request rooted at a direct child would also grant
|
||||
// that child's children, which are grandchildren of `available.target`.
|
||||
(false, false) => requested.target == available.target,
|
||||
(false, false) => requested_target == available_target,
|
||||
}
|
||||
}
|
||||
|
||||
fn rule_path_sets_overlap(left: &ResolvedRule, right: &ResolvedRule) -> bool {
|
||||
fn rule_path_sets_overlap(
|
||||
left: &ResolvedRule,
|
||||
right: &ResolvedRule,
|
||||
identity: RuleIdentity,
|
||||
) -> bool {
|
||||
let left_target = rule_target(left, identity);
|
||||
let right_target = rule_target(right, identity);
|
||||
match (left.recursive, right.recursive) {
|
||||
(true, true) => {
|
||||
left.target.starts_with(&right.target) || right.target.starts_with(&left.target)
|
||||
left_target.starts_with(right_target) || right_target.starts_with(left_target)
|
||||
}
|
||||
(true, false) => recursive_and_non_recursive_sets_overlap(left, right),
|
||||
(false, true) => recursive_and_non_recursive_sets_overlap(right, left),
|
||||
(true, false) => recursive_and_non_recursive_sets_overlap(left_target, right_target),
|
||||
(false, true) => recursive_and_non_recursive_sets_overlap(right_target, left_target),
|
||||
(false, false) => {
|
||||
left.target == right.target
|
||||
|| direct_child(&left.target, &right.target)
|
||||
|| direct_child(&right.target, &left.target)
|
||||
left_target == right_target
|
||||
|| direct_child(left_target, right_target)
|
||||
|| direct_child(right_target, left_target)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn recursive_and_non_recursive_sets_overlap(
|
||||
recursive: &ResolvedRule,
|
||||
non_recursive: &ResolvedRule,
|
||||
recursive_target: &Path,
|
||||
non_recursive_target: &Path,
|
||||
) -> bool {
|
||||
// The non-recursive set is `{target} + direct children`. It overlaps a
|
||||
// recursive subtree when either the non-recursive target is inside that
|
||||
// subtree, or the recursive subtree begins at the non-recursive target or
|
||||
// one of its direct children.
|
||||
non_recursive.target.starts_with(&recursive.target)
|
||||
|| recursive.target == non_recursive.target
|
||||
|| direct_child(&recursive.target, &non_recursive.target)
|
||||
non_recursive_target.starts_with(recursive_target)
|
||||
|| recursive_target == non_recursive_target
|
||||
|| direct_child(recursive_target, non_recursive_target)
|
||||
}
|
||||
|
||||
fn direct_child(child: &Path, parent: &Path) -> bool {
|
||||
@@ -201,23 +239,35 @@ impl Scope {
|
||||
}
|
||||
|
||||
/// Convenience constructor for tests and simple setups: a single
|
||||
/// recursive `allow(Write)` rule rooted at `root`.
|
||||
/// recursive `allow(Write)` rule rooted at `root` with the default
|
||||
/// resolved-target symlink policy.
|
||||
pub fn writable(root: impl AsRef<Path>) -> std::io::Result<Self> {
|
||||
let root = root.as_ref().canonicalize()?;
|
||||
let root = normalize_path(root.as_ref()).ok_or_else(|| {
|
||||
std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidInput,
|
||||
"scope root must be an absolute path without root traversal",
|
||||
)
|
||||
})?;
|
||||
let resolved_root = resolve_path(&root)?;
|
||||
Ok(Self {
|
||||
allow: vec![ResolvedRule {
|
||||
target: root,
|
||||
logical_target: root,
|
||||
resolved_target: resolved_root,
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: SymlinkPolicy::Resolved,
|
||||
}],
|
||||
deny: Vec::new(),
|
||||
})
|
||||
}
|
||||
|
||||
/// Resolve one rule target with the same symlink and missing-tail semantics
|
||||
/// used by scope matching.
|
||||
/// Return one rule target in the identity selected by its symlink policy.
|
||||
pub fn resolved_target(rule: &ScopeRule) -> Result<PathBuf, ScopeError> {
|
||||
Ok(resolve_rule(rule)?.target)
|
||||
let rule = resolve_rule(rule)?;
|
||||
Ok(match rule.symlink_policy {
|
||||
SymlinkPolicy::Resolved => rule.resolved_target,
|
||||
SymlinkPolicy::Logical => rule.logical_target,
|
||||
})
|
||||
}
|
||||
|
||||
/// Return whether this effective scope fully contains a requested rule.
|
||||
@@ -244,10 +294,23 @@ impl Scope {
|
||||
/// Returns `None` when `path` is outside every allow rule, or when
|
||||
/// deny rules have knocked it below `Read`.
|
||||
pub fn permission_at(&self, path: &Path) -> Option<Permission> {
|
||||
let resolved = resolve_path(path)?;
|
||||
let logical = normalize_path(path)?;
|
||||
let resolved = resolve_path(&logical).ok()?;
|
||||
self.permission_at_paths(&logical, &resolved)
|
||||
}
|
||||
|
||||
/// Effective permission for a path whose logical and provider-resolved
|
||||
/// identities were obtained inside the filesystem provider boundary.
|
||||
pub fn permission_at_paths(&self, logical: &Path, resolved: &Path) -> Option<Permission> {
|
||||
let logical = normalize_path(logical)?;
|
||||
let resolved = normalize_path(resolved)?;
|
||||
let mut effective: Option<Permission> = None;
|
||||
for rule in &self.allow {
|
||||
if rule.matches(&resolved) {
|
||||
let candidate = match rule.symlink_policy {
|
||||
SymlinkPolicy::Resolved => &resolved,
|
||||
SymlinkPolicy::Logical => &logical,
|
||||
};
|
||||
if rule.matches(candidate, rule.symlink_policy) {
|
||||
effective = match effective {
|
||||
None => Some(rule.permission),
|
||||
Some(cur) => Some(cur.max(rule.permission)),
|
||||
@@ -256,11 +319,13 @@ impl Scope {
|
||||
}
|
||||
let mut effective = effective?;
|
||||
|
||||
// Deny: min(min_deny) dictates the cap. Effective level is capped
|
||||
// strictly below that value, so deny(read) wipes access entirely.
|
||||
// Deny rules always inspect both identities. This prevents a logical
|
||||
// alias or a second symlink to the same target from bypassing a deny.
|
||||
let mut min_deny: Option<Permission> = None;
|
||||
for rule in &self.deny {
|
||||
if rule.matches(&resolved) {
|
||||
if rule.matches(&logical, SymlinkPolicy::Logical)
|
||||
|| rule.matches(&resolved, SymlinkPolicy::Resolved)
|
||||
{
|
||||
min_deny = match min_deny {
|
||||
None => Some(rule.permission),
|
||||
Some(cur) => Some(cur.min(rule.permission)),
|
||||
@@ -293,7 +358,7 @@ impl Scope {
|
||||
/// rule, preserving declaration order. Does not account for deny
|
||||
/// rules, which only cap effective permission at query time.
|
||||
pub fn readable_paths(&self) -> impl Iterator<Item = &Path> {
|
||||
self.allow.iter().map(|r| r.target.as_path())
|
||||
self.allow.iter().map(|r| r.logical_target.as_path())
|
||||
}
|
||||
|
||||
/// Allow rules with their targets resolved to absolute paths.
|
||||
@@ -305,9 +370,10 @@ impl Scope {
|
||||
self.allow
|
||||
.iter()
|
||||
.map(|r| ScopeRule {
|
||||
target: r.target.clone(),
|
||||
target: r.logical_target.clone(),
|
||||
permission: r.permission,
|
||||
recursive: r.recursive,
|
||||
symlink_policy: r.symlink_policy,
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
@@ -322,9 +388,10 @@ impl Scope {
|
||||
self.deny
|
||||
.iter()
|
||||
.map(|r| ScopeRule {
|
||||
target: r.target.clone(),
|
||||
target: r.logical_target.clone(),
|
||||
permission: r.permission,
|
||||
recursive: r.recursive,
|
||||
symlink_policy: r.symlink_policy,
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
@@ -335,7 +402,7 @@ impl Scope {
|
||||
self.allow
|
||||
.iter()
|
||||
.filter(|r| r.permission == Permission::Write)
|
||||
.map(|r| r.target.as_path())
|
||||
.map(|r| r.logical_target.as_path())
|
||||
}
|
||||
|
||||
/// Build a new [`Scope`] equal to `self` with `extra_allow` appended
|
||||
@@ -412,7 +479,10 @@ impl Scope {
|
||||
pub fn summary(&self) -> String {
|
||||
fn push_rule(out: &mut String, rule: &ResolvedRule) {
|
||||
out.push_str(" - ");
|
||||
out.push_str(&rule.target.display().to_string());
|
||||
out.push_str(&rule.logical_target.display().to_string());
|
||||
if rule.symlink_policy == SymlinkPolicy::Logical {
|
||||
out.push_str(" [logical-symlinks]");
|
||||
}
|
||||
if !rule.recursive {
|
||||
out.push_str(" [non-recursive]");
|
||||
}
|
||||
@@ -510,11 +580,15 @@ impl SharedScope {
|
||||
}
|
||||
|
||||
impl ResolvedRule {
|
||||
fn matches(&self, path: &Path) -> bool {
|
||||
fn matches(&self, path: &Path, identity: SymlinkPolicy) -> bool {
|
||||
let target = match identity {
|
||||
SymlinkPolicy::Resolved => &self.resolved_target,
|
||||
SymlinkPolicy::Logical => &self.logical_target,
|
||||
};
|
||||
if self.recursive {
|
||||
path.starts_with(&self.target)
|
||||
path.starts_with(target)
|
||||
} else {
|
||||
path == self.target || path.parent() == Some(self.target.as_path())
|
||||
path == target || path.parent() == Some(target.as_path())
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -523,48 +597,84 @@ fn resolve_rule(rule: &ScopeRule) -> Result<ResolvedRule, ScopeError> {
|
||||
if !rule.target.is_absolute() {
|
||||
return Err(ScopeError::RelativeTarget(rule.target.clone()));
|
||||
}
|
||||
let target = resolve_path(&rule.target).ok_or_else(|| ScopeError::ResolveTarget {
|
||||
let logical_target = normalize_path(&rule.target).ok_or_else(|| ScopeError::ResolveTarget {
|
||||
path: rule.target.clone(),
|
||||
source: std::io::Error::new(std::io::ErrorKind::Other, "could not absolutize target"),
|
||||
})?;
|
||||
let resolved_target =
|
||||
resolve_path(&logical_target).map_err(|source| ScopeError::ResolveTarget {
|
||||
path: rule.target.clone(),
|
||||
source,
|
||||
})?;
|
||||
Ok(ResolvedRule {
|
||||
target,
|
||||
logical_target,
|
||||
resolved_target,
|
||||
permission: rule.permission,
|
||||
recursive: rule.recursive,
|
||||
symlink_policy: rule.symlink_policy,
|
||||
})
|
||||
}
|
||||
|
||||
/// Convert `path` to an absolute form suitable for prefix comparison.
|
||||
///
|
||||
/// Tries `canonicalize` on the full path first (resolves symlinks). If
|
||||
/// the path doesn't exist yet, climbs to the closest existing ancestor,
|
||||
/// canonicalizes it, then rejoins the missing tail. Returns `None` for
|
||||
/// relative inputs that have no existing ancestor to anchor against.
|
||||
fn resolve_path(path: &Path) -> Option<PathBuf> {
|
||||
/// Resolve every existing path component while retaining a missing final tail.
|
||||
/// A dangling symlink is rejected rather than treated as an ordinary missing
|
||||
/// component because its resolved authority cannot be established.
|
||||
fn resolve_path(path: &Path) -> std::io::Result<PathBuf> {
|
||||
let mut cursor = path;
|
||||
let mut missing = Vec::<OsString>::new();
|
||||
loop {
|
||||
match std::fs::canonicalize(cursor) {
|
||||
Ok(mut resolved) => {
|
||||
for component in missing.iter().rev() {
|
||||
resolved.push(component);
|
||||
}
|
||||
return normalize_path(&resolved).ok_or_else(|| {
|
||||
std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidInput,
|
||||
"resolved target is not an absolute normalized path",
|
||||
)
|
||||
});
|
||||
}
|
||||
Err(error) if error.kind() == std::io::ErrorKind::NotFound => {
|
||||
if std::fs::symlink_metadata(cursor)
|
||||
.is_ok_and(|metadata| metadata.file_type().is_symlink())
|
||||
{
|
||||
return Err(error);
|
||||
}
|
||||
let name = cursor.file_name().ok_or(error)?;
|
||||
missing.push(name.to_os_string());
|
||||
cursor = cursor.parent().ok_or_else(|| {
|
||||
std::io::Error::new(
|
||||
std::io::ErrorKind::NotFound,
|
||||
"scope target has no existing ancestor",
|
||||
)
|
||||
})?;
|
||||
}
|
||||
Err(error) => return Err(error),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Normalize an absolute path for lexical scope comparison without consulting
|
||||
/// filesystem metadata or resolving symbolic links.
|
||||
fn normalize_path(path: &Path) -> Option<PathBuf> {
|
||||
if !path.is_absolute() {
|
||||
return None;
|
||||
}
|
||||
if let Ok(canonical) = path.canonicalize() {
|
||||
return Some(canonical);
|
||||
}
|
||||
let mut tail: Vec<OsString> = Vec::new();
|
||||
let mut cur = path.to_path_buf();
|
||||
loop {
|
||||
if let Ok(canonical) = cur.canonicalize() {
|
||||
let mut out = canonical;
|
||||
for segment in tail.iter().rev() {
|
||||
out.push(segment);
|
||||
let mut normalized = PathBuf::new();
|
||||
for component in path.components() {
|
||||
match component {
|
||||
Component::Prefix(prefix) => normalized.push(prefix.as_os_str()),
|
||||
Component::RootDir => normalized.push(component.as_os_str()),
|
||||
Component::CurDir => {}
|
||||
Component::ParentDir => {
|
||||
if !normalized.pop() {
|
||||
return None;
|
||||
}
|
||||
}
|
||||
return Some(out);
|
||||
Component::Normal(part) => normalized.push(part),
|
||||
}
|
||||
let name = cur.file_name()?.to_os_string();
|
||||
tail.push(name);
|
||||
let parent = cur.parent()?.to_path_buf();
|
||||
if parent == cur {
|
||||
return None;
|
||||
}
|
||||
cur = parent;
|
||||
}
|
||||
normalized.is_absolute().then_some(normalized)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
@@ -577,6 +687,7 @@ mod tests {
|
||||
target: target.to_path_buf(),
|
||||
permission,
|
||||
recursive,
|
||||
symlink_policy: Default::default(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -691,6 +802,7 @@ mod tests {
|
||||
target: dir.path().to_path_buf(),
|
||||
permission: Permission::Write,
|
||||
recursive: false,
|
||||
symlink_policy: Default::default(),
|
||||
}],
|
||||
deny: Vec::new(),
|
||||
};
|
||||
@@ -790,6 +902,7 @@ mod tests {
|
||||
target: PathBuf::from("relative/path"),
|
||||
permission: Permission::Read,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}],
|
||||
deny: Vec::new(),
|
||||
};
|
||||
@@ -805,6 +918,88 @@ mod tests {
|
||||
assert!(!scope.is_readable(&traversal));
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
#[test]
|
||||
fn scope_defaults_to_resolved_symlink_authority_and_logical_is_explicit() {
|
||||
use std::os::unix::fs::symlink;
|
||||
|
||||
let dir = TempDir::new().unwrap();
|
||||
let outside = TempDir::new().unwrap();
|
||||
std::fs::write(outside.path().join("outside.txt"), "visible through link").unwrap();
|
||||
symlink(outside.path(), dir.path().join("external")).unwrap();
|
||||
|
||||
let resolved = Scope::writable(dir.path()).unwrap();
|
||||
assert!(!resolved.is_readable(&dir.path().join("external/outside.txt")));
|
||||
assert!(!resolved.is_writable(&dir.path().join("external/new.txt")));
|
||||
|
||||
let logical = Scope::from_config(&ScopeConfig {
|
||||
allow: vec![ScopeRule {
|
||||
target: dir.path().to_path_buf(),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: SymlinkPolicy::Logical,
|
||||
}],
|
||||
deny: Vec::new(),
|
||||
})
|
||||
.unwrap();
|
||||
assert!(logical.is_readable(&dir.path().join("external/outside.txt")));
|
||||
assert!(logical.is_writable(&dir.path().join("external/new.txt")));
|
||||
assert!(!logical.is_readable(&outside.path().join("outside.txt")));
|
||||
assert!(!logical.is_writable(&outside.path().join("new.txt")));
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
#[test]
|
||||
fn deny_rules_match_both_logical_alias_and_resolved_target() {
|
||||
use std::os::unix::fs::symlink;
|
||||
|
||||
let root = TempDir::new().unwrap();
|
||||
let secret = root.path().join("secret");
|
||||
std::fs::create_dir(&secret).unwrap();
|
||||
std::fs::write(secret.join("key"), "hidden").unwrap();
|
||||
symlink(&secret, root.path().join("alias")).unwrap();
|
||||
let scope = Scope::from_config(&ScopeConfig {
|
||||
allow: vec![ScopeRule {
|
||||
target: root.path().to_path_buf(),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: SymlinkPolicy::Logical,
|
||||
}],
|
||||
deny: vec![ScopeRule {
|
||||
target: secret,
|
||||
permission: Permission::Read,
|
||||
recursive: true,
|
||||
symlink_policy: SymlinkPolicy::Logical,
|
||||
}],
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
assert!(!scope.is_readable(&root.path().join("alias/key")));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn delegation_symlink_policy_is_monotonically_attenuated() {
|
||||
let root = TempDir::new().unwrap();
|
||||
let mut parent_rule = allow_rule(root.path(), Permission::Write);
|
||||
parent_rule.symlink_policy = SymlinkPolicy::Logical;
|
||||
let logical_parent = DelegationScope::from_config(&ScopeConfig {
|
||||
allow: vec![parent_rule],
|
||||
deny: Vec::new(),
|
||||
})
|
||||
.unwrap();
|
||||
let resolved_child = allow_rule(&root.path().join("child"), Permission::Read);
|
||||
assert!(logical_parent.allows_rule(&resolved_child).unwrap());
|
||||
|
||||
let resolved_parent = DelegationScope::from_config(&ScopeConfig {
|
||||
allow: vec![allow_rule(root.path(), Permission::Write)],
|
||||
deny: Vec::new(),
|
||||
})
|
||||
.unwrap();
|
||||
let mut logical_child = resolved_child;
|
||||
logical_child.symlink_policy = SymlinkPolicy::Logical;
|
||||
assert!(!resolved_parent.allows_rule(&logical_child).unwrap());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn summary_lists_readable_and_writable() {
|
||||
let dir = TempDir::new().unwrap();
|
||||
@@ -851,11 +1046,13 @@ mod tests {
|
||||
target: docs.clone(),
|
||||
permission: Permission::Read,
|
||||
recursive: false,
|
||||
symlink_policy: Default::default(),
|
||||
},
|
||||
ScopeRule {
|
||||
target: dir.path().to_path_buf(),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
},
|
||||
],
|
||||
deny: Vec::new(),
|
||||
@@ -914,6 +1111,7 @@ mod tests {
|
||||
target: extra.path().to_path_buf(),
|
||||
permission: Permission::Read,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}])
|
||||
.unwrap();
|
||||
assert!(extended.is_readable(&extra.path().join("x")));
|
||||
@@ -931,6 +1129,7 @@ mod tests {
|
||||
target: sub.clone(),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}])
|
||||
.unwrap();
|
||||
let f = sub.join("a.txt");
|
||||
@@ -950,6 +1149,7 @@ mod tests {
|
||||
target: sub.clone(),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
};
|
||||
let base = Scope::writable(dir.path())
|
||||
.unwrap()
|
||||
@@ -1003,6 +1203,7 @@ mod tests {
|
||||
target: sub.clone(),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}])
|
||||
})
|
||||
.unwrap();
|
||||
@@ -1021,6 +1222,7 @@ mod tests {
|
||||
target: extra.path().to_path_buf(),
|
||||
permission: Permission::Read,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}])
|
||||
})
|
||||
.unwrap();
|
||||
|
||||
@@ -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")]
|
||||
|
||||
@@ -23,6 +23,7 @@ fn deny_write(target: &Path) -> ScopeRule {
|
||||
target: target.to_path_buf(),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
+758
-104
File diff suppressed because it is too large
Load Diff
@@ -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)]
|
||||
@@ -557,6 +573,11 @@ pub struct SubscriptionWorker {
|
||||
pub resource_key: Option<String>,
|
||||
/// Producer-owned monotonic revision for this Worker subject.
|
||||
pub subject_revision: u64,
|
||||
/// Latest revisioned foreground state observed from the Worker. This remains
|
||||
/// absent until an authoritative Worker snapshot/event has been applied.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub worker_state: Option<crate::WorkerStateSnapshot>,
|
||||
/// Runtime catalog lifecycle compatibility projection; not foreground-state authority.
|
||||
pub state: SubscriptionWorkerState,
|
||||
#[serde(default)]
|
||||
pub has_running_internal_workers: bool,
|
||||
@@ -567,7 +588,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 +610,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 +629,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 +644,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 +699,7 @@ pub enum SubscriptionSnapshot {
|
||||
events: Vec<WorkerProtocolEvent>,
|
||||
},
|
||||
WorkspaceWorkdirs {
|
||||
workdirs: Vec<SubscriptionWorkdir>,
|
||||
workdirs: Vec<WorkspaceSubscriptionWorkdir>,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -693,7 +767,7 @@ pub enum SubscriptionEventPayload {
|
||||
event: WorkerProtocolEvent,
|
||||
},
|
||||
WorkdirUpserted {
|
||||
workdir: SubscriptionWorkdir,
|
||||
workdir: WorkspaceSubscriptionWorkdir,
|
||||
},
|
||||
WorkdirRemoved {
|
||||
working_directory_id: SubscriptionWorkdirId,
|
||||
@@ -805,16 +879,49 @@ mod tests {
|
||||
runtime_id: None,
|
||||
resource_key: None,
|
||||
subject_revision: 0,
|
||||
worker_state: None,
|
||||
state: SubscriptionWorkerState::Idle,
|
||||
has_running_internal_workers: false,
|
||||
workspace_id: Some("workspace-1".to_string()),
|
||||
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 +1115,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,17 +7,22 @@ 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, SessionContentPart, SessionEntryProvenance, SessionMessageRole,
|
||||
SessionSnapshot, SessionSnapshotEntry, SessionSnapshotEntryData, SessionToolAttachment,
|
||||
ToolResultDisposition, TurnResult, WorkerEvent, WorkerStatus,
|
||||
InvokeKind, MemoryWorkerEvent, Method, PasteArtifactAvailability, PasteArtifactMediaType,
|
||||
PasteArtifactRef, PendingSubmissionSummary, PendingSubmissionsSnapshot, Permission,
|
||||
RewindSummary, RewindTarget, RewindTargetId, RunResult, ScopeRule, Segment, SessionContentPart,
|
||||
SessionEntryProvenance, SessionMessageRole, SessionSnapshot, SessionSnapshotEntry,
|
||||
SessionSnapshotEntryData, SessionToolAttachment, SubmissionDisposition, SymlinkPolicy,
|
||||
ToolResultDisposition, TurnResult, UploadedFileAvailability, UploadedFileRef, WorkerBusyState,
|
||||
WorkerCommandAcknowledgement, WorkerCommandDisposition, WorkerCommandEnvelope,
|
||||
WorkerCommandKind, WorkerEvent, WorkerMaintenanceState, WorkerRunState, WorkerState,
|
||||
WorkerStateSnapshot, 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,
|
||||
},
|
||||
};
|
||||
|
||||
@@ -44,12 +49,22 @@ pub fn generated_protocol_types() -> String {
|
||||
push_decl::<AlertSource>(&cfg, &mut output);
|
||||
push_decl::<CompletionKind>(&cfg, &mut output);
|
||||
push_decl::<WorkerStatus>(&cfg, &mut output);
|
||||
push_decl::<WorkerCommandEnvelope>(&cfg, &mut output);
|
||||
push_decl::<WorkerCommandKind>(&cfg, &mut output);
|
||||
push_decl::<WorkerCommandDisposition>(&cfg, &mut output);
|
||||
push_decl::<WorkerCommandAcknowledgement>(&cfg, &mut output);
|
||||
push_decl::<WorkerRunState>(&cfg, &mut output);
|
||||
push_decl::<WorkerMaintenanceState>(&cfg, &mut output);
|
||||
push_decl::<WorkerBusyState>(&cfg, &mut output);
|
||||
push_decl::<WorkerState>(&cfg, &mut output);
|
||||
push_decl::<WorkerStateSnapshot>(&cfg, &mut output);
|
||||
push_decl::<TurnResult>(&cfg, &mut output);
|
||||
push_decl::<InvokeKind>(&cfg, &mut output);
|
||||
push_decl::<RunResult>(&cfg, &mut output);
|
||||
push_decl::<ToolResultDisposition>(&cfg, &mut output);
|
||||
push_decl::<ErrorCode>(&cfg, &mut output);
|
||||
push_decl::<Permission>(&cfg, &mut output);
|
||||
push_decl::<SymlinkPolicy>(&cfg, &mut output);
|
||||
push_decl::<InFlightToolCallState>(&cfg, &mut output);
|
||||
push_decl::<CommandStatus>(&cfg, &mut output);
|
||||
push_decl::<CommandStream>(&cfg, &mut output);
|
||||
@@ -58,6 +73,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);
|
||||
@@ -71,6 +88,9 @@ pub fn generated_protocol_types() -> String {
|
||||
push_decl::<SessionToolAttachment>(&cfg, &mut output);
|
||||
push_decl::<SessionSnapshotEntryData>(&cfg, &mut output);
|
||||
push_decl::<SessionSnapshotEntry>(&cfg, &mut output);
|
||||
push_decl::<PendingSubmissionSummary>(&cfg, &mut output);
|
||||
push_decl::<PendingSubmissionsSnapshot>(&cfg, &mut output);
|
||||
push_decl::<SubmissionDisposition>(&cfg, &mut output);
|
||||
push_decl::<SessionSnapshot>(&cfg, &mut output);
|
||||
push_decl::<InternalWorkerKind>(&cfg, &mut output);
|
||||
push_decl::<InternalWorkerRef>(&cfg, &mut output);
|
||||
@@ -78,6 +98,9 @@ pub fn generated_protocol_types() -> String {
|
||||
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);
|
||||
@@ -88,7 +111,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);
|
||||
@@ -132,6 +155,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,20 @@
|
||||
//! 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, finalize_uploaded_file_binding,
|
||||
list_uploaded_file_refs, pin_uploaded_file, read_uploaded_file, read_uploaded_file_by_id,
|
||||
reconcile_uploaded_file_pins, release_uploaded_file_pin, uploaded_file_has_pending_owner,
|
||||
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 +120,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 +405,231 @@ 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 pin_uploaded_file(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
reference: &UploadedFileRef,
|
||||
owner_id: &str,
|
||||
) -> Result<(), StoreError> {
|
||||
let _guard = self
|
||||
.append_lock
|
||||
.lock()
|
||||
.map_err(|_| std::io::Error::other("session store append lock was poisoned"))?;
|
||||
pin_uploaded_file(&self.paste_artifact_dir(session_id), reference, owner_id)
|
||||
}
|
||||
|
||||
fn release_uploaded_file_pin(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
artifact_id: &str,
|
||||
owner_id: &str,
|
||||
) -> Result<(), StoreError> {
|
||||
let _guard = self
|
||||
.append_lock
|
||||
.lock()
|
||||
.map_err(|_| std::io::Error::other("session store append lock was poisoned"))?;
|
||||
release_uploaded_file_pin(&self.paste_artifact_dir(session_id), artifact_id, owner_id)
|
||||
}
|
||||
|
||||
fn finalize_uploaded_file_binding(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
artifact_id: &str,
|
||||
source_entry_id: &str,
|
||||
) -> Result<(), StoreError> {
|
||||
let _guard = self
|
||||
.append_lock
|
||||
.lock()
|
||||
.map_err(|_| std::io::Error::other("session store append lock was poisoned"))?;
|
||||
finalize_uploaded_file_binding(
|
||||
&self.paste_artifact_dir(session_id),
|
||||
artifact_id,
|
||||
source_entry_id,
|
||||
)
|
||||
}
|
||||
|
||||
fn reconcile_uploaded_file_pins(
|
||||
&self,
|
||||
session_id: SessionId,
|
||||
live_owner_ids: &[String],
|
||||
) -> Result<u64, StoreError> {
|
||||
let _guard = self
|
||||
.append_lock
|
||||
.lock()
|
||||
.map_err(|_| std::io::Error::other("session store append lock was poisoned"))?;
|
||||
reconcile_uploaded_file_pins(&self.paste_artifact_dir(session_id), live_owner_ids)
|
||||
}
|
||||
|
||||
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)? {
|
||||
finalize_uploaded_file_binding(&dir, &reference.artifact_id, source_entry_id)?;
|
||||
continue;
|
||||
}
|
||||
if uploaded_file_has_pending_owner(&dir, &reference.artifact_id)? {
|
||||
continue;
|
||||
}
|
||||
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 +678,524 @@ 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 pending_upload_pin_survives_cleanup_until_release_or_history_binding() {
|
||||
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: 64,
|
||||
max_session_bytes: 128,
|
||||
};
|
||||
let pending = store
|
||||
.write_uploaded_file(session_id, "pending.txt", "text/plain", b"pending", limits)
|
||||
.unwrap();
|
||||
store
|
||||
.pin_uploaded_file(session_id, &pending, "submission-1")
|
||||
.unwrap();
|
||||
assert!(matches!(
|
||||
store.pin_uploaded_file(session_id, &pending, "submission-other"),
|
||||
Err(StoreError::ArtifactAlreadyCommitted)
|
||||
));
|
||||
drop(store);
|
||||
let store = FsStore::new(tmp.path()).unwrap();
|
||||
assert_eq!(
|
||||
store.delete_uncommitted_uploaded_files(session_id).unwrap(),
|
||||
0
|
||||
);
|
||||
assert_eq!(
|
||||
store
|
||||
.read_uploaded_file_by_id(session_id, &pending.artifact_id)
|
||||
.unwrap()
|
||||
.1,
|
||||
b"pending"
|
||||
);
|
||||
|
||||
let fork_session_id = new_session_id();
|
||||
assert_eq!(
|
||||
store
|
||||
.copy_committed_uploaded_files(session_id, fork_session_id)
|
||||
.unwrap(),
|
||||
0
|
||||
);
|
||||
assert!(
|
||||
store
|
||||
.read_uploaded_file_by_id(fork_session_id, &pending.artifact_id)
|
||||
.is_err()
|
||||
);
|
||||
|
||||
let committed = store
|
||||
.bind_uploaded_file(session_id, &pending, "entry-1")
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
store.delete_uncommitted_uploaded_files(session_id).unwrap(),
|
||||
0
|
||||
);
|
||||
assert!(
|
||||
store
|
||||
.read_uploaded_file_by_id(session_id, &pending.artifact_id)
|
||||
.is_ok()
|
||||
);
|
||||
store
|
||||
.create_segment(
|
||||
session_id,
|
||||
new_segment_id(),
|
||||
&[LogEntry::InputSegmentsCheckpoint {
|
||||
ts: 1,
|
||||
user_segments: vec![vec![protocol::Segment::UploadedFile {
|
||||
file: committed.clone(),
|
||||
}]],
|
||||
}],
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
store.delete_uncommitted_uploaded_files(session_id).unwrap(),
|
||||
0
|
||||
);
|
||||
assert!(
|
||||
store
|
||||
.release_uploaded_file_pin(session_id, &pending.artifact_id, "submission-1")
|
||||
.is_err()
|
||||
);
|
||||
|
||||
let releasable = store
|
||||
.write_uploaded_file(session_id, "cancelled.txt", "text/plain", b"cancel", limits)
|
||||
.unwrap();
|
||||
store
|
||||
.pin_uploaded_file(session_id, &releasable, "submission-2")
|
||||
.unwrap();
|
||||
store
|
||||
.release_uploaded_file_pin(session_id, &releasable.artifact_id, "submission-2")
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
store.delete_uncommitted_uploaded_files(session_id).unwrap(),
|
||||
1
|
||||
);
|
||||
assert!(
|
||||
store
|
||||
.read_uploaded_file_by_id(session_id, &releasable.artifact_id)
|
||||
.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(_))
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -183,6 +183,7 @@ fn canonicalize_history_entry(
|
||||
item,
|
||||
metadata: legacy_metadata(segment_id, line_index, 0),
|
||||
},
|
||||
extensions: Vec::new(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -27,6 +27,7 @@
|
||||
//! system_prompt: None,
|
||||
//! config: &config,
|
||||
//! history: Vec::new(),
|
||||
//! user_segments: Vec::new(),
|
||||
//! })?;
|
||||
//! ```
|
||||
|
||||
@@ -35,11 +36,13 @@ 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;
|
||||
|
||||
@@ -53,6 +56,7 @@ pub use 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_logged_history_entry,
|
||||
create_compacted_segment, create_segment, create_segment_with_ids, ensure_head_or_fork, fork,
|
||||
@@ -64,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
|
||||
}
|
||||
}
|
||||
@@ -41,6 +41,24 @@ pub fn project_session_snapshot(session_id: SessionId, log: &[LogEntry]) -> Sess
|
||||
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,
|
||||
@@ -53,7 +71,7 @@ pub fn project_session_snapshot(session_id: SessionId, log: &[LogEntry]) -> Sess
|
||||
entries.push(history_entry(entry, *ts, data));
|
||||
}
|
||||
}
|
||||
LogEntry::AnnotatedSystemItem { ts, entry } => entries.push(system_entry(
|
||||
LogEntry::AnnotatedSystemItem { ts, entry, .. } => entries.push(system_entry(
|
||||
&entry.item,
|
||||
entry.metadata.entry_id.0.clone(),
|
||||
*ts,
|
||||
@@ -82,7 +100,10 @@ pub fn project_session_snapshot(session_id: SessionId, log: &[LogEntry]) -> Sess
|
||||
}
|
||||
}
|
||||
|
||||
SessionSnapshot { entries }
|
||||
SessionSnapshot {
|
||||
pending_submissions: protocol::PendingSubmissionsSnapshot::default(),
|
||||
entries,
|
||||
}
|
||||
}
|
||||
|
||||
fn extend_history(
|
||||
@@ -357,6 +378,63 @@ mod tests {
|
||||
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();
|
||||
|
||||
@@ -17,6 +17,33 @@ pub struct SegmentStartState<'a> {
|
||||
pub system_prompt: Option<&'a str>,
|
||||
pub config: &'a RequestConfig,
|
||||
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
|
||||
@@ -42,16 +69,8 @@ pub fn create_segment_with_ids(
|
||||
segment_id: SegmentId,
|
||||
state: SegmentStartState<'_>,
|
||||
) -> Result<(), StoreError> {
|
||||
let entry = LogEntry::AnnotatedSegmentStart {
|
||||
ts: segment_log::now_millis(),
|
||||
session_id,
|
||||
system_prompt: state.system_prompt.map(String::from),
|
||||
config: state.config.clone(),
|
||||
history: state.history.to_vec(),
|
||||
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
|
||||
@@ -68,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::AnnotatedSegmentStart {
|
||||
ts: segment_log::now_millis(),
|
||||
session_id: source_session_id,
|
||||
system_prompt: state.system_prompt.map(String::from),
|
||||
config: state.config.clone(),
|
||||
history: state.history.to_vec(),
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -152,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::AnnotatedSegmentStart {
|
||||
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: state.history.to_vec(),
|
||||
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(())
|
||||
}
|
||||
|
||||
@@ -272,6 +287,7 @@ pub fn append_system_item(
|
||||
LogEntry::AnnotatedSystemItem {
|
||||
ts: segment_log::now_millis(),
|
||||
entry,
|
||||
extensions: Vec::new(),
|
||||
},
|
||||
)
|
||||
}
|
||||
@@ -420,20 +436,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::AnnotatedSegmentStart {
|
||||
ts: segment_log::now_millis(),
|
||||
session_id,
|
||||
system_prompt: state.system_prompt.map(String::from),
|
||||
config: state.config.clone(),
|
||||
history: state.history.to_vec(),
|
||||
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))
|
||||
}
|
||||
|
||||
@@ -460,11 +470,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::AnnotatedSegmentStart { .. }))
|
||||
.position(|entry| {
|
||||
!matches!(
|
||||
entry,
|
||||
LogEntry::AnnotatedSegmentStart { .. }
|
||||
| LogEntry::InputSegmentsCheckpoint { .. }
|
||||
)
|
||||
})
|
||||
.unwrap_or(entries.len())
|
||||
} else {
|
||||
entries
|
||||
@@ -476,8 +493,9 @@ pub fn fork_at(
|
||||
let state = segment_log::collect_state(&entries[..cut]);
|
||||
|
||||
let fork_id = crate::new_segment_id();
|
||||
let ts = segment_log::now_millis();
|
||||
let entry = LogEntry::AnnotatedSegmentStart {
|
||||
ts: segment_log::now_millis(),
|
||||
ts,
|
||||
session_id: source_session_id,
|
||||
system_prompt: state.system_prompt,
|
||||
config: state.config,
|
||||
@@ -488,7 +506,14 @@ pub fn fork_at(
|
||||
}),
|
||||
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)
|
||||
}
|
||||
|
||||
|
||||
@@ -63,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
|
||||
@@ -104,6 +112,8 @@ pub enum LogEntry {
|
||||
AnnotatedSystemItem {
|
||||
ts: u64,
|
||||
entry: LoggedSystemHistoryEntry,
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
extensions: Vec<SessionExtension>,
|
||||
},
|
||||
|
||||
/// Turn boundary. Records the turn count after increment.
|
||||
@@ -273,6 +283,9 @@ pub fn collect_state(entries: &[LogEntry]) -> RestoredState {
|
||||
.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.
|
||||
@@ -301,12 +314,19 @@ pub fn collect_state(entries: &[LogEntry]) -> RestoredState {
|
||||
state.annotated_history.push(entry.clone());
|
||||
state.history.push(Item::from(entry.item.clone()));
|
||||
}
|
||||
LogEntry::AnnotatedSystemItem { entry, .. } => {
|
||||
LogEntry::AnnotatedSystemItem {
|
||||
entry, extensions, ..
|
||||
} => {
|
||||
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());
|
||||
state.extensions.extend(
|
||||
extensions
|
||||
.iter()
|
||||
.map(|extension| (extension.domain.clone(), extension.payload.clone())),
|
||||
);
|
||||
}
|
||||
LogEntry::TurnEnd { turn_count, .. } => {
|
||||
if let Some(active_turn_count) = &mut state.active_run_turn_count {
|
||||
|
||||
@@ -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,138 @@ 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)
|
||||
}
|
||||
|
||||
/// Retain an uploaded file while a durable pending operation owns it.
|
||||
fn pin_uploaded_file(
|
||||
&self,
|
||||
_session_id: SessionId,
|
||||
_reference: &UploadedFileRef,
|
||||
_owner_id: &str,
|
||||
) -> Result<(), StoreError> {
|
||||
Err(StoreError::PasteArtifactUnsupported)
|
||||
}
|
||||
|
||||
/// Release a pending-operation pin without changing committed ownership.
|
||||
fn release_uploaded_file_pin(
|
||||
&self,
|
||||
_session_id: SessionId,
|
||||
_artifact_id: &str,
|
||||
_owner_id: &str,
|
||||
) -> Result<(), StoreError> {
|
||||
Err(StoreError::PasteArtifactUnsupported)
|
||||
}
|
||||
|
||||
/// Complete the pending-to-history handoff after the history entry commits.
|
||||
fn finalize_uploaded_file_binding(
|
||||
&self,
|
||||
_session_id: SessionId,
|
||||
_artifact_id: &str,
|
||||
_source_entry_id: &str,
|
||||
) -> Result<(), StoreError> {
|
||||
Err(StoreError::PasteArtifactUnsupported)
|
||||
}
|
||||
|
||||
/// Clear pending-operation pins that have no owner in restored durable
|
||||
/// Worker Session state. This repairs an interrupted pin-before-checkpoint
|
||||
/// acceptance without disturbing live queue owners or committed history.
|
||||
fn reconcile_uploaded_file_pins(
|
||||
&self,
|
||||
_session_id: SessionId,
|
||||
_live_owner_ids: &[String],
|
||||
) -> Result<u64, StoreError> {
|
||||
Ok(0)
|
||||
}
|
||||
|
||||
/// 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,
|
||||
|
||||
@@ -0,0 +1,675 @@
|
||||
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;
|
||||
fn validate_pending_owner_id(owner_id: &str) -> Result<()> {
|
||||
if owner_id.is_empty() || owner_id.len() > 256 {
|
||||
return Err(StoreError::ArtifactIntegrityMismatch);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[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")]
|
||||
pending_owner_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,
|
||||
pending_owner_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 uploaded_file_has_pending_owner(dir: &Path, artifact_id: &str) -> Result<bool> {
|
||||
let path = record_path(dir, artifact_id)?;
|
||||
let stored: StoredUploadedFile = serde_json::from_slice(&fs::read(path)?)?;
|
||||
Ok(stored.pending_owner_id.is_some())
|
||||
}
|
||||
|
||||
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 pin_uploaded_file(
|
||||
dir: &Path,
|
||||
reference: &UploadedFileRef,
|
||||
owner_id: &str,
|
||||
) -> Result<()> {
|
||||
validate_pending_owner_id(owner_id)?;
|
||||
if reference.source_entry_id.is_some() {
|
||||
return Err(StoreError::ArtifactAlreadyCommitted);
|
||||
}
|
||||
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, &reference.artifact_id)?;
|
||||
let mut stored: StoredUploadedFile = serde_json::from_slice(&fs::read(&path)?)?;
|
||||
if stored.file_name != reference.file_name
|
||||
|| stored.media_type != reference.media_type
|
||||
|| stored.created_at_ms != reference.created_at_ms
|
||||
|| stored.byte_len != reference.byte_len
|
||||
|| stored.sha256 != reference.sha256
|
||||
{
|
||||
return Err(StoreError::ArtifactIntegrityMismatch);
|
||||
}
|
||||
if stored.source_entry_id.is_some() {
|
||||
return Err(StoreError::ArtifactAlreadyCommitted);
|
||||
}
|
||||
if let Some(existing_owner) = stored.pending_owner_id.as_deref() {
|
||||
return if existing_owner == owner_id {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(StoreError::ArtifactAlreadyCommitted)
|
||||
};
|
||||
}
|
||||
stored.pending_owner_id = Some(owner_id.to_owned());
|
||||
let temp = dir.join(format!(".{}.file.pin.tmp", reference.artifact_id));
|
||||
fs::write(&temp, serde_json::to_vec(&stored)?)?;
|
||||
fs::rename(temp, path)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) fn release_uploaded_file_pin(
|
||||
dir: &Path,
|
||||
artifact_id: &str,
|
||||
owner_id: &str,
|
||||
) -> Result<()> {
|
||||
validate_pending_owner_id(owner_id)?;
|
||||
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.pending_owner_id.as_deref() != Some(owner_id) {
|
||||
return Err(StoreError::ArtifactIntegrityMismatch);
|
||||
}
|
||||
stored.pending_owner_id = None;
|
||||
let temp = dir.join(format!(".{artifact_id}.file.unpin.tmp"));
|
||||
fs::write(&temp, serde_json::to_vec(&stored)?)?;
|
||||
fs::rename(temp, path)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) fn finalize_uploaded_file_binding(
|
||||
dir: &Path,
|
||||
artifact_id: &str,
|
||||
source_entry_id: &str,
|
||||
) -> Result<()> {
|
||||
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(source_entry_id) {
|
||||
return Err(StoreError::ArtifactIntegrityMismatch);
|
||||
}
|
||||
if stored.pending_owner_id.is_none() {
|
||||
return Ok(());
|
||||
}
|
||||
stored.pending_owner_id = None;
|
||||
let temp = dir.join(format!(".{artifact_id}.file.finalize.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 reconcile_uploaded_file_pins(dir: &Path, live_owner_ids: &[String]) -> 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 reconciled = 0_u64;
|
||||
for entry in fs::read_dir(dir)? {
|
||||
let entry = entry?;
|
||||
let path = entry.path();
|
||||
let Some(file_name) = path.file_name().and_then(|name| name.to_str()) else {
|
||||
continue;
|
||||
};
|
||||
let Some(artifact_id) = file_name.strip_suffix(".file.json") else {
|
||||
continue;
|
||||
};
|
||||
let mut stored: StoredUploadedFile = serde_json::from_slice(&fs::read(&path)?)?;
|
||||
let Some(owner_id) = stored.pending_owner_id.as_deref() else {
|
||||
continue;
|
||||
};
|
||||
if live_owner_ids.iter().any(|live| live == owner_id) {
|
||||
continue;
|
||||
}
|
||||
stored.pending_owner_id = None;
|
||||
let temp = dir.join(format!(".{artifact_id}.file.reconcile.tmp"));
|
||||
fs::write(&temp, serde_json::to_vec(&stored)?)?;
|
||||
fs::rename(temp, path)?;
|
||||
reconciled = reconciled.saturating_add(1);
|
||||
}
|
||||
Ok(reconciled)
|
||||
}
|
||||
|
||||
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() && stored.pending_owner_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() || stored.pending_owner_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()),
|
||||
}
|
||||
}
|
||||
@@ -63,6 +63,8 @@ pub struct WorkerSpawnedScopeRule {
|
||||
pub target: PathBuf,
|
||||
pub permission: String,
|
||||
pub recursive: bool,
|
||||
#[serde(default)]
|
||||
pub symlink_policy: protocol::SymlinkPolicy,
|
||||
}
|
||||
|
||||
/// One child Worker spawned by this Worker and persisted with the spawner's
|
||||
@@ -608,6 +610,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,
|
||||
@@ -664,6 +684,25 @@ mod tests {
|
||||
assert_eq!(restored, metadata);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn spawned_scope_rule_defaults_resolved_and_roundtrips_logical_policy() {
|
||||
let legacy: WorkerSpawnedScopeRule = serde_json::from_value(serde_json::json!({
|
||||
"target": "/workspace/src",
|
||||
"permission": "read",
|
||||
"recursive": true
|
||||
}))
|
||||
.unwrap();
|
||||
assert_eq!(legacy.symlink_policy, protocol::SymlinkPolicy::Resolved);
|
||||
|
||||
let logical = WorkerSpawnedScopeRule {
|
||||
symlink_policy: protocol::SymlinkPolicy::Logical,
|
||||
..legacy
|
||||
};
|
||||
let restored: WorkerSpawnedScopeRule =
|
||||
serde_json::from_value(serde_json::to_value(&logical).unwrap()).unwrap();
|
||||
assert_eq!(restored, logical);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn worker_aggregate_store_writes_one_fixed_metadata_identity() {
|
||||
let tmp = tempfile::tempdir().unwrap();
|
||||
@@ -817,6 +856,7 @@ mod tests {
|
||||
target: std::path::Path::new("/tmp/delegated").into(),
|
||||
permission: "write".into(),
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
};
|
||||
store
|
||||
.set_spawned_children(
|
||||
|
||||
@@ -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};
|
||||
@@ -25,6 +27,7 @@ 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 {
|
||||
@@ -317,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,
|
||||
@@ -601,6 +633,45 @@ mod tests {
|
||||
assert_eq!(store.list_sessions().unwrap(), vec![session_id]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
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();
|
||||
|
||||
@@ -3,13 +3,14 @@ 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};
|
||||
|
||||
// =============================================================================
|
||||
@@ -99,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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -194,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,
|
||||
@@ -236,6 +240,7 @@ async fn session_run_logs_entries() {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: annotated(&worker.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -284,6 +289,7 @@ async fn session_restore_round_trip() {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: annotated(&worker.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -323,6 +329,7 @@ async fn session_run_with_tool_call() {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: annotated(&worker.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -346,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());
|
||||
@@ -358,6 +366,7 @@ async fn session_resume_after_pause() {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: annotated(&worker.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -381,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]
|
||||
@@ -397,6 +406,7 @@ async fn session_fork_creates_new_session() {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: annotated(&worker.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -404,28 +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: 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_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"));
|
||||
}
|
||||
|
||||
@@ -441,6 +461,7 @@ async fn session_fork_at_truncates_within_session() {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: annotated(&worker.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -454,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");
|
||||
@@ -466,6 +491,7 @@ 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,
|
||||
@@ -491,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();
|
||||
@@ -503,6 +607,7 @@ async fn session_config_changed_logged() {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: annotated(&worker.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -536,6 +641,7 @@ async fn session_auto_forks_on_conflict() {
|
||||
system_prompt: worker_a.get_system_prompt(),
|
||||
config: worker_a.request_config(),
|
||||
history: annotated(&worker_a.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -564,6 +670,7 @@ async fn session_auto_forks_on_conflict() {
|
||||
system_prompt: worker_a.get_system_prompt(),
|
||||
config: worker_a.request_config(),
|
||||
history: annotated(&worker_a.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -623,6 +730,7 @@ async fn nested_past_fork_leaves_ancestors_immutable() {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: annotated(&worker.history()),
|
||||
user_segments: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -653,12 +761,19 @@ 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] {
|
||||
// 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),
|
||||
..
|
||||
|
||||
@@ -4,6 +4,7 @@ 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::{
|
||||
@@ -18,6 +19,7 @@ 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;
|
||||
@@ -43,7 +45,7 @@ pub struct StandaloneHost {
|
||||
lease: Option<StandaloneWorkerLease>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Error)]
|
||||
pub enum StandaloneStartupError {
|
||||
#[error("the standalone state store could not be opened or validated")]
|
||||
StateStore,
|
||||
@@ -53,6 +55,16 @@ pub enum StandaloneStartupError {
|
||||
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")]
|
||||
@@ -306,7 +318,11 @@ impl StandaloneHost {
|
||||
}
|
||||
|
||||
pub async fn shutdown(mut self) -> Result<(), StandaloneShutdownError> {
|
||||
let _ = self.handle.send(Method::Shutdown).await;
|
||||
let command = protocol::WorkerCommandEnvelope::for_snapshot(
|
||||
u64::MAX,
|
||||
&self.handle.shared_state.snapshot(),
|
||||
);
|
||||
let _ = self.handle.send(Method::Shutdown { command }).await;
|
||||
let Some(shutdown) = self.shutdown.take() else {
|
||||
self.retain_lease();
|
||||
return Err(StandaloneShutdownError::ConfirmationLost);
|
||||
@@ -488,7 +504,11 @@ fn active_pointer(
|
||||
}
|
||||
|
||||
async fn stop_started_worker(started: BootstrappedWorker) {
|
||||
let _ = started.handle.send(Method::Shutdown).await;
|
||||
let command = protocol::WorkerCommandEnvelope::for_snapshot(
|
||||
u64::MAX,
|
||||
&started.handle.shared_state.snapshot(),
|
||||
);
|
||||
let _ = started.handle.send(Method::Shutdown { command }).await;
|
||||
let _ = tokio::time::timeout(Duration::from_secs(2), started.shutdown).await;
|
||||
}
|
||||
|
||||
@@ -509,6 +529,15 @@ fn classify_store_startup_error(error: StandaloneStoreError) -> StandaloneStartu
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
@@ -191,8 +191,7 @@ impl StandaloneWorkerStore {
|
||||
StandaloneStoreError::Io(error)
|
||||
}
|
||||
})?;
|
||||
let record: StandaloneWorkerRecord = serde_json::from_slice(&bytes)
|
||||
.map_err(|source| StandaloneStoreError::CorruptRecord { id, source })?;
|
||||
let record = decode_worker_record(id, &bytes)?;
|
||||
if record.schema_version > SCHEMA_VERSION {
|
||||
return Err(StandaloneStoreError::NewerSchema {
|
||||
id,
|
||||
@@ -408,7 +407,7 @@ impl StandaloneWorkerStore {
|
||||
.create_new(true)
|
||||
.open(&temporary)
|
||||
.map_err(StandaloneStoreError::Io)?;
|
||||
serde_json::to_writer_pretty(&mut file, next).map_err(StandaloneStoreError::Json)?;
|
||||
write_worker_record(&mut file, next)?;
|
||||
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)?;
|
||||
@@ -428,8 +427,7 @@ impl StandaloneWorkerStore {
|
||||
) -> 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 })
|
||||
decode_worker_record(id, &bytes)
|
||||
}
|
||||
|
||||
fn worker_dir(&self, id: WorkerId) -> PathBuf {
|
||||
@@ -634,6 +632,50 @@ fn observe_process(pid: u32) -> ProcessObservation {
|
||||
}
|
||||
}
|
||||
|
||||
fn decode_worker_record(
|
||||
id: WorkerId,
|
||||
bytes: &[u8],
|
||||
) -> Result<StandaloneWorkerRecord, StandaloneStoreError> {
|
||||
let decode = || -> Result<StandaloneWorkerRecord, serde_json::Error> {
|
||||
let mut snapshot: serde_json::Value = serde_json::from_slice(bytes)?;
|
||||
let object = snapshot.as_object_mut().ok_or_else(|| {
|
||||
serde_json::Error::io(io::Error::new(
|
||||
io::ErrorKind::InvalidData,
|
||||
"standalone Worker record must be an object",
|
||||
))
|
||||
})?;
|
||||
let persisted_manifest = object.remove("manifest").ok_or_else(|| {
|
||||
serde_json::Error::io(io::Error::new(
|
||||
io::ErrorKind::InvalidData,
|
||||
"standalone Worker record is missing manifest",
|
||||
))
|
||||
})?;
|
||||
let manifest = manifest::read_persisted_worker_manifest_snapshot(persisted_manifest)?;
|
||||
object.insert("manifest".to_string(), serde_json::to_value(manifest)?);
|
||||
serde_json::from_value(snapshot)
|
||||
};
|
||||
decode().map_err(|source| StandaloneStoreError::CorruptRecord { id, source })
|
||||
}
|
||||
|
||||
fn write_worker_record(
|
||||
writer: &mut impl Write,
|
||||
record: &StandaloneWorkerRecord,
|
||||
) -> Result<(), StandaloneStoreError> {
|
||||
let mut snapshot = serde_json::to_value(record).map_err(StandaloneStoreError::Json)?;
|
||||
let object = snapshot.as_object_mut().ok_or_else(|| {
|
||||
StandaloneStoreError::Json(serde_json::Error::io(io::Error::new(
|
||||
io::ErrorKind::InvalidData,
|
||||
"standalone Worker record must be an object",
|
||||
)))
|
||||
})?;
|
||||
object.insert(
|
||||
"manifest".to_string(),
|
||||
manifest::write_persisted_worker_manifest_snapshot(&record.manifest)
|
||||
.map_err(StandaloneStoreError::Json)?,
|
||||
);
|
||||
serde_json::to_writer_pretty(writer, &snapshot).map_err(StandaloneStoreError::Json)
|
||||
}
|
||||
|
||||
fn now_unix_ms() -> Result<u64, StandaloneStoreError> {
|
||||
let duration = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
@@ -709,7 +751,70 @@ pub enum StandaloneStoreError {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{LeaseLiveness, ProcessObservation, classify_lease_liveness};
|
||||
use super::*;
|
||||
|
||||
fn test_manifest() -> WorkerManifest {
|
||||
WorkerManifest::from_toml(
|
||||
r#"
|
||||
[worker]
|
||||
name = "standalone-test"
|
||||
|
||||
[model]
|
||||
scheme = "anthropic"
|
||||
model_id = "claude-sonnet-4-20250514"
|
||||
|
||||
[engine]
|
||||
|
||||
[[scope.allow]]
|
||||
target = "/tmp"
|
||||
permission = "write"
|
||||
"#,
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn standalone_record_uses_versioned_manifest_adapter_for_legacy_memory() {
|
||||
let worker_id = "01a05782-d5dd-78f1-b9cd-ce37535bdb9d".parse().unwrap();
|
||||
let manifest = test_manifest();
|
||||
let record = StandaloneWorkerRecord {
|
||||
schema_version: SCHEMA_VERSION,
|
||||
revision: 6,
|
||||
worker_id,
|
||||
worker_name: manifest.worker.name.clone(),
|
||||
storage_key: "standalone-test".to_string(),
|
||||
cwd: StandaloneCwdIdentity {
|
||||
canonical_path: PathBuf::from("/tmp"),
|
||||
device: None,
|
||||
inode: None,
|
||||
},
|
||||
manifest,
|
||||
active_session_id: "01a05782-d5dd-78f1-b9cd-ce37535bdb9e".parse().unwrap(),
|
||||
active_segment_id: None,
|
||||
status: StandaloneWorkerStatus::Stopped,
|
||||
created_at_unix_ms: 1,
|
||||
updated_at_unix_ms: 2,
|
||||
shutdown_reason: None,
|
||||
};
|
||||
let mut legacy = serde_json::to_value(&record).unwrap();
|
||||
legacy["manifest"]["feature"]["memory"] = serde_json::json!({
|
||||
"enabled": false,
|
||||
"staging": false,
|
||||
});
|
||||
|
||||
let decoded =
|
||||
decode_worker_record(worker_id, &serde_json::to_vec(&legacy).unwrap()).unwrap();
|
||||
assert!(!decoded.manifest.feature.memory.profile.enabled);
|
||||
|
||||
let mut persisted = Vec::new();
|
||||
write_worker_record(&mut persisted, &decoded).unwrap();
|
||||
let persisted: serde_json::Value = serde_json::from_slice(&persisted).unwrap();
|
||||
assert_eq!(persisted["manifest"]["schema_version"], 2);
|
||||
assert_eq!(
|
||||
persisted["manifest"]["manifest"]["feature"]["memory"]["profile"]["enabled"],
|
||||
false
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn lease_liveness_requires_positive_live_or_stale_evidence() {
|
||||
|
||||
@@ -99,7 +99,10 @@ async fn in_process_host_runs_text_and_read_tool_then_shuts_down() {
|
||||
let mut protocol_client = host.connect();
|
||||
|
||||
protocol_client
|
||||
.send(&Method::run_text("read the probe"))
|
||||
.send(&Method::submit_text(
|
||||
protocol::new_submission_request_id(),
|
||||
"read the probe",
|
||||
))
|
||||
.await
|
||||
.expect("submit input");
|
||||
|
||||
@@ -164,6 +167,65 @@ async fn in_process_host_runs_text_and_read_tool_then_shuts_down() {
|
||||
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");
|
||||
@@ -277,11 +339,15 @@ async fn standalone_restore_preserves_history_tasks_notifications_and_cwd_scope(
|
||||
let worker_id = host.worker_id();
|
||||
let mut protocol_client = host.connect();
|
||||
protocol_client
|
||||
.send(&Method::run_text("first request"))
|
||||
.send(&Method::submit_text(
|
||||
protocol::new_submission_request_id(),
|
||||
"first request",
|
||||
))
|
||||
.await?;
|
||||
wait_for_run_end(&mut protocol_client).await?;
|
||||
protocol_client
|
||||
.send(&Method::Notify {
|
||||
notification_request_id: protocol::new_submission_request_id(),
|
||||
message: "persisted notification".to_string(),
|
||||
auto_run: true,
|
||||
})
|
||||
@@ -335,7 +401,10 @@ async fn standalone_restore_preserves_history_tasks_notifications_and_cwd_scope(
|
||||
assert!(snapshot.contains("persisted notification"), "{snapshot}");
|
||||
|
||||
protocol_client
|
||||
.send(&Method::run_text("continue after restore"))
|
||||
.send(&Method::submit_text(
|
||||
protocol::new_submission_request_id(),
|
||||
"continue after restore",
|
||||
))
|
||||
.await?;
|
||||
wait_for_run_end(&mut protocol_client).await?;
|
||||
let request = second_inspection
|
||||
|
||||
@@ -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
|
||||
}),
|
||||
@@ -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(),
|
||||
})
|
||||
}
|
||||
|
||||
@@ -118,6 +118,7 @@ impl Tool for BashTool {
|
||||
command: params.command,
|
||||
timeout_secs,
|
||||
output_limit: INLINE_BYTE_BUDGET,
|
||||
cwd: None,
|
||||
spill_dir: Some(self.output_dir.clone()),
|
||||
tool_call_id: Some(call_id.clone()),
|
||||
})
|
||||
@@ -299,11 +300,13 @@ mod tests {
|
||||
target: root.path().to_path_buf(),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
},
|
||||
ScopeRule {
|
||||
target: output.path().to_path_buf(),
|
||||
permission: Permission::Read,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
},
|
||||
],
|
||||
deny: Vec::new(),
|
||||
|
||||
@@ -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(_)
|
||||
|
||||
@@ -40,6 +40,7 @@ fn setup() -> (TempDir, TempDir, Registry) {
|
||||
target: spill.path().to_path_buf(),
|
||||
permission: Permission::Read,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
});
|
||||
let scope = Scope::from_config(&config).unwrap();
|
||||
let fs: WorkdirSessionHandle =
|
||||
|
||||
@@ -27,6 +27,7 @@ fn scope_with_spill(workspace: &Path, spill: &Path) -> Scope {
|
||||
target: spill.to_path_buf(),
|
||||
permission: Permission::Read,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
});
|
||||
Scope::from_config(&config).unwrap()
|
||||
}
|
||||
|
||||
@@ -30,4 +30,5 @@ pulldown-cmark = { version = "0.13.3", default-features = false }
|
||||
agen.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
async-trait.workspace = true
|
||||
tempfile = { workspace = true }
|
||||
|
||||
+296
-187
@@ -5,7 +5,7 @@ use std::time::{Duration, Instant};
|
||||
use protocol::{
|
||||
AlertLevel, AlertSource, CompletionEntry, CompletionKind, ErrorCode, Event, InFlightBlock,
|
||||
InFlightSnapshot, InFlightToolCallState, InternalWorkerRef, InternalWorkerSnapshot, Method,
|
||||
RewindTarget, RunResult, Segment, WorkerStatus,
|
||||
RewindTarget, RunResult, Segment, WorkerCommandEnvelope, WorkerStateSnapshot, WorkerStatus,
|
||||
};
|
||||
|
||||
use crate::block::{
|
||||
@@ -102,23 +102,6 @@ struct RollbackSubmitState {
|
||||
turn_before: usize,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct QueuedInput {
|
||||
segments: Vec<Segment>,
|
||||
preview: String,
|
||||
}
|
||||
|
||||
impl QueuedInput {
|
||||
fn new(segments: Vec<Segment>) -> Self {
|
||||
let preview = Segment::flatten_to_text(&segments);
|
||||
Self { segments, preview }
|
||||
}
|
||||
|
||||
pub fn preview(&self) -> &str {
|
||||
&self.preview
|
||||
}
|
||||
}
|
||||
|
||||
struct ComposerInputHistory {
|
||||
entries: VecDeque<Vec<Segment>>,
|
||||
browse: Option<ComposerInputHistoryBrowse>,
|
||||
@@ -242,8 +225,10 @@ pub struct WorkerViewTab {
|
||||
pub struct App {
|
||||
pub worker_name: String,
|
||||
pub connected: bool,
|
||||
/// Last controller status reported by the Worker. Drives the status line
|
||||
/// and Ctrl-key routing; do not infer this solely from replayed history.
|
||||
/// Latest authoritative revisioned live execution state.
|
||||
pub worker_state: WorkerStateSnapshot,
|
||||
next_command_id: u64,
|
||||
/// Derived Runtime-catalog compatibility projection used by existing UI.
|
||||
pub worker_status: WorkerStatus,
|
||||
/// True while the Worker is in `WorkerStatus::Running`.
|
||||
pub running: bool,
|
||||
@@ -272,7 +257,7 @@ pub struct App {
|
||||
/// Current transient actionbar notice. Notices are local UI state only:
|
||||
/// they are never appended to transcript/session history or LLM context.
|
||||
actionbar_notice: Option<ActionbarNotice>,
|
||||
/// Normal composer input that is submitted as `Method::Run`.
|
||||
/// Normal composer input that is submitted as `Method::Submit`.
|
||||
pub input: InputBuffer,
|
||||
/// Separate command-line input. It is never submitted as a user message.
|
||||
pub command_input: InputBuffer,
|
||||
@@ -333,9 +318,8 @@ pub struct App {
|
||||
/// Top entry index of the task pane's visible window. Clamped on
|
||||
/// render so it never points past the end of the list.
|
||||
pub task_pane_scroll: usize,
|
||||
/// TUI-local FIFO of user inputs submitted while the Worker is already running.
|
||||
/// Entries have not been sent to the Worker yet, so they remain editable/cancellable locally.
|
||||
queued_inputs: VecDeque<QueuedInput>,
|
||||
/// Authoritative WorkerSession FIFO summary received from snapshot/live events.
|
||||
pending_submissions: protocol::PendingSubmissionsSnapshot,
|
||||
/// TUI-local readline-style composer input history. This is intentionally
|
||||
/// client-side only: recalled entries are plain drafts until submitted again.
|
||||
input_history: ComposerInputHistory,
|
||||
@@ -355,6 +339,8 @@ impl App {
|
||||
Self {
|
||||
worker_name,
|
||||
connected: false,
|
||||
worker_state: WorkerStateSnapshot::initial(1),
|
||||
next_command_id: 1,
|
||||
worker_status: WorkerStatus::Idle,
|
||||
running: false,
|
||||
paused: false,
|
||||
@@ -395,7 +381,7 @@ impl App {
|
||||
text_selection: TextSelectionState::default(),
|
||||
task_pane_open: false,
|
||||
task_pane_scroll: 0,
|
||||
queued_inputs: VecDeque::new(),
|
||||
pending_submissions: protocol::PendingSubmissionsSnapshot::default(),
|
||||
input_history: ComposerInputHistory::new(),
|
||||
input_history_store: None,
|
||||
pending_submit_rollback: None,
|
||||
@@ -763,21 +749,53 @@ impl App {
|
||||
if self.paused {
|
||||
self.input_history.cancel_browse();
|
||||
self.input.clear();
|
||||
return Some(Method::Resume);
|
||||
let command = self.next_command_envelope();
|
||||
return Some(Method::Resume { command });
|
||||
}
|
||||
return None;
|
||||
}
|
||||
self.record_input_history(segments.clone());
|
||||
if self.running {
|
||||
self.queued_inputs.push_back(QueuedInput::new(segments));
|
||||
self.input.clear();
|
||||
self.completion = None;
|
||||
return None;
|
||||
}
|
||||
self.input.clear();
|
||||
Some(self.method_for_run(segments))
|
||||
}
|
||||
|
||||
pub fn submit_notify_input(&mut self) -> Option<Method> {
|
||||
let segments = self.input.submit_segments();
|
||||
if segments_are_blank(&segments) {
|
||||
return None;
|
||||
}
|
||||
if segments
|
||||
.iter()
|
||||
.any(|segment| matches!(segment, Segment::UploadedFile { .. }))
|
||||
{
|
||||
self.push_error("Notify accepts text only; remove attachments or queue a Submit.");
|
||||
return None;
|
||||
}
|
||||
let message = Segment::flatten_to_text(&segments);
|
||||
self.record_input_history(segments);
|
||||
self.input.clear();
|
||||
Some(Method::Notify {
|
||||
notification_request_id: protocol::new_submission_request_id(),
|
||||
message,
|
||||
auto_run: true,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn restore_unsent_run(&mut self, method: &Method) {
|
||||
let Method::Submit { 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.push_error(
|
||||
"Submit transport failed; current Composer was preserved and the unsent input was not queued.",
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
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::AnnotatedUserInput`.
|
||||
@@ -790,7 +808,10 @@ impl App {
|
||||
block_start: self.blocks.len(),
|
||||
turn_before: self.turn_index,
|
||||
});
|
||||
Method::Run { input: segments }
|
||||
Method::Submit {
|
||||
submission_request_id: protocol::new_submission_request_id(),
|
||||
input: segments,
|
||||
}
|
||||
}
|
||||
|
||||
fn record_input_history(&mut self, segments: Vec<Segment>) {
|
||||
@@ -811,7 +832,7 @@ impl App {
|
||||
}
|
||||
|
||||
pub fn queued_input_count(&self) -> usize {
|
||||
self.queued_inputs.len()
|
||||
self.pending_submissions.submissions.len()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
@@ -896,36 +917,35 @@ impl App {
|
||||
}
|
||||
}
|
||||
|
||||
pub fn continue_pending_method(&self) -> Option<Method> {
|
||||
Some(Method::ContinuePending {
|
||||
expected_revision: self.pending_submissions.revision,
|
||||
expected_head_id: self.pending_submissions.head_id.clone()?,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn clear_pending_method(&self) -> Method {
|
||||
Method::ClearPendingSubmissions {
|
||||
expected_revision: self.pending_submissions.revision,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn cancel_pending_method(&self, submission_id: String) -> Method {
|
||||
Method::CancelPendingSubmission {
|
||||
submission_id,
|
||||
expected_revision: self.pending_submissions.revision,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn next_queued_input_preview(&self) -> Option<&str> {
|
||||
self.queued_inputs.front().map(QueuedInput::preview)
|
||||
self.pending_submissions
|
||||
.submissions
|
||||
.first()
|
||||
.map(|submission| submission.submission_id.as_str())
|
||||
}
|
||||
|
||||
pub fn clear_queued_inputs(&mut self) -> usize {
|
||||
let cleared = self.queued_inputs.len();
|
||||
self.queued_inputs.clear();
|
||||
cleared
|
||||
}
|
||||
|
||||
pub fn restore_next_queued_input_to_composer(&mut self) -> bool {
|
||||
if self.queued_inputs.is_empty() {
|
||||
return false;
|
||||
}
|
||||
if !self.input.is_empty() {
|
||||
self.push_error("Composer is not empty; clear it before editing queued input.");
|
||||
return false;
|
||||
}
|
||||
let Some(queued) = self.queued_inputs.pop_front() else {
|
||||
return false;
|
||||
};
|
||||
self.input_history.cancel_browse();
|
||||
self.input.replace_with_segments(&queued.segments);
|
||||
self.completion = None;
|
||||
true
|
||||
}
|
||||
|
||||
fn pop_next_queued_run(&mut self) -> Option<Method> {
|
||||
let queued = self.queued_inputs.pop_front()?;
|
||||
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>) {
|
||||
@@ -1099,12 +1119,42 @@ impl App {
|
||||
}
|
||||
}
|
||||
|
||||
pub fn next_command_envelope(&mut self) -> WorkerCommandEnvelope {
|
||||
let command_id = self
|
||||
.next_command_id
|
||||
.max(self.worker_state.last_command_id.saturating_add(1));
|
||||
let command = WorkerCommandEnvelope::for_snapshot(command_id, &self.worker_state);
|
||||
self.next_command_id = command_id.saturating_add(1);
|
||||
command
|
||||
}
|
||||
|
||||
fn apply_worker_state_snapshot(&mut self, snapshot: &WorkerStateSnapshot) {
|
||||
match protocol::apply_worker_state_snapshot(&mut self.worker_state, snapshot) {
|
||||
Ok(protocol::WorkerStateSnapshotApply::Applied) => {
|
||||
self.set_worker_status(self.worker_state.catalog_status());
|
||||
}
|
||||
Ok(
|
||||
protocol::WorkerStateSnapshotApply::Duplicate
|
||||
| protocol::WorkerStateSnapshotApply::Stale,
|
||||
) => {}
|
||||
Err(error) => self.handle_error(
|
||||
ErrorCode::Internal,
|
||||
format!("worker state stream rejected: {error}"),
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn handle_worker_event(&mut self, event: Event) -> Option<Method> {
|
||||
if self.rewind_refresh_fence && event_is_stale_after_rewind(&event) {
|
||||
return None;
|
||||
}
|
||||
|
||||
match event {
|
||||
Event::SubmissionAccepted { .. } => {}
|
||||
Event::SubmissionRejected { message, .. } => self.push_error(message),
|
||||
Event::PendingSubmissionsChanged { pending } => {
|
||||
self.pending_submissions = pending;
|
||||
}
|
||||
Event::UserMessage { segments } => {
|
||||
self.turn_index += 1;
|
||||
self.blocks.push(Block::TurnHeader {
|
||||
@@ -1130,18 +1180,14 @@ impl App {
|
||||
self.assistant_streaming = false;
|
||||
}
|
||||
Event::TurnStart { .. } => {
|
||||
self.set_worker_status(WorkerStatus::Running);
|
||||
self.run_requests += 1;
|
||||
self.current_tool = None;
|
||||
self.latest_llm_wait_event = None;
|
||||
self.assistant_streaming = false;
|
||||
}
|
||||
Event::InvokeStart { .. } => {
|
||||
self.set_worker_status(WorkerStatus::Running);
|
||||
}
|
||||
Event::InvokeStart { .. } => {}
|
||||
// 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.
|
||||
// authoritative run state comes only from WorkerStateSnapshot.
|
||||
Event::LlmCallStart { .. } | Event::LlmCallEnd { .. } => {
|
||||
self.latest_llm_wait_event = None;
|
||||
}
|
||||
@@ -1348,15 +1394,7 @@ impl App {
|
||||
output_tokens: self.run_output_tokens,
|
||||
});
|
||||
self.pending_submit_rollback = None;
|
||||
self.reset_run_state(match result {
|
||||
RunResult::Paused => WorkerStatus::Paused,
|
||||
RunResult::Finished | RunResult::LimitReached | RunResult::RolledBack => {
|
||||
WorkerStatus::Idle
|
||||
}
|
||||
});
|
||||
if matches!(result, RunResult::Finished | RunResult::LimitReached) {
|
||||
return self.pop_next_queued_run();
|
||||
}
|
||||
self.reset_run_state();
|
||||
}
|
||||
}
|
||||
Event::CompactStart { .. } => {
|
||||
@@ -1426,14 +1464,15 @@ impl App {
|
||||
Event::Snapshot {
|
||||
session,
|
||||
greeting,
|
||||
status,
|
||||
state,
|
||||
in_flight,
|
||||
internal_workers,
|
||||
} => {
|
||||
self.rewind_refresh_fence = false;
|
||||
self.pending_submissions = session.pending_submissions.clone();
|
||||
self.restore_snapshot(&session, greeting, in_flight);
|
||||
self.replace_internal_worker_snapshots(internal_workers);
|
||||
self.set_worker_status(status);
|
||||
self.apply_worker_state_snapshot(&state);
|
||||
}
|
||||
Event::InternalWorker {
|
||||
worker,
|
||||
@@ -1443,9 +1482,12 @@ impl App {
|
||||
Event::InternalWorkerRemoved { worker, revision } => {
|
||||
self.remove_internal_worker(worker, revision)
|
||||
}
|
||||
Event::Status { status } => {
|
||||
Event::WorkerState { snapshot } => {
|
||||
self.rewind_refresh_fence = false;
|
||||
self.set_worker_status(status);
|
||||
self.apply_worker_state_snapshot(&snapshot);
|
||||
}
|
||||
Event::CommandAcknowledged { acknowledgement } => {
|
||||
self.apply_worker_state_snapshot(&acknowledgement.state);
|
||||
}
|
||||
// Command telemetry is an operational Web Console surface. The
|
||||
// TUI continues to render the final Bash ToolResult from history.
|
||||
@@ -1485,7 +1527,7 @@ impl App {
|
||||
};
|
||||
self.completion = None;
|
||||
self.close_rewind_picker();
|
||||
self.reset_run_state(self.worker_status);
|
||||
self.reset_run_state();
|
||||
let mut message = if restored_composer {
|
||||
format!(
|
||||
"Rewound session: discarded {} log entries; restored selected input to composer.",
|
||||
@@ -1533,8 +1575,7 @@ impl App {
|
||||
None
|
||||
}
|
||||
|
||||
fn reset_run_state(&mut self, status: WorkerStatus) {
|
||||
self.set_worker_status(status);
|
||||
fn reset_run_state(&mut self) {
|
||||
self.run_requests = 0;
|
||||
self.run_upload_tokens = 0;
|
||||
self.run_output_tokens = 0;
|
||||
@@ -1564,7 +1605,7 @@ impl App {
|
||||
"Rolled back empty assistant turn; no local submitted input was available to restore."
|
||||
.to_owned()
|
||||
};
|
||||
self.reset_run_state(WorkerStatus::Idle);
|
||||
self.reset_run_state();
|
||||
self.blocks.push(Block::Alert {
|
||||
level: AlertLevel::Warn,
|
||||
source: AlertSource::Worker,
|
||||
@@ -2008,12 +2049,18 @@ impl App {
|
||||
self.input_mode = CommandInputMode::Composer;
|
||||
self.command_completion_selected = None;
|
||||
}
|
||||
if let Some(Method::ListRewindTargets) = result.method.as_ref() {
|
||||
let mut method = result.method;
|
||||
if let Some(Method::Compact { .. }) = method {
|
||||
method = Some(Method::Compact {
|
||||
command: self.next_command_envelope(),
|
||||
});
|
||||
}
|
||||
if let Some(Method::ListRewindTargets) = method.as_ref() {
|
||||
self.completion = None;
|
||||
self.rewind_picker = None;
|
||||
self.rewind_request_pending = true;
|
||||
}
|
||||
result.method
|
||||
method
|
||||
}
|
||||
|
||||
fn push_command_diagnostic(&mut self, message: impl Into<String>) {
|
||||
@@ -2663,7 +2710,10 @@ mod rewind_refresh_tests {
|
||||
});
|
||||
|
||||
app.handle_worker_event(Event::RewindApplied {
|
||||
session: protocol::SessionSnapshot { entries: vec![] },
|
||||
session: protocol::SessionSnapshot {
|
||||
pending_submissions: protocol::PendingSubmissionsSnapshot::default(),
|
||||
entries: vec![],
|
||||
},
|
||||
input: vec![Segment::text("selected rewind input")],
|
||||
summary: summary(3),
|
||||
});
|
||||
@@ -2682,7 +2732,10 @@ mod rewind_refresh_tests {
|
||||
});
|
||||
|
||||
app.handle_worker_event(Event::RewindApplied {
|
||||
session: protocol::SessionSnapshot { entries: vec![] },
|
||||
session: protocol::SessionSnapshot {
|
||||
pending_submissions: protocol::PendingSubmissionsSnapshot::default(),
|
||||
entries: vec![],
|
||||
},
|
||||
input: vec![Segment::text("rewound input")],
|
||||
summary: summary(1),
|
||||
});
|
||||
@@ -2725,7 +2778,10 @@ mod rewind_refresh_tests {
|
||||
});
|
||||
|
||||
app.handle_worker_event(Event::RewindApplied {
|
||||
session: protocol::SessionSnapshot { entries: vec![] },
|
||||
session: protocol::SessionSnapshot {
|
||||
pending_submissions: protocol::PendingSubmissionsSnapshot::default(),
|
||||
entries: vec![],
|
||||
},
|
||||
input: vec![Segment::text("rewound input")],
|
||||
summary: summary(2),
|
||||
});
|
||||
@@ -2734,8 +2790,8 @@ mod rewind_refresh_tests {
|
||||
});
|
||||
assert!(!blocks_contain(&app, "stale tail after rewind"));
|
||||
|
||||
app.handle_worker_event(Event::Status {
|
||||
status: WorkerStatus::Idle,
|
||||
app.handle_worker_event(Event::WorkerState {
|
||||
snapshot: WorkerStatus::Idle.into(),
|
||||
});
|
||||
app.handle_worker_event(Event::TextDelta {
|
||||
text: "new live tail after status".into(),
|
||||
@@ -2859,7 +2915,7 @@ mod composer_history_persistence_tests {
|
||||
path: "src/lib.rs".into(),
|
||||
},
|
||||
]);
|
||||
assert!(matches!(app.submit_input(), Some(Method::Run { .. })));
|
||||
assert!(matches!(app.submit_input(), Some(Method::Submit { .. })));
|
||||
|
||||
let mut reloaded = App::new_with_input_history_store("test".into(), store);
|
||||
assert!(reloaded.browse_input_history_older());
|
||||
@@ -2940,7 +2996,7 @@ mod composer_history_persistence_tests {
|
||||
app.insert_char(c);
|
||||
}
|
||||
match app.submit_input() {
|
||||
Some(Method::Run { input }) => input,
|
||||
Some(Method::Submit { input, .. }) => input,
|
||||
other => panic!("expected Run, got {other:?}"),
|
||||
}
|
||||
}
|
||||
@@ -3406,72 +3462,44 @@ mod completion_flow_tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn running_submit_is_queued_locally_and_clears_composer() {
|
||||
fn running_submit_is_sent_to_the_worker_and_not_queued_locally() {
|
||||
let mut app = App::new("test".into());
|
||||
app.set_worker_status(WorkerStatus::Running);
|
||||
insert_text(&mut app, "queued turn");
|
||||
|
||||
assert!(app.submit_input().is_none());
|
||||
let method = app.submit_input();
|
||||
|
||||
assert_eq!(app.queued_input_count(), 1);
|
||||
assert_eq!(app.next_queued_input_preview(), Some("queued turn"));
|
||||
assert!(matches!(method, Some(Method::Submit { .. })));
|
||||
assert_eq!(app.queued_input_count(), 0);
|
||||
assert_eq!(input_text(&app), "");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn finished_run_auto_sends_next_queued_input() {
|
||||
fn pending_submission_projection_is_worker_authoritative() {
|
||||
let mut app = App::new("test".into());
|
||||
app.set_worker_status(WorkerStatus::Running);
|
||||
insert_text(&mut app, "next turn");
|
||||
assert!(app.submit_input().is_none());
|
||||
|
||||
let method = app.handle_worker_event(Event::RunEnd {
|
||||
result: RunResult::Finished,
|
||||
app.handle_worker_event(Event::PendingSubmissionsChanged {
|
||||
pending: protocol::PendingSubmissionsSnapshot {
|
||||
revision: 3,
|
||||
notification_count: 0,
|
||||
head_id: Some("submission-1".into()),
|
||||
submissions: vec![protocol::PendingSubmissionSummary {
|
||||
submission_id: "submission-1".into(),
|
||||
accepted_at_ms: 7,
|
||||
segment_count: 2,
|
||||
byte_len: 42,
|
||||
}],
|
||||
},
|
||||
});
|
||||
|
||||
match method {
|
||||
Some(Method::Run { input }) => {
|
||||
assert_eq!(Segment::flatten_to_text(&input), "next turn");
|
||||
}
|
||||
other => panic!("expected queued Run, got {other:?}"),
|
||||
}
|
||||
assert_eq!(app.queued_input_count(), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn limit_reached_run_auto_sends_next_queued_input() {
|
||||
let mut app = App::new("test".into());
|
||||
app.set_worker_status(WorkerStatus::Running);
|
||||
insert_text(&mut app, "next after limit");
|
||||
assert!(app.submit_input().is_none());
|
||||
|
||||
let method = app.handle_worker_event(Event::RunEnd {
|
||||
result: RunResult::LimitReached,
|
||||
});
|
||||
|
||||
match method {
|
||||
Some(Method::Run { input }) => {
|
||||
assert_eq!(Segment::flatten_to_text(&input), "next after limit");
|
||||
}
|
||||
other => panic!("expected queued Run, got {other:?}"),
|
||||
}
|
||||
assert_eq!(app.queued_input_count(), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn paused_and_rolled_back_run_do_not_auto_send_queue() {
|
||||
for result in [RunResult::Paused, RunResult::RolledBack] {
|
||||
let mut app = App::new("test".into());
|
||||
app.set_worker_status(WorkerStatus::Running);
|
||||
insert_text(&mut app, "held turn");
|
||||
assert!(app.submit_input().is_none());
|
||||
|
||||
let method = app.handle_worker_event(Event::RunEnd { result });
|
||||
|
||||
assert!(method.is_none());
|
||||
assert_eq!(app.queued_input_count(), 1);
|
||||
assert_eq!(app.next_queued_input_preview(), Some("held turn"));
|
||||
}
|
||||
assert_eq!(app.queued_input_count(), 1);
|
||||
assert_eq!(app.next_queued_input_preview(), Some("submission-1"));
|
||||
assert!(
|
||||
app.handle_worker_event(Event::RunEnd {
|
||||
result: RunResult::Finished,
|
||||
})
|
||||
.is_none()
|
||||
);
|
||||
assert_eq!(app.queued_input_count(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -3479,25 +3507,7 @@ mod completion_flow_tests {
|
||||
let mut app = App::new("test".into());
|
||||
app.set_worker_status(WorkerStatus::Paused);
|
||||
|
||||
assert!(matches!(app.submit_input(), Some(Method::Resume)));
|
||||
assert_eq!(app.queued_input_count(), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn queued_input_can_be_restored_to_composer_or_cleared() {
|
||||
let mut app = App::new("test".into());
|
||||
app.set_worker_status(WorkerStatus::Running);
|
||||
insert_text(&mut app, "edit me");
|
||||
assert!(app.submit_input().is_none());
|
||||
|
||||
assert!(app.restore_next_queued_input_to_composer());
|
||||
assert_eq!(app.queued_input_count(), 0);
|
||||
assert_eq!(input_text(&app), "edit me");
|
||||
|
||||
app.input.clear();
|
||||
insert_text(&mut app, "clear me");
|
||||
assert!(app.submit_input().is_none());
|
||||
assert_eq!(app.clear_queued_inputs(), 1);
|
||||
assert!(matches!(app.submit_input(), Some(Method::Resume { .. })));
|
||||
assert_eq!(app.queued_input_count(), 0);
|
||||
}
|
||||
|
||||
@@ -3512,7 +3522,7 @@ mod completion_flow_tests {
|
||||
app.insert_char(c);
|
||||
}
|
||||
match app.submit_input() {
|
||||
Some(Method::Run { input }) => input,
|
||||
Some(Method::Submit { input, .. }) => input,
|
||||
other => panic!("expected Run, got {other:?}"),
|
||||
}
|
||||
}
|
||||
@@ -3552,7 +3562,7 @@ mod completion_flow_tests {
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
session: public_session(vec![session_start_value]),
|
||||
status: WorkerStatus::Running,
|
||||
state: test_worker_state(WorkerStatus::Running),
|
||||
in_flight: Default::default(),
|
||||
internal_workers: Vec::new(),
|
||||
});
|
||||
@@ -3563,6 +3573,90 @@ mod completion_flow_tests {
|
||||
assert!(matches!(app.blocks.first(), Some(Block::Greeting(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn occurrence_events_do_not_infer_foreground_worker_state() {
|
||||
let mut app = App::new("test".into());
|
||||
app.handle_worker_event(Event::TurnStart { turn: 1 });
|
||||
app.handle_worker_event(Event::InvokeStart {
|
||||
kind: protocol::InvokeKind::UserSend,
|
||||
});
|
||||
app.handle_worker_event(Event::RunEnd {
|
||||
result: RunResult::Paused,
|
||||
});
|
||||
assert_eq!(app.worker_state.state, protocol::WorkerState::Idle);
|
||||
assert_eq!(app.worker_status, WorkerStatus::Idle);
|
||||
|
||||
let running = WorkerStateSnapshot {
|
||||
execution_generation: 1,
|
||||
revision: 1,
|
||||
state: protocol::WorkerState::Busy(protocol::WorkerBusyState::Run(
|
||||
protocol::WorkerRunState::Running,
|
||||
)),
|
||||
last_command_id: 0,
|
||||
};
|
||||
app.handle_worker_event(Event::WorkerState {
|
||||
snapshot: running.clone(),
|
||||
});
|
||||
app.handle_worker_event(Event::RunEnd {
|
||||
result: RunResult::Finished,
|
||||
});
|
||||
assert_eq!(app.worker_state, running);
|
||||
assert_eq!(app.worker_status, WorkerStatus::Running);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn worker_state_events_and_acknowledgements_share_monotonic_application() {
|
||||
let mut app = App::new("test".into());
|
||||
let running = WorkerStateSnapshot {
|
||||
execution_generation: 4,
|
||||
revision: 3,
|
||||
state: protocol::WorkerState::Busy(protocol::WorkerBusyState::Run(
|
||||
protocol::WorkerRunState::Running,
|
||||
)),
|
||||
last_command_id: 2,
|
||||
};
|
||||
app.handle_worker_event(Event::WorkerState {
|
||||
snapshot: running.clone(),
|
||||
});
|
||||
app.handle_worker_event(Event::WorkerState {
|
||||
snapshot: WorkerStateSnapshot {
|
||||
revision: 2,
|
||||
state: protocol::WorkerState::Idle,
|
||||
..running.clone()
|
||||
},
|
||||
});
|
||||
assert_eq!(app.worker_state, running);
|
||||
|
||||
let paused = WorkerStateSnapshot {
|
||||
revision: 4,
|
||||
state: protocol::WorkerState::Busy(protocol::WorkerBusyState::Run(
|
||||
protocol::WorkerRunState::Paused,
|
||||
)),
|
||||
last_command_id: 3,
|
||||
..running.clone()
|
||||
};
|
||||
app.handle_worker_event(Event::CommandAcknowledged {
|
||||
acknowledgement: protocol::WorkerCommandAcknowledgement {
|
||||
command_id: 3,
|
||||
command: protocol::WorkerCommandKind::Pause,
|
||||
disposition: protocol::WorkerCommandDisposition::Accepted,
|
||||
state: paused.clone(),
|
||||
},
|
||||
});
|
||||
assert_eq!(app.worker_state, paused);
|
||||
|
||||
app.handle_worker_event(Event::WorkerState {
|
||||
snapshot: WorkerStateSnapshot {
|
||||
state: protocol::WorkerState::Idle,
|
||||
..paused.clone()
|
||||
},
|
||||
});
|
||||
assert_eq!(app.worker_state, paused);
|
||||
assert!(app.run_error_messages.iter().any(|message| {
|
||||
message.contains("conflicting worker state snapshots at generation 4 revision 4")
|
||||
}));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn snapshot_replaces_live_error_with_one_durable_run_error_block() {
|
||||
let mut app = App::new("test".into());
|
||||
@@ -3570,8 +3664,8 @@ mod completion_flow_tests {
|
||||
code: ErrorCode::ProviderError,
|
||||
message: "provider unavailable".into(),
|
||||
});
|
||||
app.handle_worker_event(Event::Status {
|
||||
status: WorkerStatus::Idle,
|
||||
app.handle_worker_event(Event::WorkerState {
|
||||
snapshot: WorkerStatus::Idle.into(),
|
||||
});
|
||||
|
||||
let live_errors = app
|
||||
@@ -3596,7 +3690,7 @@ mod completion_flow_tests {
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
session: public_session(vec![serde_json::to_value(run_errored).unwrap()]),
|
||||
status: WorkerStatus::Idle,
|
||||
state: test_worker_state(WorkerStatus::Idle),
|
||||
in_flight: Default::default(),
|
||||
internal_workers: Vec::new(),
|
||||
});
|
||||
@@ -3657,9 +3751,10 @@ mod completion_flow_tests {
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
session: protocol::SessionSnapshot {
|
||||
pending_submissions: protocol::PendingSubmissionsSnapshot::default(),
|
||||
entries: Vec::new(),
|
||||
},
|
||||
status: WorkerStatus::Running,
|
||||
state: test_worker_state(WorkerStatus::Running),
|
||||
in_flight: InFlightSnapshot {
|
||||
blocks: vec![
|
||||
InFlightBlock::Thinking {
|
||||
@@ -3765,6 +3860,7 @@ mod completion_flow_tests {
|
||||
revision,
|
||||
status: WorkerStatus::Idle,
|
||||
session: protocol::SessionSnapshot {
|
||||
pending_submissions: protocol::PendingSubmissionsSnapshot::default(),
|
||||
entries: Vec::new(),
|
||||
},
|
||||
in_flight: protocol::InFlightSnapshot::default(),
|
||||
@@ -3982,9 +4078,10 @@ mod completion_flow_tests {
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
session: protocol::SessionSnapshot {
|
||||
pending_submissions: protocol::PendingSubmissionsSnapshot::default(),
|
||||
entries: Vec::new(),
|
||||
},
|
||||
status: WorkerStatus::Idle,
|
||||
state: test_worker_state(WorkerStatus::Idle),
|
||||
in_flight: Default::default(),
|
||||
internal_workers: Vec::new(),
|
||||
});
|
||||
@@ -4033,9 +4130,10 @@ mod completion_flow_tests {
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
session: protocol::SessionSnapshot {
|
||||
pending_submissions: protocol::PendingSubmissionsSnapshot::default(),
|
||||
entries: Vec::new(),
|
||||
},
|
||||
status: WorkerStatus::Idle,
|
||||
state: test_worker_state(WorkerStatus::Idle),
|
||||
in_flight: Default::default(),
|
||||
internal_workers: vec![InternalWorkerSnapshot {
|
||||
worker: InternalWorkerRef {
|
||||
@@ -4046,6 +4144,7 @@ mod completion_flow_tests {
|
||||
},
|
||||
revision: 4,
|
||||
session: protocol::SessionSnapshot {
|
||||
pending_submissions: protocol::PendingSubmissionsSnapshot::default(),
|
||||
entries: Vec::new(),
|
||||
},
|
||||
status: WorkerStatus::Running,
|
||||
@@ -4182,6 +4281,13 @@ mod completion_flow_tests {
|
||||
.count()
|
||||
}
|
||||
|
||||
fn test_worker_state(status: WorkerStatus) -> WorkerStateSnapshot {
|
||||
let mut snapshot = WorkerStateSnapshot::from(status);
|
||||
snapshot.execution_generation = 1;
|
||||
snapshot.revision = 1;
|
||||
snapshot
|
||||
}
|
||||
|
||||
fn test_greeting() -> protocol::Greeting {
|
||||
protocol::Greeting {
|
||||
worker_name: "test".into(),
|
||||
@@ -4204,10 +4310,11 @@ mod completion_flow_tests {
|
||||
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
session: protocol::SessionSnapshot {
|
||||
pending_submissions: protocol::PendingSubmissionsSnapshot::default(),
|
||||
entries: Vec::new(),
|
||||
},
|
||||
greeting,
|
||||
status: WorkerStatus::Idle,
|
||||
state: test_worker_state(WorkerStatus::Idle),
|
||||
in_flight: Default::default(),
|
||||
internal_workers: Vec::new(),
|
||||
});
|
||||
@@ -4406,7 +4513,7 @@ mod completion_flow_tests {
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
session: public_session(assistant_item_entries),
|
||||
status: WorkerStatus::Running,
|
||||
state: test_worker_state(WorkerStatus::Running),
|
||||
in_flight: Default::default(),
|
||||
internal_workers: Vec::new(),
|
||||
});
|
||||
@@ -4419,23 +4526,23 @@ mod completion_flow_tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn input_history_records_queued_inputs_and_suppresses_consecutive_duplicates() {
|
||||
fn input_history_records_running_submits_and_suppresses_consecutive_duplicates() {
|
||||
let mut app = App::new("test".into());
|
||||
app.running = true;
|
||||
|
||||
for c in "repeat".chars() {
|
||||
app.insert_char(c);
|
||||
}
|
||||
assert!(app.submit_input().is_none());
|
||||
assert!(app.submit_input().is_some());
|
||||
assert_eq!(app.input_history_len(), 1);
|
||||
assert_eq!(app.queued_input_count(), 1);
|
||||
assert_eq!(app.queued_input_count(), 0);
|
||||
|
||||
for c in "repeat".chars() {
|
||||
app.insert_char(c);
|
||||
}
|
||||
assert!(app.submit_input().is_none());
|
||||
assert!(app.submit_input().is_some());
|
||||
assert_eq!(app.input_history_len(), 1);
|
||||
assert_eq!(app.queued_input_count(), 2);
|
||||
assert_eq!(app.queued_input_count(), 0);
|
||||
|
||||
app.insert_char(' ');
|
||||
assert!(app.submit_input().is_none());
|
||||
@@ -4463,7 +4570,7 @@ mod completion_flow_tests {
|
||||
},
|
||||
];
|
||||
app.input.replace_with_segments(&original);
|
||||
assert!(matches!(app.submit_input(), Some(Method::Run { .. })));
|
||||
assert!(matches!(app.submit_input(), Some(Method::Submit { .. })));
|
||||
|
||||
assert!(app.browse_input_history_older());
|
||||
assert_eq!(app.input.submit_segments(), original);
|
||||
@@ -4475,7 +4582,7 @@ mod completion_flow_tests {
|
||||
for c in "sent".chars() {
|
||||
app.insert_char(c);
|
||||
}
|
||||
assert!(matches!(app.submit_input(), Some(Method::Run { .. })));
|
||||
assert!(matches!(app.submit_input(), Some(Method::Submit { .. })));
|
||||
|
||||
for c in "draft".chars() {
|
||||
app.insert_char(c);
|
||||
@@ -4493,7 +4600,7 @@ mod completion_flow_tests {
|
||||
for c in "sent".chars() {
|
||||
app.insert_char(c);
|
||||
}
|
||||
assert!(matches!(app.submit_input(), Some(Method::Run { .. })));
|
||||
assert!(matches!(app.submit_input(), Some(Method::Submit { .. })));
|
||||
|
||||
assert!(app.browse_input_history_older());
|
||||
assert!(app.input_history_is_browsing());
|
||||
@@ -4510,17 +4617,19 @@ mod completion_flow_tests {
|
||||
for c in "first".chars() {
|
||||
app.insert_char(c);
|
||||
}
|
||||
assert!(matches!(app.submit_input(), Some(Method::Run { .. })));
|
||||
assert!(matches!(app.submit_input(), Some(Method::Submit { .. })));
|
||||
for c in "second".chars() {
|
||||
app.insert_char(c);
|
||||
}
|
||||
assert!(matches!(app.submit_input(), Some(Method::Run { .. })));
|
||||
assert!(matches!(app.submit_input(), Some(Method::Submit { .. })));
|
||||
|
||||
assert!(app.browse_input_history_older());
|
||||
assert!(app.browse_input_history_older());
|
||||
let method = app.submit_input();
|
||||
match method {
|
||||
Some(Method::Run { input }) => assert_eq!(Segment::flatten_to_text(&input), "first"),
|
||||
Some(Method::Submit { input, .. }) => {
|
||||
assert_eq!(Segment::flatten_to_text(&input), "first")
|
||||
}
|
||||
other => panic!("expected recalled run, got {other:?}"),
|
||||
}
|
||||
assert_eq!(app.input_history_len(), 3);
|
||||
|
||||
@@ -0,0 +1,483 @@
|
||||
use client::{
|
||||
BackendCreateWorkerRequest, BackendWorkerLaunchOptions, BackendWorkerLaunchProfileCandidate,
|
||||
BackendWorkerLaunchRuntimeOption, BackendWorkerLaunchTarget, create_backend_worker,
|
||||
get_backend_worker_launch_options,
|
||||
};
|
||||
use crossterm::event::{self, Event, KeyCode, KeyEventKind, KeyModifiers};
|
||||
use ratatui::layout::{Constraint, Direction, Layout};
|
||||
use ratatui::style::{Color, Modifier, Style};
|
||||
use ratatui::text::{Line, Span};
|
||||
use ratatui::widgets::{Block, Borders, Paragraph, Wrap};
|
||||
|
||||
use crate::backend_workspace_picker::select_backend_workspace;
|
||||
use crate::console;
|
||||
use crate::inline_terminal::{InlineTerminal, with_inline_terminal};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum Field {
|
||||
Name,
|
||||
Runtime,
|
||||
Profile,
|
||||
}
|
||||
|
||||
impl Field {
|
||||
fn next(self) -> Self {
|
||||
match self {
|
||||
Self::Name => Self::Runtime,
|
||||
Self::Runtime => Self::Profile,
|
||||
Self::Profile => Self::Name,
|
||||
}
|
||||
}
|
||||
|
||||
fn previous(self) -> Self {
|
||||
match self {
|
||||
Self::Name => Self::Profile,
|
||||
Self::Runtime => Self::Name,
|
||||
Self::Profile => Self::Runtime,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
struct Selection {
|
||||
runtime_id: String,
|
||||
display_name: String,
|
||||
profile: String,
|
||||
}
|
||||
|
||||
struct FormState {
|
||||
field: Field,
|
||||
display_name: String,
|
||||
runtime_index: usize,
|
||||
profile_index: usize,
|
||||
status: String,
|
||||
}
|
||||
|
||||
impl FormState {
|
||||
fn new(options: &BackendWorkerLaunchOptions) -> Self {
|
||||
let runtime_index = options
|
||||
.runtimes
|
||||
.iter()
|
||||
.position(runtime_supports_workdirless_creation)
|
||||
.unwrap_or(0);
|
||||
let profile_index = options
|
||||
.default_profile
|
||||
.as_deref()
|
||||
.and_then(|default| {
|
||||
options
|
||||
.profiles
|
||||
.iter()
|
||||
.position(|candidate| candidate.id == default)
|
||||
})
|
||||
.unwrap_or(0);
|
||||
Self {
|
||||
field: Field::Name,
|
||||
display_name: "Worker".to_string(),
|
||||
runtime_index,
|
||||
profile_index,
|
||||
status: String::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn current_runtime<'a>(
|
||||
&self,
|
||||
options: &'a BackendWorkerLaunchOptions,
|
||||
) -> Option<&'a BackendWorkerLaunchRuntimeOption> {
|
||||
options.runtimes.get(self.runtime_index)
|
||||
}
|
||||
|
||||
fn current_profile<'a>(
|
||||
&self,
|
||||
options: &'a BackendWorkerLaunchOptions,
|
||||
) -> Option<&'a BackendWorkerLaunchProfileCandidate> {
|
||||
options.profiles.get(self.profile_index)
|
||||
}
|
||||
|
||||
fn cycle_runtime(&mut self, options: &BackendWorkerLaunchOptions, delta: isize) {
|
||||
self.runtime_index = cycle_index(self.runtime_index, options.runtimes.len(), delta);
|
||||
self.status.clear();
|
||||
}
|
||||
|
||||
fn cycle_profile(&mut self, options: &BackendWorkerLaunchOptions, delta: isize) {
|
||||
self.profile_index = cycle_index(self.profile_index, options.profiles.len(), delta);
|
||||
self.status.clear();
|
||||
}
|
||||
|
||||
fn submit(&mut self, options: &BackendWorkerLaunchOptions) -> Option<Selection> {
|
||||
let display_name = self.display_name.trim();
|
||||
if display_name.is_empty() {
|
||||
self.status = "Worker name is required.".to_string();
|
||||
self.field = Field::Name;
|
||||
return None;
|
||||
}
|
||||
let Some(runtime) = self.current_runtime(options) else {
|
||||
self.status = "No Runtime is available in this Workspace.".to_string();
|
||||
self.field = Field::Runtime;
|
||||
return None;
|
||||
};
|
||||
if !runtime.worker_creation_available {
|
||||
self.status = "The selected Runtime cannot create Workers right now.".to_string();
|
||||
self.field = Field::Runtime;
|
||||
return None;
|
||||
}
|
||||
if runtime.working_directory_required {
|
||||
self.status =
|
||||
"The selected Runtime requires a workdir; this launch flow does not select one yet."
|
||||
.to_string();
|
||||
self.field = Field::Runtime;
|
||||
return None;
|
||||
}
|
||||
let Some(profile) = self.current_profile(options) else {
|
||||
self.status = "No Worker profile is available.".to_string();
|
||||
self.field = Field::Profile;
|
||||
return None;
|
||||
};
|
||||
Some(Selection {
|
||||
runtime_id: runtime.runtime_id.clone(),
|
||||
display_name: display_name.to_string(),
|
||||
profile: profile.id.clone(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn run(mut target: BackendWorkerLaunchTarget) -> Result<(), Box<dyn std::error::Error>> {
|
||||
if target.workspace_id().is_none() {
|
||||
let Some(workspace) = select_backend_workspace(&target.base_url).await? else {
|
||||
return Ok(());
|
||||
};
|
||||
target.select_workspace(workspace);
|
||||
}
|
||||
|
||||
let options = get_backend_worker_launch_options(&target).await?;
|
||||
let Some(selection) = select_worker(&options)? else {
|
||||
return Ok(());
|
||||
};
|
||||
let request = request_from_selection(selection);
|
||||
let created = create_backend_worker(&target, &request).await?;
|
||||
let runtime_target = target.runtime_target(created.runtime_id, created.worker_id)?;
|
||||
console::run_backend_runtime(runtime_target).await
|
||||
}
|
||||
|
||||
fn request_from_selection(selection: Selection) -> BackendCreateWorkerRequest {
|
||||
BackendCreateWorkerRequest {
|
||||
runtime_id: selection.runtime_id,
|
||||
display_name: selection.display_name,
|
||||
profile: Some(selection.profile),
|
||||
initial_submit: Vec::new(),
|
||||
working_directory: None,
|
||||
ticket_assignment: None,
|
||||
control_operation_id: None,
|
||||
}
|
||||
}
|
||||
|
||||
const VIEWPORT_LINES: u16 = 14;
|
||||
|
||||
fn select_worker(
|
||||
options: &BackendWorkerLaunchOptions,
|
||||
) -> Result<Option<Selection>, Box<dyn std::error::Error>> {
|
||||
with_inline_terminal(VIEWPORT_LINES, |terminal| run_form(terminal, options))
|
||||
}
|
||||
|
||||
fn run_form(
|
||||
terminal: &mut InlineTerminal,
|
||||
options: &BackendWorkerLaunchOptions,
|
||||
) -> Result<Option<Selection>, Box<dyn std::error::Error>> {
|
||||
let mut state = FormState::new(options);
|
||||
|
||||
loop {
|
||||
terminal.draw(|frame| render(frame, &state, options))?;
|
||||
let event = event::read()?;
|
||||
let Event::Key(key) = event else {
|
||||
continue;
|
||||
};
|
||||
if key.kind != KeyEventKind::Press {
|
||||
continue;
|
||||
}
|
||||
if key.code == KeyCode::Char('c') && key.modifiers.contains(KeyModifiers::CONTROL) {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
match key.code {
|
||||
KeyCode::Esc => {
|
||||
return Ok(None);
|
||||
}
|
||||
KeyCode::Tab | KeyCode::Down => {
|
||||
state.field = state.field.next();
|
||||
state.status.clear();
|
||||
}
|
||||
KeyCode::BackTab | KeyCode::Up => {
|
||||
state.field = state.field.previous();
|
||||
state.status.clear();
|
||||
}
|
||||
KeyCode::Left => match state.field {
|
||||
Field::Runtime => state.cycle_runtime(options, -1),
|
||||
Field::Profile => state.cycle_profile(options, -1),
|
||||
Field::Name => {}
|
||||
},
|
||||
KeyCode::Right => match state.field {
|
||||
Field::Runtime => state.cycle_runtime(options, 1),
|
||||
Field::Profile => state.cycle_profile(options, 1),
|
||||
Field::Name => {}
|
||||
},
|
||||
KeyCode::Enter => {
|
||||
if let Some(selection) = state.submit(options) {
|
||||
return Ok(Some(selection));
|
||||
}
|
||||
}
|
||||
KeyCode::Backspace if state.field == Field::Name => {
|
||||
state.display_name.pop();
|
||||
state.status.clear();
|
||||
}
|
||||
KeyCode::Char(character)
|
||||
if state.field == Field::Name
|
||||
&& !key.modifiers.contains(KeyModifiers::CONTROL)
|
||||
&& !character.is_control() =>
|
||||
{
|
||||
state.display_name.push(character);
|
||||
state.status.clear();
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn render(frame: &mut ratatui::Frame<'_>, state: &FormState, options: &BackendWorkerLaunchOptions) {
|
||||
let area = frame.area();
|
||||
let vertical = Layout::default()
|
||||
.direction(Direction::Vertical)
|
||||
.constraints([
|
||||
Constraint::Length(1),
|
||||
Constraint::Length(3),
|
||||
Constraint::Length(3),
|
||||
Constraint::Length(3),
|
||||
Constraint::Length(3),
|
||||
Constraint::Min(1),
|
||||
])
|
||||
.split(area);
|
||||
|
||||
frame.render_widget(
|
||||
Paragraph::new(Line::from(vec![
|
||||
Span::styled(
|
||||
"New Backend Worker",
|
||||
Style::default().add_modifier(Modifier::BOLD),
|
||||
),
|
||||
Span::raw(format!(" Workspace: {}", options.workspace_id)),
|
||||
])),
|
||||
vertical[0],
|
||||
);
|
||||
|
||||
let focused = Style::default().fg(Color::Cyan);
|
||||
frame.render_widget(
|
||||
Paragraph::new(state.display_name.as_str()).block(
|
||||
Block::default()
|
||||
.borders(Borders::ALL)
|
||||
.title(" Name ")
|
||||
.border_style(if state.field == Field::Name {
|
||||
focused
|
||||
} else {
|
||||
Style::default()
|
||||
}),
|
||||
),
|
||||
vertical[1],
|
||||
);
|
||||
|
||||
let runtime_text = state
|
||||
.current_runtime(options)
|
||||
.map(runtime_label)
|
||||
.unwrap_or_else(|| "No Runtime available".to_string());
|
||||
frame.render_widget(
|
||||
Paragraph::new(runtime_text).block(
|
||||
Block::default()
|
||||
.borders(Borders::ALL)
|
||||
.title(runtime_title(state, options))
|
||||
.border_style(if state.field == Field::Runtime {
|
||||
focused
|
||||
} else {
|
||||
Style::default()
|
||||
}),
|
||||
),
|
||||
vertical[2],
|
||||
);
|
||||
|
||||
let profile_text = state
|
||||
.current_profile(options)
|
||||
.map(|profile| {
|
||||
if profile.description.is_empty() {
|
||||
profile.label.clone()
|
||||
} else {
|
||||
format!("{} — {}", profile.label, profile.description)
|
||||
}
|
||||
})
|
||||
.unwrap_or_else(|| "No profile available".to_string());
|
||||
frame.render_widget(
|
||||
Paragraph::new(profile_text).block(
|
||||
Block::default()
|
||||
.borders(Borders::ALL)
|
||||
.title(profile_title(state, options))
|
||||
.border_style(if state.field == Field::Profile {
|
||||
focused
|
||||
} else {
|
||||
Style::default()
|
||||
}),
|
||||
),
|
||||
vertical[3],
|
||||
);
|
||||
|
||||
let status = if state.status.is_empty() {
|
||||
"Tab/↑/↓: field ←/→: choice Enter: create Esc/Ctrl-C: cancel"
|
||||
} else {
|
||||
state.status.as_str()
|
||||
};
|
||||
frame.render_widget(
|
||||
Paragraph::new(status)
|
||||
.style(if state.status.is_empty() {
|
||||
Style::default().fg(Color::DarkGray)
|
||||
} else {
|
||||
Style::default().fg(Color::Yellow)
|
||||
})
|
||||
.wrap(Wrap { trim: true }),
|
||||
vertical[4],
|
||||
);
|
||||
|
||||
if state.field == Field::Name {
|
||||
let max_cursor = vertical[1].width.saturating_sub(2) as usize;
|
||||
frame.set_cursor_position((
|
||||
vertical[1].x + 1 + state.display_name.chars().count().min(max_cursor) as u16,
|
||||
vertical[1].y + 1,
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
fn runtime_title(state: &FormState, options: &BackendWorkerLaunchOptions) -> String {
|
||||
if options.runtimes.is_empty() {
|
||||
" Runtime ".to_string()
|
||||
} else {
|
||||
format!(
|
||||
" Runtime ({}/{}) ",
|
||||
state.runtime_index + 1,
|
||||
options.runtimes.len()
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
fn profile_title(state: &FormState, options: &BackendWorkerLaunchOptions) -> String {
|
||||
if options.profiles.is_empty() {
|
||||
" Profile ".to_string()
|
||||
} else {
|
||||
format!(
|
||||
" Profile ({}/{}) ",
|
||||
state.profile_index + 1,
|
||||
options.profiles.len()
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
fn runtime_label(runtime: &BackendWorkerLaunchRuntimeOption) -> String {
|
||||
let availability = if !runtime.worker_creation_available {
|
||||
"unavailable"
|
||||
} else if runtime.working_directory_required {
|
||||
"workdir required"
|
||||
} else {
|
||||
"no workdir"
|
||||
};
|
||||
format!(
|
||||
"{} [{}] — {availability}",
|
||||
runtime.display_name, runtime.runtime_id
|
||||
)
|
||||
}
|
||||
|
||||
fn runtime_supports_workdirless_creation(runtime: &BackendWorkerLaunchRuntimeOption) -> bool {
|
||||
runtime.worker_creation_available && !runtime.working_directory_required
|
||||
}
|
||||
|
||||
fn cycle_index(current: usize, len: usize, delta: isize) -> usize {
|
||||
if len == 0 {
|
||||
return 0;
|
||||
}
|
||||
(current as isize + delta).rem_euclid(len as isize) as usize
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use client::{BackendDiagnostic, BackendWorkerLaunchOptions};
|
||||
|
||||
fn options() -> BackendWorkerLaunchOptions {
|
||||
BackendWorkerLaunchOptions {
|
||||
workspace_id: "workspace-1".to_string(),
|
||||
runtimes: vec![
|
||||
BackendWorkerLaunchRuntimeOption {
|
||||
runtime_id: "external".to_string(),
|
||||
display_name: "External".to_string(),
|
||||
built_in: false,
|
||||
worker_creation_available: true,
|
||||
working_directory_required: true,
|
||||
status: "online".to_string(),
|
||||
diagnostics: Vec::new(),
|
||||
},
|
||||
BackendWorkerLaunchRuntimeOption {
|
||||
runtime_id: "embedded".to_string(),
|
||||
display_name: "Embedded".to_string(),
|
||||
built_in: true,
|
||||
worker_creation_available: true,
|
||||
working_directory_required: false,
|
||||
status: "online".to_string(),
|
||||
diagnostics: Vec::new(),
|
||||
},
|
||||
],
|
||||
profiles: vec![
|
||||
BackendWorkerLaunchProfileCandidate {
|
||||
id: "builtin:default".to_string(),
|
||||
label: "Default".to_string(),
|
||||
description: String::new(),
|
||||
},
|
||||
BackendWorkerLaunchProfileCandidate {
|
||||
id: "builtin:coder".to_string(),
|
||||
label: "Coder".to_string(),
|
||||
description: "Ticket implementation".to_string(),
|
||||
},
|
||||
],
|
||||
default_profile: Some("builtin:coder".to_string()),
|
||||
repositories: Vec::new(),
|
||||
working_directories: Vec::new(),
|
||||
diagnostics: Vec::<BackendDiagnostic>::new(),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn defaults_to_workdirless_runtime_and_backend_default_profile() {
|
||||
let options = options();
|
||||
let state = FormState::new(&options);
|
||||
assert_eq!(
|
||||
state.current_runtime(&options).unwrap().runtime_id,
|
||||
"embedded"
|
||||
);
|
||||
assert_eq!(state.current_profile(&options).unwrap().id, "builtin:coder");
|
||||
assert_eq!(state.display_name, "Worker");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn workdir_required_runtime_cannot_be_submitted() {
|
||||
let options = options();
|
||||
let mut state = FormState::new(&options);
|
||||
state.runtime_index = 0;
|
||||
assert_eq!(state.submit(&options), None);
|
||||
assert!(state.status.contains("requires a workdir"));
|
||||
assert_eq!(state.field, Field::Runtime);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn selection_builds_workdirless_create_request() {
|
||||
let request = request_from_selection(Selection {
|
||||
runtime_id: "embedded".to_string(),
|
||||
display_name: "Coder one".to_string(),
|
||||
profile: "builtin:coder".to_string(),
|
||||
});
|
||||
assert_eq!(request.runtime_id, "embedded");
|
||||
assert_eq!(request.display_name, "Coder one");
|
||||
assert_eq!(request.profile.as_deref(), Some("builtin:coder"));
|
||||
assert!(request.initial_submit.is_empty());
|
||||
assert!(request.working_directory.is_none());
|
||||
assert!(request.ticket_assignment.is_none());
|
||||
}
|
||||
}
|
||||
@@ -3,19 +3,21 @@ use std::io;
|
||||
use std::time::Duration;
|
||||
|
||||
use client::{
|
||||
BackendRuntimeListTarget, BackendWorkerSummary, list_backend_stopped_workers,
|
||||
list_backend_workers, restore_backend_worker,
|
||||
BackendRuntimeListTarget, BackendWorkerOperationState, BackendWorkerRestoreResponse,
|
||||
BackendWorkerSummary, list_backend_stopped_workers, 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 unicode_width::UnicodeWidthStr;
|
||||
|
||||
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;
|
||||
@@ -83,17 +85,20 @@ pub(crate) async fn run(
|
||||
let restore_target = target
|
||||
.runtime_target(selected.runtime_id.clone(), selected.worker_id.clone())
|
||||
.map_err(|error| io::Error::other(error.to_string()))?;
|
||||
restore_backend_worker(&restore_target)
|
||||
let restore = restore_backend_worker(&restore_target)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
io::Error::other(format!(
|
||||
"failed to restore Backend worker {}/{}: {error}",
|
||||
selected.runtime_id, selected.worker_id
|
||||
))
|
||||
})?
|
||||
.result
|
||||
.worker
|
||||
.unwrap_or(selected)
|
||||
})?;
|
||||
restored_worker(restore).map_err(|error| {
|
||||
io::Error::other(format!(
|
||||
"failed to restore Backend worker {}/{}: {error}",
|
||||
selected.runtime_id, selected.worker_id
|
||||
))
|
||||
})?
|
||||
} else {
|
||||
selected
|
||||
};
|
||||
@@ -104,6 +109,33 @@ pub(crate) async fn run(
|
||||
}
|
||||
}
|
||||
|
||||
fn restored_worker(response: BackendWorkerRestoreResponse) -> Result<BackendWorkerSummary, String> {
|
||||
if response.result.state != BackendWorkerOperationState::Accepted {
|
||||
let diagnostics = response
|
||||
.result
|
||||
.diagnostics
|
||||
.iter()
|
||||
.map(|diagnostic| format!("{}: {}", diagnostic.code, diagnostic.message))
|
||||
.collect::<Vec<_>>()
|
||||
.join("; ");
|
||||
let state = match response.result.state {
|
||||
BackendWorkerOperationState::Accepted => unreachable!(),
|
||||
BackendWorkerOperationState::Rejected => "rejected",
|
||||
BackendWorkerOperationState::Unsupported => "unsupported",
|
||||
};
|
||||
return Err(if diagnostics.is_empty() {
|
||||
format!("restore was {state} without a diagnostic")
|
||||
} else {
|
||||
format!("restore was {state}: {diagnostics}")
|
||||
});
|
||||
}
|
||||
|
||||
response
|
||||
.result
|
||||
.worker
|
||||
.ok_or_else(|| "restore was accepted without a Worker snapshot".to_string())
|
||||
}
|
||||
|
||||
fn dedup_workers(workers: &mut Vec<BackendWorkerSummary>) {
|
||||
let mut seen = std::collections::HashSet::new();
|
||||
workers.retain(|worker| seen.insert((worker.runtime_id.clone(), worker.worker_id.clone())));
|
||||
@@ -127,31 +159,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 +217,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,
|
||||
@@ -255,9 +267,10 @@ fn draw(frame: &mut Frame<'_>, state: &BackendWorkerPickerState) {
|
||||
layout[0],
|
||||
);
|
||||
|
||||
let column_widths = WorkerColumnWidths::from_workers(&state.workers);
|
||||
for (i, worker) in state.workers.iter().enumerate() {
|
||||
frame.render_widget(
|
||||
Paragraph::new(row_line(worker, i == state.selected)),
|
||||
Paragraph::new(row_line(worker, &column_widths, i == state.selected)),
|
||||
layout[i + 1],
|
||||
);
|
||||
}
|
||||
@@ -292,7 +305,28 @@ fn picker_title(target: &BackendRuntimeListTarget) -> String {
|
||||
format!("backend workers workspace: {workspace} runtime: {runtime}")
|
||||
}
|
||||
|
||||
fn row_line(worker: &BackendWorkerSummary, selected: bool) -> Line<'static> {
|
||||
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
|
||||
struct WorkerColumnWidths {
|
||||
identity: usize,
|
||||
name: usize,
|
||||
state: usize,
|
||||
}
|
||||
|
||||
impl WorkerColumnWidths {
|
||||
fn from_workers(workers: &[BackendWorkerSummary]) -> Self {
|
||||
workers.iter().fold(Self::default(), |widths, worker| Self {
|
||||
identity: widths.identity.max(text_width(&short_worker_id(worker))),
|
||||
name: widths.name.max(text_width(worker_name(worker))),
|
||||
state: widths.state.max(text_width(&worker_state(worker))),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn row_line(
|
||||
worker: &BackendWorkerSummary,
|
||||
widths: &WorkerColumnWidths,
|
||||
selected: bool,
|
||||
) -> Line<'static> {
|
||||
let marker = if selected { "▶ " } else { " " };
|
||||
let id_style = if selected {
|
||||
Style::default()
|
||||
@@ -301,42 +335,70 @@ fn row_line(worker: &BackendWorkerSummary, selected: bool) -> Line<'static> {
|
||||
} else {
|
||||
Style::default().fg(Color::Cyan)
|
||||
};
|
||||
let preview_style = if selected {
|
||||
let name_style = if selected {
|
||||
Style::default().fg(Color::White)
|
||||
} else {
|
||||
Style::default().fg(Color::DarkGray)
|
||||
};
|
||||
|
||||
let label = if worker.label.is_empty() {
|
||||
worker.worker_id.as_str()
|
||||
} else {
|
||||
worker.label.as_str()
|
||||
};
|
||||
let profile = worker.profile.as_deref().unwrap_or("-");
|
||||
|
||||
Line::from(vec![
|
||||
Span::raw(marker),
|
||||
Span::styled(short_worker_id(worker), id_style),
|
||||
Span::raw(" "),
|
||||
Span::styled(
|
||||
format!("[{}]", worker.state),
|
||||
state_style(worker.state.as_str()),
|
||||
pad_column(&short_worker_id(worker), widths.identity),
|
||||
id_style,
|
||||
),
|
||||
Span::raw(" "),
|
||||
Span::styled(pad_column(worker_name(worker), widths.name), name_style),
|
||||
Span::raw(" "),
|
||||
Span::styled(
|
||||
format!("profile:{profile}"),
|
||||
Style::default().fg(Color::DarkGray),
|
||||
pad_column(&worker_state(worker), widths.state),
|
||||
state_style(worker_state_label(worker)),
|
||||
),
|
||||
Span::raw(" "),
|
||||
Span::styled(
|
||||
working_directory_text(worker),
|
||||
Style::default().fg(Color::DarkGray),
|
||||
),
|
||||
Span::raw(" "),
|
||||
Span::styled(label.to_string(), preview_style),
|
||||
])
|
||||
}
|
||||
|
||||
fn worker_name(worker: &BackendWorkerSummary) -> &str {
|
||||
if !worker.label.is_empty() {
|
||||
worker.label.as_str()
|
||||
} else if !worker.display_name.is_empty() {
|
||||
worker.display_name.as_str()
|
||||
} else {
|
||||
worker.worker_id.as_str()
|
||||
}
|
||||
}
|
||||
|
||||
fn worker_state_label(worker: &BackendWorkerSummary) -> &str {
|
||||
match worker.worker_state.as_ref().map(|state| &state.state) {
|
||||
Some(protocol::WorkerState::Idle) => "idle",
|
||||
Some(protocol::WorkerState::Busy(protocol::WorkerBusyState::Run(
|
||||
protocol::WorkerRunState::Paused,
|
||||
))) => "paused",
|
||||
Some(protocol::WorkerState::Busy(_)) => "running",
|
||||
None if worker.state == "stopped" => "stopped",
|
||||
None => "unknown",
|
||||
}
|
||||
}
|
||||
|
||||
fn worker_state(worker: &BackendWorkerSummary) -> String {
|
||||
format!("[{}]", worker_state_label(worker))
|
||||
}
|
||||
|
||||
fn text_width(value: &str) -> usize {
|
||||
UnicodeWidthStr::width(value)
|
||||
}
|
||||
|
||||
fn pad_column(value: &str, width: usize) -> String {
|
||||
format!(
|
||||
"{value}{}",
|
||||
" ".repeat(width.saturating_sub(text_width(value)))
|
||||
)
|
||||
}
|
||||
|
||||
fn state_style(state: &str) -> Style {
|
||||
match state {
|
||||
"running" | "idle" | "active" => Style::default()
|
||||
@@ -367,18 +429,15 @@ fn working_directory_text(worker: &BackendWorkerSummary) -> String {
|
||||
let Some(wd) = worker.working_directory.as_ref() else {
|
||||
return "wd:—".to_string();
|
||||
};
|
||||
let cleanliness = wd.cleanliness.as_deref().unwrap_or("unknown");
|
||||
format!(
|
||||
"wd:{}:{} {} {}",
|
||||
wd.repository_id, wd.working_directory_id, wd.status, cleanliness
|
||||
)
|
||||
format!("wd:{}・{}", wd.repository_key, wd.working_directory_id)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use client::{
|
||||
BackendWorkerCapabilitySummary, BackendWorkerImplementationSummary,
|
||||
BackendDiagnostic, BackendDiagnosticSeverity, BackendWorkerCapabilitySummary,
|
||||
BackendWorkerImplementationSummary, BackendWorkerRestoreResult,
|
||||
BackendWorkerWorkspaceSummary,
|
||||
};
|
||||
|
||||
@@ -398,7 +457,15 @@ mod tests {
|
||||
identity: "ws".to_string(),
|
||||
workspace_id: Some("ws".to_string()),
|
||||
},
|
||||
state: "running".to_string(),
|
||||
state: "idle".to_string(),
|
||||
worker_state: Some(protocol::WorkerStateSnapshot {
|
||||
execution_generation: 1,
|
||||
revision: 1,
|
||||
state: protocol::WorkerState::Busy(protocol::WorkerBusyState::Run(
|
||||
protocol::WorkerRunState::Running,
|
||||
)),
|
||||
last_command_id: 0,
|
||||
}),
|
||||
last_seen_at: None,
|
||||
pinned: false,
|
||||
retention_state: String::new(),
|
||||
@@ -415,18 +482,159 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn worker_row_matches_inline_picker_shape() {
|
||||
let row = row_line(&worker("runtime-a", "worker-b", Some("default")), true);
|
||||
let text = row
|
||||
fn row_text(worker: &BackendWorkerSummary, widths: &WorkerColumnWidths) -> String {
|
||||
row_line(worker, widths, false)
|
||||
.spans
|
||||
.into_iter()
|
||||
.map(|span| span.content)
|
||||
.collect::<String>();
|
||||
assert!(text.starts_with("▶ W-1"));
|
||||
assert!(text.contains("[running]"));
|
||||
assert!(text.contains("profile:default"));
|
||||
assert!(text.contains("wd:—"));
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn display_column(text: &str, value: &str) -> usize {
|
||||
let byte_offset = text.find(value).expect("value in rendered row");
|
||||
text_width(&text[..byte_offset])
|
||||
}
|
||||
|
||||
fn restore_response(
|
||||
state: BackendWorkerOperationState,
|
||||
worker: Option<BackendWorkerSummary>,
|
||||
diagnostics: Vec<BackendDiagnostic>,
|
||||
) -> BackendWorkerRestoreResponse {
|
||||
BackendWorkerRestoreResponse {
|
||||
workspace_id: "workspace-a".to_string(),
|
||||
runtime_id: "runtime-a".to_string(),
|
||||
worker_id: "worker-a".to_string(),
|
||||
result: BackendWorkerRestoreResult {
|
||||
state,
|
||||
worker,
|
||||
diagnostics,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejected_restore_surfaces_diagnostic_instead_of_attaching_selected_worker() {
|
||||
let error = restored_worker(restore_response(
|
||||
BackendWorkerOperationState::Rejected,
|
||||
None,
|
||||
vec![BackendDiagnostic {
|
||||
code: "working_directory_not_found".to_string(),
|
||||
severity: BackendDiagnosticSeverity::Error,
|
||||
message: "working directory was not found".to_string(),
|
||||
}],
|
||||
))
|
||||
.expect_err("rejected restore must not produce a Worker to attach");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
"restore was rejected: working_directory_not_found: working directory was not found"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accepted_restore_requires_returned_worker_snapshot() {
|
||||
let error = restored_worker(restore_response(
|
||||
BackendWorkerOperationState::Accepted,
|
||||
None,
|
||||
Vec::new(),
|
||||
))
|
||||
.expect_err("accepted restore without a Worker must not attach the stale selection");
|
||||
|
||||
assert_eq!(error, "restore was accepted without a Worker snapshot");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accepted_restore_returns_authoritative_worker_snapshot() {
|
||||
let worker = worker("runtime-a", "worker-a", Some("builtin:companion"));
|
||||
let restored = restored_worker(restore_response(
|
||||
BackendWorkerOperationState::Accepted,
|
||||
Some(worker.clone()),
|
||||
Vec::new(),
|
||||
))
|
||||
.expect("accepted restore should return its Worker snapshot");
|
||||
|
||||
assert_eq!(restored, worker);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn worker_row_orders_and_simplifies_columns() {
|
||||
let mut worker = worker("runtime-a", "worker-b", Some("builtin:coder"));
|
||||
worker.resource_key = "W-90".to_string();
|
||||
worker.display_name = "Coder".to_string();
|
||||
worker.label = "Coder · T-585".to_string();
|
||||
worker.state = "stopped".to_string();
|
||||
worker.worker_state = None;
|
||||
worker.working_directory = Some(
|
||||
serde_json::from_value(serde_json::json!({
|
||||
"working_directory_id": "001a06a9f0202000000",
|
||||
"repository_key": "main",
|
||||
"materializer_kind": "runtime_git_clone",
|
||||
"status": "active",
|
||||
"cleanliness": "clean"
|
||||
}))
|
||||
.unwrap(),
|
||||
);
|
||||
let widths = WorkerColumnWidths::from_workers(std::slice::from_ref(&worker));
|
||||
let text = row_text(&worker, &widths);
|
||||
|
||||
assert_eq!(
|
||||
text,
|
||||
" W-90 Coder · T-585 [stopped] wd:main・001a06a9f0202000000"
|
||||
);
|
||||
assert!(!text.contains("profile:"));
|
||||
assert!(!text.contains("active clean"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn worker_rows_align_identity_name_state_and_workdir_columns() {
|
||||
let mut short = worker("runtime-a", "worker-a", None);
|
||||
short.resource_key = "W-2".to_string();
|
||||
short.label = "Coder".to_string();
|
||||
short.display_name = short.label.clone();
|
||||
short.state = "idle".to_string();
|
||||
short.worker_state = Some(protocol::WorkerStateSnapshot {
|
||||
execution_generation: 1,
|
||||
revision: 2,
|
||||
state: protocol::WorkerState::Idle,
|
||||
last_command_id: 0,
|
||||
});
|
||||
|
||||
let mut long = worker("runtime-a", "worker-b", None);
|
||||
long.resource_key = "W-100".to_string();
|
||||
long.label = "Longer worker · T-9".to_string();
|
||||
long.display_name = long.label.clone();
|
||||
long.state = "stopped".to_string();
|
||||
long.worker_state = None;
|
||||
|
||||
for worker in [&mut short, &mut long] {
|
||||
worker.working_directory = Some(
|
||||
serde_json::from_value(serde_json::json!({
|
||||
"working_directory_id": "workdir-1",
|
||||
"repository_key": "main",
|
||||
"materializer_kind": "runtime_git_clone",
|
||||
"status": "active"
|
||||
}))
|
||||
.unwrap(),
|
||||
);
|
||||
}
|
||||
|
||||
let workers = vec![short, long];
|
||||
let widths = WorkerColumnWidths::from_workers(&workers);
|
||||
let first = row_text(&workers[0], &widths);
|
||||
let second = row_text(&workers[1], &widths);
|
||||
|
||||
assert_eq!(
|
||||
display_column(&first, "Coder"),
|
||||
display_column(&second, "Longer")
|
||||
);
|
||||
assert_eq!(
|
||||
display_column(&first, "[idle]"),
|
||||
display_column(&second, "[stopped]")
|
||||
);
|
||||
assert_eq!(
|
||||
display_column(&first, "wd:main"),
|
||||
display_column(&second, "wd:main")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -189,7 +189,7 @@ fn prompt_create_request_inner() -> PickerResult<Option<CreateBackendWorkspaceRe
|
||||
println!("Repository path/URI is required.");
|
||||
return Ok(None);
|
||||
}
|
||||
let repository_name = prompt_line("Repository display name [Main]: ")?;
|
||||
let repository_key = prompt_line("Repository key [main]: ")?;
|
||||
let default_ref = prompt_line("Default ref [repository default]: ")?;
|
||||
let operation_key = format!(
|
||||
"tui-workspace-create-{}-{}",
|
||||
@@ -204,11 +204,11 @@ fn prompt_create_request_inner() -> PickerResult<Option<CreateBackendWorkspaceRe
|
||||
display_name,
|
||||
repository: CreateBackendWorkspaceRepository {
|
||||
uri,
|
||||
display_name: Some(if repository_name.is_empty() {
|
||||
"Main".to_string()
|
||||
repository_key: if repository_key.is_empty() {
|
||||
"main".to_string()
|
||||
} else {
|
||||
repository_name
|
||||
}),
|
||||
repository_key
|
||||
},
|
||||
default_ref: (!default_ref.is_empty()).then_some(default_ref),
|
||||
},
|
||||
}))
|
||||
|
||||
@@ -409,7 +409,12 @@ fn compact_command(invocation: CommandInvocation<'_>) -> CommandExecution {
|
||||
let _ = invocation.environment;
|
||||
let _ = invocation.args.raw();
|
||||
CommandExecution {
|
||||
method: Some(Method::Compact),
|
||||
method: Some(Method::Compact {
|
||||
command: protocol::WorkerCommandEnvelope::for_snapshot(
|
||||
0,
|
||||
&protocol::WorkerStateSnapshot::initial(1),
|
||||
),
|
||||
}),
|
||||
diagnostics: vec![CommandDiagnostic::new("compact requested")],
|
||||
exit_command_mode: true,
|
||||
clear_input: true,
|
||||
@@ -483,7 +488,7 @@ mod tests {
|
||||
fn compact_command_returns_compact_method_not_run() {
|
||||
let registry = CommandRegistry::builtins();
|
||||
let result = registry.dispatch("compact", &env());
|
||||
assert!(matches!(result.method, Some(Method::Compact)));
|
||||
assert!(matches!(result.method, Some(Method::Compact { .. })));
|
||||
assert!(result.exit_command_mode);
|
||||
assert!(result.clear_input);
|
||||
assert!(result.diagnostics[0].message.contains("compact requested"));
|
||||
|
||||
+771
-129
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<()> {
|
||||
|
||||
+24
-11
@@ -1,5 +1,6 @@
|
||||
mod app;
|
||||
mod backend_dashboard;
|
||||
mod backend_spawn;
|
||||
mod backend_worker_picker;
|
||||
mod backend_workspace_picker;
|
||||
mod block;
|
||||
@@ -10,12 +11,14 @@ mod composer_keys;
|
||||
mod console;
|
||||
#[cfg(feature = "e2e-test")]
|
||||
mod e2e_observer;
|
||||
mod inline_terminal;
|
||||
mod input;
|
||||
pub mod keys;
|
||||
mod markdown;
|
||||
mod scroll;
|
||||
pub mod setup_model;
|
||||
mod standalone_picker;
|
||||
mod standalone_spawn;
|
||||
mod task;
|
||||
mod text_selection;
|
||||
mod tool;
|
||||
@@ -49,6 +52,8 @@ pub enum LaunchMode {
|
||||
/// 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 },
|
||||
/// Create one Backend Worker and attach to it.
|
||||
BackendSpawn,
|
||||
/// List Backend Workers and attach to the selected Worker.
|
||||
Workers {
|
||||
runtime_id: Option<String>,
|
||||
@@ -136,17 +141,21 @@ pub async fn launch(options: LaunchOptions) -> ExitCode {
|
||||
LaunchMode::Spawn {
|
||||
worker_name,
|
||||
profile,
|
||||
} => match target.spawn_worker() {
|
||||
Ok(spawn) => {
|
||||
console::run_standalone(
|
||||
workspace_root.clone(),
|
||||
spawn.state_dir,
|
||||
worker_name,
|
||||
profile,
|
||||
)
|
||||
.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::StandaloneResume { include_all } => {
|
||||
match standalone_picker::pick(target.as_ref(), include_all) {
|
||||
@@ -155,6 +164,10 @@ pub async fn launch(options: LaunchOptions) -> ExitCode {
|
||||
Err(error) => Err(Box::new(error) as Box<dyn std::error::Error>),
|
||||
}
|
||||
}
|
||||
LaunchMode::BackendSpawn => match target.launch_backend_worker() {
|
||||
Ok(launch) => backend_spawn::run(launch.target).await,
|
||||
Err(e) => Err(Box::new(e) as Box<dyn std::error::Error>),
|
||||
},
|
||||
LaunchMode::Workers {
|
||||
runtime_id,
|
||||
include_stopped,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -3,15 +3,14 @@ use std::time::Duration;
|
||||
|
||||
use client::{StandaloneWorkerListIntent, StandaloneWorkerResumeIntent, Target};
|
||||
use crossterm::event::{self, Event as TermEvent, KeyCode, KeyEventKind, KeyModifiers};
|
||||
use ratatui::Terminal;
|
||||
use ratatui::backend::CrosstermBackend;
|
||||
use ratatui::layout::{Constraint, Layout};
|
||||
use ratatui::prelude::{Color, Line, Modifier, Span, Style};
|
||||
use ratatui::widgets::Paragraph;
|
||||
use ratatui::{TerminalOptions, Viewport};
|
||||
use standalone::{StandaloneListScope, StandaloneWorkerRecord, StandaloneWorkerStore};
|
||||
use thiserror::Error;
|
||||
|
||||
use crate::inline_terminal::with_inline_terminal;
|
||||
|
||||
const LIMIT: usize = 100;
|
||||
|
||||
pub(crate) fn pick(
|
||||
@@ -57,41 +56,36 @@ 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);
|
||||
let mut terminal = Terminal::with_options(
|
||||
CrosstermBackend::new(io::stdout()),
|
||||
TerminalOptions {
|
||||
viewport: Viewport::Inline(height),
|
||||
},
|
||||
)
|
||||
.map_err(StandalonePickerError::Io)?;
|
||||
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);
|
||||
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;
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if !ctrl => {
|
||||
selected = (selected + 1).min(records.len() - 1);
|
||||
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),
|
||||
_ => {}
|
||||
}
|
||||
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) {
|
||||
@@ -145,11 +139,11 @@ pub(crate) enum StandalonePickerError {
|
||||
#[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"
|
||||
"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(#[source] io::Error),
|
||||
Io(#[from] io::Error),
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
+60
-14
@@ -1296,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),
|
||||
@@ -1314,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(),
|
||||
@@ -1842,7 +1880,7 @@ fn actionbar_left_item(app: &App, now: Instant) -> Option<(String, Style)> {
|
||||
}
|
||||
if app.queued_input_count() > 0 {
|
||||
return Some((
|
||||
"Alt-q edit queued Alt-c clear queued".to_string(),
|
||||
"Alt-n notify Alt-q continue Alt-d cancel queued Alt-c clear queued".to_string(),
|
||||
Style::default().fg(Color::DarkGray),
|
||||
));
|
||||
}
|
||||
@@ -2098,9 +2136,25 @@ mod tests {
|
||||
use super::*;
|
||||
use crate::app::{ActionbarNoticeLevel, ActionbarNoticeSource, App};
|
||||
use crate::block::{ToolCallBlock, ToolCallState};
|
||||
use protocol::WorkerStatus;
|
||||
use protocol::Event;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
fn set_pending_submission(app: &mut App, id: &str) {
|
||||
app.handle_worker_event(Event::PendingSubmissionsChanged {
|
||||
pending: protocol::PendingSubmissionsSnapshot {
|
||||
revision: 1,
|
||||
notification_count: 0,
|
||||
head_id: Some(id.into()),
|
||||
submissions: vec![protocol::PendingSubmissionSummary {
|
||||
submission_id: id.into(),
|
||||
accepted_at_ms: 1,
|
||||
segment_count: 1,
|
||||
byte_len: 1,
|
||||
}],
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn run_status_line_matches_console_metrics_and_spinner_frame() {
|
||||
let now = Instant::now();
|
||||
@@ -2213,15 +2267,11 @@ mod tests {
|
||||
#[test]
|
||||
fn queue_status_text_includes_count_and_preview() {
|
||||
let mut app = App::new("test".into());
|
||||
app.set_worker_status(WorkerStatus::Running);
|
||||
for c in "queued preview".chars() {
|
||||
app.insert_char(c);
|
||||
}
|
||||
assert!(app.submit_input().is_none());
|
||||
set_pending_submission(&mut app, "submission-1");
|
||||
|
||||
assert_eq!(
|
||||
queue_status_text(&app),
|
||||
Some("queued: 1 — queued preview".to_string())
|
||||
Some("queued: 1 — submission-1".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
@@ -2251,14 +2301,10 @@ mod tests {
|
||||
Some("Worker keeps running. Press Ctrl-C again to exit TUI.".into())
|
||||
);
|
||||
|
||||
app.set_worker_status(WorkerStatus::Running);
|
||||
for c in "queued turn".chars() {
|
||||
app.insert_char(c);
|
||||
}
|
||||
assert!(app.submit_input().is_none());
|
||||
set_pending_submission(&mut app, "submission-1");
|
||||
assert_eq!(
|
||||
actionbar_left_item(&app, now).map(|(text, _)| text),
|
||||
Some("Alt-q edit queued Alt-c clear queued".into())
|
||||
Some("Alt-n notify Alt-q continue Alt-d cancel queued Alt-c clear queued".into())
|
||||
);
|
||||
|
||||
app.enter_command_mode();
|
||||
|
||||
@@ -14,10 +14,12 @@ fs-operation.workspace = true
|
||||
manifest.workspace = true
|
||||
reqwest = { version = "0.13", default-features = false, features = ["json", "rustls"], optional = true }
|
||||
serde = { workspace = true, features = ["derive"] }
|
||||
serde_json.workspace = true
|
||||
sha2.workspace = true
|
||||
tempfile.workspace = true
|
||||
thiserror.workspace = true
|
||||
tokio = { workspace = true, features = ["process", "rt", "sync", "time"] }
|
||||
workspace-api = { workspace = true }
|
||||
|
||||
[dev-dependencies]
|
||||
serde_json.workspace = true
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+204
-60
@@ -11,7 +11,8 @@ use crate::{
|
||||
CommandHandle, CommandOutput, CommandOutputRequest, CommandRequest, CommandStatus, EditRequest,
|
||||
EditResult, GlobRequest, GlobResult, GrepRequest, GrepResult, ListRequest, ListResult,
|
||||
ReadRequest, ReadResult, StatRequest, StatResult, WorkdirError, WorkdirId,
|
||||
WorkdirSessionCapabilities, WriteRequest, WriteResult,
|
||||
WorkdirScopeAuthorizationRequest, WorkdirScopeOverlapRequest, WorkdirSessionCapabilities,
|
||||
WriteRequest, WriteResult,
|
||||
};
|
||||
|
||||
/// Opaque Runtime-owned identifier for one ephemeral Workdir session.
|
||||
@@ -55,6 +56,8 @@ pub struct OpenWorkdirSessionResponse {
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(tag = "operation", content = "request", rename_all = "snake_case")]
|
||||
pub enum WorkdirSessionOperation {
|
||||
AuthorizeScope(WorkdirScopeAuthorizationRequest),
|
||||
ScopeRulesOverlap(WorkdirScopeOverlapRequest),
|
||||
Stat(StatRequest),
|
||||
Read(ReadRequest),
|
||||
Write(WriteRequest),
|
||||
@@ -68,12 +71,10 @@ pub enum WorkdirSessionOperation {
|
||||
CommandCancel(CommandHandle),
|
||||
}
|
||||
|
||||
/// Wire envelope for an operation and its optional provider-enforced child scope.
|
||||
/// Wire envelope for one provider operation.
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct WorkdirSessionOperationRequest {
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
pub delegations: Vec<crate::WorkdirDelegationRequest>,
|
||||
pub operation: WorkdirSessionOperation,
|
||||
}
|
||||
|
||||
@@ -81,6 +82,8 @@ pub struct WorkdirSessionOperationRequest {
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(tag = "operation", content = "result", rename_all = "snake_case")]
|
||||
pub enum WorkdirSessionOperationResult {
|
||||
AuthorizeScope,
|
||||
ScopeRulesOverlap { overlaps: bool },
|
||||
Stat(StatResult),
|
||||
Read(ReadResult),
|
||||
Write(WriteResult),
|
||||
@@ -102,8 +105,18 @@ pub enum WorkdirTransportErrorCode {
|
||||
Conflict,
|
||||
Unsupported,
|
||||
InvalidRequest,
|
||||
Denied,
|
||||
OutOfScope,
|
||||
SymlinkOutOfScope,
|
||||
BrokenSymlink,
|
||||
SymlinkTargetIsDirectory,
|
||||
ReadOnly,
|
||||
IsDirectory,
|
||||
SymlinkDirectoryNotTraversed,
|
||||
UnknownCommand,
|
||||
Unavailable,
|
||||
Io,
|
||||
Transport,
|
||||
Internal,
|
||||
}
|
||||
|
||||
@@ -114,8 +127,18 @@ impl WorkdirTransportErrorCode {
|
||||
Self::Conflict => "conflict",
|
||||
Self::Unsupported => "unsupported",
|
||||
Self::InvalidRequest => "invalid_request",
|
||||
Self::Denied => "denied",
|
||||
Self::OutOfScope => "out_of_scope",
|
||||
Self::SymlinkOutOfScope => "symlink_out_of_scope",
|
||||
Self::BrokenSymlink => "broken_symlink",
|
||||
Self::SymlinkTargetIsDirectory => "symlink_target_is_directory",
|
||||
Self::ReadOnly => "read_only",
|
||||
Self::IsDirectory => "is_directory",
|
||||
Self::SymlinkDirectoryNotTraversed => "symlink_directory_not_traversed",
|
||||
Self::UnknownCommand => "unknown_command",
|
||||
Self::Unavailable => "unavailable",
|
||||
Self::Io => "io",
|
||||
Self::Transport => "transport",
|
||||
Self::Internal => "internal",
|
||||
}
|
||||
}
|
||||
@@ -125,9 +148,16 @@ impl WorkdirTransportErrorCode {
|
||||
match self {
|
||||
Self::NotFound | Self::UnknownCommand => 404,
|
||||
Self::Conflict => 409,
|
||||
Self::Unsupported | Self::InvalidRequest => 400,
|
||||
Self::Denied | Self::OutOfScope | Self::SymlinkOutOfScope | Self::ReadOnly => 403,
|
||||
Self::Unsupported
|
||||
| Self::InvalidRequest
|
||||
| Self::BrokenSymlink
|
||||
| Self::SymlinkTargetIsDirectory
|
||||
| Self::IsDirectory
|
||||
| Self::SymlinkDirectoryNotTraversed => 400,
|
||||
Self::Unavailable => 503,
|
||||
Self::Internal => 500,
|
||||
Self::Io | Self::Internal => 500,
|
||||
Self::Transport => 502,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -160,8 +190,41 @@ impl WorkdirTransportError {
|
||||
WorkdirError::Unavailable(_) | WorkdirError::SessionClosed => {
|
||||
(Code::Unavailable, "Workdir session is unavailable")
|
||||
}
|
||||
WorkdirError::Denied(_) => (Code::InvalidRequest, "Workdir operation was denied"),
|
||||
WorkdirError::Transport(_) => (Code::Internal, "Workdir transport failed"),
|
||||
WorkdirError::Denied(_) => (Code::Denied, "Workdir operation was denied"),
|
||||
WorkdirError::OutOfScope(_) => (Code::OutOfScope, "Workdir path is out of scope"),
|
||||
WorkdirError::SymlinkOutOfScope { .. } => (
|
||||
Code::SymlinkOutOfScope,
|
||||
"Workdir symlink target is out of scope",
|
||||
),
|
||||
WorkdirError::BrokenSymlink { .. } => {
|
||||
(Code::BrokenSymlink, "Workdir symlink target does not exist")
|
||||
}
|
||||
WorkdirError::SymlinkTargetIsDirectory { .. } => (
|
||||
Code::SymlinkTargetIsDirectory,
|
||||
"Workdir symlink target is a directory",
|
||||
),
|
||||
WorkdirError::ReadOnly(_) => (Code::ReadOnly, "Workdir path is read-only"),
|
||||
WorkdirError::IsDirectory(_) => (Code::IsDirectory, "Workdir path is a directory"),
|
||||
WorkdirError::SymlinkDirectoryNotTraversed { .. } => (
|
||||
Code::SymlinkDirectoryNotTraversed,
|
||||
"Workdir symlink directory was not traversed",
|
||||
),
|
||||
WorkdirError::Io { source, .. } => match source.kind() {
|
||||
std::io::ErrorKind::NotFound => (Code::NotFound, "Workdir path was not found"),
|
||||
std::io::ErrorKind::PermissionDenied => {
|
||||
(Code::Denied, "Workdir operation was denied")
|
||||
}
|
||||
std::io::ErrorKind::AlreadyExists => {
|
||||
(Code::Conflict, "Workdir resource already exists")
|
||||
}
|
||||
std::io::ErrorKind::InvalidInput | std::io::ErrorKind::InvalidData => {
|
||||
(Code::InvalidRequest, "Workdir operation request is invalid")
|
||||
}
|
||||
std::io::ErrorKind::TimedOut => (Code::Unavailable, "Workdir operation timed out"),
|
||||
_ => (Code::Io, "Workdir I/O operation failed"),
|
||||
},
|
||||
WorkdirError::OperationFailed => (Code::Internal, "Workdir operation failed"),
|
||||
WorkdirError::Transport(_) => (Code::Transport, "Workdir transport failed"),
|
||||
WorkdirError::InvalidPath(_)
|
||||
| WorkdirError::RelativePath(_)
|
||||
| WorkdirError::InvalidGlob(_)
|
||||
@@ -169,14 +232,6 @@ impl WorkdirTransportError {
|
||||
| WorkdirError::InvalidArgument(_) => {
|
||||
(Code::InvalidRequest, "Workdir operation request is invalid")
|
||||
}
|
||||
WorkdirError::OutOfScope(_)
|
||||
| WorkdirError::SymlinkOutOfScope { .. }
|
||||
| WorkdirError::BrokenSymlink { .. }
|
||||
| WorkdirError::SymlinkTargetIsDirectory { .. }
|
||||
| WorkdirError::ReadOnly(_)
|
||||
| WorkdirError::IsDirectory(_)
|
||||
| WorkdirError::SymlinkDirectoryNotTraversed { .. }
|
||||
| WorkdirError::Io { .. } => (Code::Internal, "Workdir operation failed"),
|
||||
};
|
||||
Self {
|
||||
code,
|
||||
@@ -190,10 +245,38 @@ impl WorkdirTransportError {
|
||||
Code::NotFound => WorkdirError::NotFound("<remote>".into()),
|
||||
Code::Conflict => WorkdirError::Conflict(self.message),
|
||||
Code::Unsupported => WorkdirError::UnsupportedOperation(self.message),
|
||||
Code::UnknownCommand => WorkdirError::UnknownCommand("<remote>".to_string()),
|
||||
Code::InvalidRequest => WorkdirError::InvalidArgument(self.message),
|
||||
Code::Denied => WorkdirError::Denied(self.message),
|
||||
Code::OutOfScope => WorkdirError::OutOfScope("<remote>".into()),
|
||||
Code::SymlinkOutOfScope => WorkdirError::SymlinkOutOfScope {
|
||||
path: "<remote>".into(),
|
||||
target: "<remote-target>".into(),
|
||||
required_permission: "requested",
|
||||
},
|
||||
Code::BrokenSymlink => WorkdirError::BrokenSymlink {
|
||||
path: "<remote>".into(),
|
||||
link: "<remote-link>".into(),
|
||||
target: "<remote-target>".into(),
|
||||
},
|
||||
Code::SymlinkTargetIsDirectory => WorkdirError::SymlinkTargetIsDirectory {
|
||||
path: "<remote>".into(),
|
||||
target: "<remote-target>".into(),
|
||||
},
|
||||
Code::ReadOnly => WorkdirError::ReadOnly("<remote>".into()),
|
||||
Code::IsDirectory => WorkdirError::IsDirectory("<remote>".into()),
|
||||
Code::SymlinkDirectoryNotTraversed => WorkdirError::SymlinkDirectoryNotTraversed {
|
||||
tool: "remote operation",
|
||||
path: "<remote>".into(),
|
||||
target: "<remote-target>".into(),
|
||||
},
|
||||
Code::UnknownCommand => WorkdirError::UnknownCommand("<remote>".to_string()),
|
||||
Code::Unavailable => WorkdirError::Unavailable(self.message),
|
||||
Code::Internal => WorkdirError::Transport(self.message),
|
||||
Code::Io => WorkdirError::Io {
|
||||
path: "<remote>".into(),
|
||||
source: std::io::Error::other(self.message),
|
||||
},
|
||||
Code::Transport => WorkdirError::Transport(self.message),
|
||||
Code::Internal => WorkdirError::OperationFailed,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -209,13 +292,18 @@ mod client {
|
||||
use reqwest::{Client, StatusCode, Url};
|
||||
|
||||
use super::*;
|
||||
use crate::{Workdir, WorkdirSession, WorkdirSessionHandle};
|
||||
use crate::{Workdir, WorkdirSession};
|
||||
|
||||
/// Provides a fresh bearer token for each Runtime request. Backend
|
||||
/// implementations can mint short-lived capability tokens without making a
|
||||
/// Worker-bound session expire with the token used to open it.
|
||||
pub trait WorkdirHttpAuthorization: std::fmt::Debug + Send + Sync {
|
||||
fn bearer_token(&self) -> Result<String, WorkdirError>;
|
||||
fn bearer_token(
|
||||
&self,
|
||||
method: &str,
|
||||
path_and_query: &str,
|
||||
body: &[u8],
|
||||
) -> Result<String, WorkdirError>;
|
||||
}
|
||||
|
||||
struct FixedBearerToken(Arc<str>);
|
||||
@@ -227,7 +315,12 @@ mod client {
|
||||
}
|
||||
|
||||
impl WorkdirHttpAuthorization for FixedBearerToken {
|
||||
fn bearer_token(&self) -> Result<String, WorkdirError> {
|
||||
fn bearer_token(
|
||||
&self,
|
||||
_method: &str,
|
||||
_path_and_query: &str,
|
||||
_body: &[u8],
|
||||
) -> Result<String, WorkdirError> {
|
||||
Ok(self.0.to_string())
|
||||
}
|
||||
}
|
||||
@@ -244,7 +337,6 @@ mod client {
|
||||
workdir: Workdir,
|
||||
session_id: WorkdirSessionId,
|
||||
capabilities: WorkdirSessionCapabilities,
|
||||
delegations: Vec<crate::WorkdirDelegationRequest>,
|
||||
closed: AtomicBool,
|
||||
}
|
||||
|
||||
@@ -277,10 +369,14 @@ mod client {
|
||||
&base_url,
|
||||
&["v1", "working-directories", workdir_id.as_str(), "sessions"],
|
||||
)?;
|
||||
let body = serde_json::to_vec(&request)
|
||||
.map_err(|error| WorkdirError::Unavailable(error.to_string()))?;
|
||||
let token = authorization.bearer_token("POST", url.path(), &body)?;
|
||||
let response = client
|
||||
.post(url)
|
||||
.bearer_auth(authorization.bearer_token()?)
|
||||
.json(&request)
|
||||
.bearer_auth(token)
|
||||
.header("content-type", "application/json")
|
||||
.body(body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(http_unavailable)?;
|
||||
@@ -297,7 +393,6 @@ mod client {
|
||||
workdir: Workdir::new(opened.workdir_id.as_str()),
|
||||
session_id: opened.session_id,
|
||||
capabilities: opened.capabilities,
|
||||
delegations: Vec::new(),
|
||||
closed: AtomicBool::new(false),
|
||||
})
|
||||
}
|
||||
@@ -324,15 +419,16 @@ mod client {
|
||||
"operations",
|
||||
],
|
||||
)?;
|
||||
let operation = WorkdirSessionOperationRequest {
|
||||
delegations: self.delegations.clone(),
|
||||
operation,
|
||||
};
|
||||
let operation = WorkdirSessionOperationRequest { operation };
|
||||
let body = serde_json::to_vec(&operation)
|
||||
.map_err(|error| WorkdirError::Unavailable(error.to_string()))?;
|
||||
let token = self.authorization.bearer_token("POST", url.path(), &body)?;
|
||||
let response = self
|
||||
.client
|
||||
.post(url)
|
||||
.bearer_auth(self.authorization.bearer_token()?)
|
||||
.json(&operation)
|
||||
.bearer_auth(token)
|
||||
.header("content-type", "application/json")
|
||||
.body(body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(http_unavailable)?;
|
||||
@@ -356,35 +452,30 @@ mod client {
|
||||
self.capabilities
|
||||
}
|
||||
|
||||
fn transports_delegation_context(&self) -> bool {
|
||||
true
|
||||
async fn authorize_scope_path(
|
||||
&self,
|
||||
request: WorkdirScopeAuthorizationRequest,
|
||||
) -> Result<(), WorkdirError> {
|
||||
match self
|
||||
.operate(WorkdirSessionOperation::AuthorizeScope(request))
|
||||
.await?
|
||||
{
|
||||
WorkdirSessionOperationResult::AuthorizeScope => Ok(()),
|
||||
_ => Err(Self::mismatch("authorize_scope")),
|
||||
}
|
||||
}
|
||||
|
||||
async fn capture_delegation_source(
|
||||
async fn scope_rules_overlap(
|
||||
&self,
|
||||
request: &crate::WorkdirDelegationRequest,
|
||||
) -> Result<WorkdirSessionHandle, WorkdirError> {
|
||||
if self.closed.load(Ordering::Acquire) {
|
||||
return Err(WorkdirError::SessionClosed);
|
||||
request: WorkdirScopeOverlapRequest,
|
||||
) -> Result<bool, WorkdirError> {
|
||||
match self
|
||||
.operate(WorkdirSessionOperation::ScopeRulesOverlap(request))
|
||||
.await?
|
||||
{
|
||||
WorkdirSessionOperationResult::ScopeRulesOverlap { overlaps } => Ok(overlaps),
|
||||
_ => Err(Self::mismatch("scope_rules_overlap")),
|
||||
}
|
||||
let mut delegations = self.delegations.clone();
|
||||
delegations.push(request.clone());
|
||||
let candidate = Arc::new(Self {
|
||||
client: self.client.clone(),
|
||||
base_url: self.base_url.clone(),
|
||||
authorization: self.authorization.clone(),
|
||||
workdir: self.workdir.clone(),
|
||||
session_id: self.session_id.clone(),
|
||||
capabilities: self.capabilities,
|
||||
delegations,
|
||||
closed: AtomicBool::new(false),
|
||||
});
|
||||
candidate
|
||||
.stat(StatRequest {
|
||||
path: fs_operation::FsPath::new("").expect("empty Workdir path is valid"),
|
||||
})
|
||||
.await?;
|
||||
Ok(candidate)
|
||||
}
|
||||
|
||||
async fn stat(&self, request: StatRequest) -> Result<StatResult, WorkdirError> {
|
||||
@@ -501,10 +592,11 @@ mod client {
|
||||
&self.base_url,
|
||||
&["v1", "workdir-sessions", self.session_id.as_str()],
|
||||
)?;
|
||||
let token = self.authorization.bearer_token("DELETE", url.path(), &[])?;
|
||||
let response = self
|
||||
.client
|
||||
.delete(url)
|
||||
.bearer_auth(self.authorization.bearer_token()?)
|
||||
.bearer_auth(token)
|
||||
.send()
|
||||
.await
|
||||
.map_err(http_unavailable)?;
|
||||
@@ -584,8 +676,42 @@ mod tests {
|
||||
"modified externally",
|
||||
),
|
||||
(WorkdirTransportErrorCode::Unsupported, 400, "unsupported"),
|
||||
(WorkdirTransportErrorCode::Denied, 403, "denied"),
|
||||
(
|
||||
WorkdirTransportErrorCode::OutOfScope,
|
||||
403,
|
||||
"outside allowed scope",
|
||||
),
|
||||
(
|
||||
WorkdirTransportErrorCode::SymlinkOutOfScope,
|
||||
403,
|
||||
"outside allowed requested scope",
|
||||
),
|
||||
(
|
||||
WorkdirTransportErrorCode::BrokenSymlink,
|
||||
400,
|
||||
"broken symlink",
|
||||
),
|
||||
(
|
||||
WorkdirTransportErrorCode::SymlinkTargetIsDirectory,
|
||||
400,
|
||||
"symlink to a directory",
|
||||
),
|
||||
(WorkdirTransportErrorCode::ReadOnly, 403, "read-only"),
|
||||
(WorkdirTransportErrorCode::IsDirectory, 400, "expected file"),
|
||||
(
|
||||
WorkdirTransportErrorCode::SymlinkDirectoryNotTraversed,
|
||||
400,
|
||||
"does not follow symlink directories",
|
||||
),
|
||||
(WorkdirTransportErrorCode::Unavailable, 503, "unavailable"),
|
||||
(WorkdirTransportErrorCode::Internal, 500, "transport failed"),
|
||||
(WorkdirTransportErrorCode::Io, 500, "I/O error"),
|
||||
(
|
||||
WorkdirTransportErrorCode::Transport,
|
||||
502,
|
||||
"transport failed",
|
||||
),
|
||||
(WorkdirTransportErrorCode::Internal, 500, "operation failed"),
|
||||
] {
|
||||
let transport = WorkdirTransportError {
|
||||
code,
|
||||
@@ -620,7 +746,8 @@ mod tests {
|
||||
let transport = WorkdirTransportError::from_workdir_error(&WorkdirError::Transport(
|
||||
"Workspace API request timed out".to_string(),
|
||||
));
|
||||
assert_eq!(transport.code, WorkdirTransportErrorCode::Internal);
|
||||
assert_eq!(transport.code, WorkdirTransportErrorCode::Transport);
|
||||
assert_eq!(transport.code.http_status(), 502);
|
||||
assert_eq!(transport.message, "Workdir transport failed");
|
||||
assert!(matches!(
|
||||
transport.into_workdir_error(),
|
||||
@@ -635,8 +762,25 @@ mod tests {
|
||||
source: std::io::Error::new(std::io::ErrorKind::PermissionDenied, "host detail"),
|
||||
};
|
||||
let transport = WorkdirTransportError::from_workdir_error(&error);
|
||||
assert_eq!(transport.code, WorkdirTransportErrorCode::Internal);
|
||||
assert_eq!(transport.code, WorkdirTransportErrorCode::Denied);
|
||||
assert!(!transport.message.contains("/secret"));
|
||||
assert!(!transport.message.contains("host detail"));
|
||||
assert!(matches!(
|
||||
transport.into_workdir_error(),
|
||||
WorkdirError::Denied(_)
|
||||
));
|
||||
|
||||
let error = WorkdirError::Io {
|
||||
path: "/secret/runtime/root/file".into(),
|
||||
source: std::io::Error::other("host detail"),
|
||||
};
|
||||
let transport = WorkdirTransportError::from_workdir_error(&error);
|
||||
assert_eq!(transport.code, WorkdirTransportErrorCode::Io);
|
||||
assert!(!transport.message.contains("/secret"));
|
||||
assert!(!transport.message.contains("host detail"));
|
||||
assert!(matches!(
|
||||
transport.into_workdir_error(),
|
||||
WorkdirError::Io { .. }
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
+29
-31
@@ -5,10 +5,10 @@
|
||||
//! bound to one Worker. Tools consume sessions; they do not own Workdir
|
||||
//! materialization or cleanup.
|
||||
|
||||
mod delegation;
|
||||
pub mod http;
|
||||
mod local;
|
||||
mod operation;
|
||||
mod scope;
|
||||
pub mod workspace;
|
||||
|
||||
use std::path::{Path, PathBuf};
|
||||
@@ -18,11 +18,6 @@ use async_trait::async_trait;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tokio::sync::broadcast;
|
||||
|
||||
pub use delegation::{
|
||||
AppliedWorkdirDelegation, ReadOnlyWorkdirSession, WorkdirDelegation,
|
||||
WorkdirDelegationPermission, WorkdirDelegationRequest, WorkdirDelegationRule,
|
||||
apply_delegation_chain, delegation_capable_session,
|
||||
};
|
||||
pub use fs_operation::{
|
||||
ContentHash, EditRequest, EditResult, EntryKind, FsPath as WorkdirPath, GlobRequest,
|
||||
GlobResult, GrepOutputMode, GrepRequest, GrepResult, ListEntry, ListRequest, ListResult,
|
||||
@@ -32,6 +27,11 @@ pub use local::{
|
||||
LocalWorkdirSession, SymlinkInfo, WorkdirSessionResource, direct_symlink, first_symlink,
|
||||
};
|
||||
pub use operation::*;
|
||||
pub use scope::{
|
||||
ReadOnlyWorkdirSession, WorkdirScopeAuthorizationRequest, WorkdirScopeLease,
|
||||
WorkdirScopeOverlapRequest, WorkdirToolBroker, WorkdirToolScope, WorkdirToolScopePermission,
|
||||
WorkdirToolScopeRule,
|
||||
};
|
||||
|
||||
/// Persistent, opaque identity of one materialized Workdir.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||||
@@ -148,36 +148,31 @@ pub trait WorkdirSession: std::fmt::Debug + Send + Sync {
|
||||
fn workdir(&self) -> &Workdir;
|
||||
fn capabilities(&self) -> WorkdirSessionCapabilities;
|
||||
|
||||
fn is_delegation_capable(&self) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
/// Whether this session transports the delegation chain to another
|
||||
/// provider boundary that will apply logical cwd/path resolution there.
|
||||
fn transports_delegation_context(&self) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
/// Capture a provider-specific source for a delegated child session.
|
||||
/// Remote providers use this boundary to pin attachment identity without
|
||||
/// exposing transport handles or host paths.
|
||||
async fn capture_delegation_source(
|
||||
/// Validate an attenuated filesystem rule at the provider boundary without
|
||||
/// exposing the resolved host path. Providers that cannot resolve symbolic
|
||||
/// links must reject resolved-policy checks rather than downgrade them.
|
||||
async fn authorize_scope_path(
|
||||
&self,
|
||||
_request: &WorkdirDelegationRequest,
|
||||
) -> Result<WorkdirSessionHandle, WorkdirError> {
|
||||
Err(WorkdirError::Denied(
|
||||
"workdir provider does not support delegated sessions".into(),
|
||||
))
|
||||
request: WorkdirScopeAuthorizationRequest,
|
||||
) -> Result<(), WorkdirError> {
|
||||
if request.rules.iter().any(|rule| {
|
||||
rule.symlink_policy == manifest::SymlinkPolicy::Logical
|
||||
&& scope::rule_allows_path(rule, &request.path, request.permission)
|
||||
}) {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(WorkdirError::Denied(
|
||||
"Workdir provider cannot establish resolved scope authority".to_string(),
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
/// Attenuate this session into a revocable child lease. Only sessions
|
||||
/// created with [`delegation_capable_session`] implement this operation.
|
||||
async fn delegate(
|
||||
async fn scope_rules_overlap(
|
||||
&self,
|
||||
_request: WorkdirDelegationRequest,
|
||||
) -> Result<WorkdirDelegation, WorkdirError> {
|
||||
_request: WorkdirScopeOverlapRequest,
|
||||
) -> Result<bool, WorkdirError> {
|
||||
Err(WorkdirError::Denied(
|
||||
"workdir session is not delegation-capable".into(),
|
||||
"Workdir provider cannot compare resolved scope authority".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
@@ -234,6 +229,9 @@ pub enum WorkdirError {
|
||||
#[error("Workdir session is unavailable: {0}")]
|
||||
Unavailable(String),
|
||||
|
||||
#[error("Workdir operation failed")]
|
||||
OperationFailed,
|
||||
|
||||
#[error("Workdir transport failed: {0}")]
|
||||
Transport(String),
|
||||
|
||||
|
||||
+358
-87
@@ -18,7 +18,7 @@ use std::sync::{Arc, Mutex as StdMutex};
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use manifest::{Permission, Scope, ScopeConfig, ScopeRule, SharedScope};
|
||||
use manifest::{Permission, Scope, SharedScope, SymlinkPolicy};
|
||||
use sha2::{Digest, Sha256};
|
||||
use tokio::process::Command;
|
||||
use tokio::sync::{Mutex, broadcast, watch};
|
||||
@@ -28,9 +28,9 @@ use crate::{
|
||||
CommandEvent, CommandHandle, CommandOutput, CommandOutputRequest, CommandRequest,
|
||||
CommandSnapshot, CommandStatus, CommandStream, CommandStreamSlice, EditRequest, EditResult,
|
||||
GlobRequest, GlobResult, GrepRequest, GrepResult, ListRequest, ListResult, ReadRequest,
|
||||
ReadResult, StatRequest, StatResult, Workdir, WorkdirDelegationPermission,
|
||||
WorkdirDelegationRequest, WorkdirError, WorkdirPath, WorkdirSession,
|
||||
WorkdirSessionCapabilities, WorkdirSessionCapability, WorkdirSessionHandle, WriteRequest,
|
||||
ReadResult, StatRequest, StatResult, Workdir, WorkdirError, WorkdirPath,
|
||||
WorkdirScopeAuthorizationRequest, WorkdirScopeOverlapRequest, WorkdirSession,
|
||||
WorkdirSessionCapabilities, WorkdirSessionCapability, WorkdirToolScopePermission, WriteRequest,
|
||||
WriteResult,
|
||||
};
|
||||
#[cfg(test)]
|
||||
@@ -213,6 +213,52 @@ impl fs_operation::FsAccessPolicy for ScopeAccess {
|
||||
fn is_writable(&self, path: &Path) -> bool {
|
||||
self.0.is_writable(path)
|
||||
}
|
||||
|
||||
fn is_readable_paths(&self, logical: &Path, resolved: &Path) -> bool {
|
||||
matches!(
|
||||
self.0.permission_at_paths(logical, resolved),
|
||||
Some(Permission::Read | Permission::Write)
|
||||
)
|
||||
}
|
||||
|
||||
fn is_writable_paths(&self, logical: &Path, resolved: &Path) -> bool {
|
||||
self.0.permission_at_paths(logical, resolved) == Some(Permission::Write)
|
||||
}
|
||||
}
|
||||
|
||||
fn path_sets_overlap(
|
||||
left: &Path,
|
||||
left_recursive: bool,
|
||||
right: &Path,
|
||||
right_recursive: bool,
|
||||
) -> bool {
|
||||
match (left_recursive, right_recursive) {
|
||||
(true, true) => left.starts_with(right) || right.starts_with(left),
|
||||
(true, false) => {
|
||||
right.starts_with(left)
|
||||
|| left == right
|
||||
|| left.parent().is_some_and(|parent| parent == right)
|
||||
}
|
||||
(false, true) => {
|
||||
left.starts_with(right)
|
||||
|| left == right
|
||||
|| right.parent().is_some_and(|parent| parent == left)
|
||||
}
|
||||
(false, false) => {
|
||||
left == right
|
||||
|| left.parent().is_some_and(|parent| parent == right)
|
||||
|| right.parent().is_some_and(|parent| parent == left)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn rule_targets(
|
||||
root: &Path,
|
||||
rule: &crate::WorkdirToolScopeRule,
|
||||
) -> std::io::Result<(PathBuf, PathBuf)> {
|
||||
let logical = root.join(rule.target.as_str());
|
||||
let resolved = fs_operation::resolve_access_path(&logical)?;
|
||||
Ok((logical, resolved))
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
@@ -399,6 +445,11 @@ impl LocalWorkdirSession {
|
||||
return Err(WorkdirError::RelativePath(path.to_path_buf()));
|
||||
}
|
||||
let symlink = first_symlink(path);
|
||||
if let Some(info) = symlink.as_ref()
|
||||
&& !info.target_exists
|
||||
{
|
||||
return Err(broken_symlink_error(path, info));
|
||||
}
|
||||
let scope = self.inner.scope.load();
|
||||
if !scope.is_readable(path) {
|
||||
return Err(symlink_out_of_scope_or_plain(
|
||||
@@ -408,11 +459,6 @@ impl LocalWorkdirSession {
|
||||
&scope,
|
||||
));
|
||||
}
|
||||
if let Some(info) = symlink.as_ref() {
|
||||
if !info.target_exists {
|
||||
return Err(broken_symlink_error(path, info));
|
||||
}
|
||||
}
|
||||
let meta = std::fs::metadata(path).map_err(|e| match e.kind() {
|
||||
std::io::ErrorKind::NotFound => WorkdirError::NotFound(path.to_path_buf()),
|
||||
_ => WorkdirError::io(path, e),
|
||||
@@ -558,67 +604,84 @@ impl WorkdirSession for LocalWorkdirSession {
|
||||
self.inner.capabilities
|
||||
}
|
||||
|
||||
async fn capture_delegation_source(
|
||||
async fn authorize_scope_path(
|
||||
&self,
|
||||
request: &WorkdirDelegationRequest,
|
||||
) -> Result<WorkdirSessionHandle, WorkdirError> {
|
||||
let host_rules = request
|
||||
.rules
|
||||
.iter()
|
||||
.map(|rule| ScopeRule {
|
||||
target: self.inner.root.join(rule.target.as_str()),
|
||||
permission: match rule.permission {
|
||||
WorkdirDelegationPermission::Read => Permission::Read,
|
||||
WorkdirDelegationPermission::Write => Permission::Write,
|
||||
},
|
||||
recursive: rule.recursive,
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
for (logical, host) in request.rules.iter().zip(&host_rules) {
|
||||
if logical.permission == WorkdirDelegationPermission::Write {
|
||||
let resolved = Scope::resolved_target(host)
|
||||
.map_err(|error| WorkdirError::Denied(error.to_string()))?;
|
||||
if resolved != host.target {
|
||||
return Err(WorkdirError::Denied(format!(
|
||||
"write delegation target `{}` traverses a symlink",
|
||||
logical.target
|
||||
)));
|
||||
}
|
||||
}
|
||||
}
|
||||
let parent_scope = self.inner.scope.snapshot();
|
||||
for rule in &host_rules {
|
||||
if !parent_scope
|
||||
.allows_rule(rule)
|
||||
.map_err(|error| WorkdirError::Denied(error.to_string()))?
|
||||
{
|
||||
return Err(WorkdirError::Denied(format!(
|
||||
"delegated provider scope `{}` exceeds the parent session",
|
||||
rule.target.display()
|
||||
)));
|
||||
}
|
||||
}
|
||||
let child_scope = Scope::from_config(&ScopeConfig {
|
||||
allow: host_rules,
|
||||
deny: Vec::new(),
|
||||
})
|
||||
.map_err(|error| WorkdirError::Denied(error.to_string()))?;
|
||||
let child_cwd = self.inner.root.join(request.cwd.as_str());
|
||||
if !child_scope.is_readable(&child_cwd)
|
||||
|| !std::fs::metadata(&child_cwd).is_ok_and(|metadata| metadata.is_dir())
|
||||
{
|
||||
request: WorkdirScopeAuthorizationRequest,
|
||||
) -> Result<(), WorkdirError> {
|
||||
self.ensure_open()?;
|
||||
let logical = self.inner.root.join(request.path.as_str());
|
||||
let resolved = fs_operation::resolve_access_path(&logical)
|
||||
.map_err(|error| WorkdirError::io(&logical, error))?;
|
||||
let parent_permission = self
|
||||
.inner
|
||||
.scope
|
||||
.load()
|
||||
.permission_at_paths(&logical, &resolved);
|
||||
let parent_allows = match request.permission {
|
||||
WorkdirToolScopePermission::Read => matches!(
|
||||
parent_permission,
|
||||
Some(Permission::Read | Permission::Write)
|
||||
),
|
||||
WorkdirToolScopePermission::Write => parent_permission == Some(Permission::Write),
|
||||
};
|
||||
if !parent_allows {
|
||||
return Err(WorkdirError::Denied(format!(
|
||||
"delegated cwd `{}` is not a readable Workdir directory",
|
||||
request.cwd
|
||||
"Workdir path `{}` exceeds the provider attachment scope",
|
||||
request.path
|
||||
)));
|
||||
}
|
||||
Ok(Arc::new(LocalWorkdirSession::materialized_bound(
|
||||
self.inner.workdir.clone(),
|
||||
self.inner.root.clone(),
|
||||
self.inner.root.clone(),
|
||||
SharedScope::new(child_scope),
|
||||
self.inner.capabilities,
|
||||
)))
|
||||
let allowed = request.rules.iter().any(|rule| {
|
||||
if request.permission == WorkdirToolScopePermission::Write
|
||||
&& rule.permission != WorkdirToolScopePermission::Write
|
||||
{
|
||||
return false;
|
||||
}
|
||||
let logical_target = self.inner.root.join(rule.target.as_str());
|
||||
let (candidate, target) = match rule.symlink_policy {
|
||||
SymlinkPolicy::Logical => (logical.as_path(), logical_target),
|
||||
SymlinkPolicy::Resolved => {
|
||||
let Ok(target) = fs_operation::resolve_access_path(&logical_target) else {
|
||||
return false;
|
||||
};
|
||||
(resolved.as_path(), target)
|
||||
}
|
||||
};
|
||||
if rule.recursive {
|
||||
candidate.starts_with(target)
|
||||
} else {
|
||||
candidate == target || candidate.parent() == Some(target.as_path())
|
||||
}
|
||||
});
|
||||
if allowed {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(WorkdirError::Denied(format!(
|
||||
"Workdir path `{}` is outside the provider-resolved delegated scope",
|
||||
request.path
|
||||
)))
|
||||
}
|
||||
}
|
||||
|
||||
async fn scope_rules_overlap(
|
||||
&self,
|
||||
request: WorkdirScopeOverlapRequest,
|
||||
) -> Result<bool, WorkdirError> {
|
||||
self.ensure_open()?;
|
||||
let (left_logical, left_resolved) = rule_targets(&self.inner.root, &request.left)
|
||||
.map_err(|error| WorkdirError::io(&self.inner.root, error))?;
|
||||
let (right_logical, right_resolved) = rule_targets(&self.inner.root, &request.right)
|
||||
.map_err(|error| WorkdirError::io(&self.inner.root, error))?;
|
||||
Ok(path_sets_overlap(
|
||||
&left_logical,
|
||||
request.left.recursive,
|
||||
&right_logical,
|
||||
request.right.recursive,
|
||||
) || path_sets_overlap(
|
||||
&left_resolved,
|
||||
request.left.recursive,
|
||||
&right_resolved,
|
||||
request.right.recursive,
|
||||
))
|
||||
}
|
||||
|
||||
async fn stat(&self, request: StatRequest) -> Result<StatResult, WorkdirError> {
|
||||
@@ -694,9 +757,20 @@ impl WorkdirSession for LocalWorkdirSession {
|
||||
{
|
||||
return Err(WorkdirError::OutOfScope(spill_dir.to_path_buf()));
|
||||
}
|
||||
let cwd = if let Some(logical_cwd) = request.cwd.as_ref() {
|
||||
let cwd = self.resolve(logical_cwd);
|
||||
let scope = self.inner.scope.snapshot();
|
||||
if !scope.is_readable(&cwd)
|
||||
|| !std::fs::metadata(&cwd).is_ok_and(|metadata| metadata.is_dir())
|
||||
{
|
||||
return Err(WorkdirError::OutOfScope(cwd));
|
||||
}
|
||||
cwd
|
||||
} else {
|
||||
self.inner.cwd.clone()
|
||||
};
|
||||
let id = self.inner.next_command_id.fetch_add(1, Ordering::Relaxed);
|
||||
let handle = CommandHandle(format!("command-{id}"));
|
||||
let cwd = self.inner.cwd.clone();
|
||||
let (completion_tx, completion) = watch::channel(false);
|
||||
let command_id = handle.0.clone();
|
||||
let telemetry = self.inner.command_telemetry.clone();
|
||||
@@ -1388,6 +1462,22 @@ mod tests {
|
||||
)
|
||||
}
|
||||
|
||||
fn make_logical_fs(dir: &TempDir) -> LocalWorkdirSession {
|
||||
LocalWorkdirSession::new(
|
||||
Scope::from_config(&ScopeConfig {
|
||||
allow: vec![ScopeRule {
|
||||
target: dir.path().to_path_buf(),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: SymlinkPolicy::Logical,
|
||||
}],
|
||||
deny: Vec::new(),
|
||||
})
|
||||
.unwrap(),
|
||||
dir.path().to_path_buf(),
|
||||
)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn logical_provider_operations_cover_read_write_edit_stat_and_list() {
|
||||
let dir = TempDir::new().unwrap();
|
||||
@@ -1516,6 +1606,7 @@ mod tests {
|
||||
command: "sleep 30".to_owned(),
|
||||
timeout_secs: 60,
|
||||
output_limit: 1024,
|
||||
cwd: None,
|
||||
spill_dir: None,
|
||||
tool_call_id: None,
|
||||
},
|
||||
@@ -1586,6 +1677,102 @@ mod tests {
|
||||
assert_eq!(read.bytes, b"persisted");
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
#[tokio::test]
|
||||
async fn resolved_provider_scope_rejects_read_and_write_through_outside_alias() {
|
||||
use std::os::unix::fs::symlink;
|
||||
|
||||
let root = TempDir::new().unwrap();
|
||||
let outside = TempDir::new().unwrap();
|
||||
let target = outside.path().join("target.txt");
|
||||
fs::write(&target, "secret").unwrap();
|
||||
symlink(&target, root.path().join("alias.txt")).unwrap();
|
||||
symlink(outside.path(), root.path().join("alias-dir")).unwrap();
|
||||
let workdir = make_fs(&root);
|
||||
|
||||
assert!(matches!(
|
||||
WorkdirSession::read(
|
||||
&workdir,
|
||||
ReadRequest {
|
||||
path: WorkdirPath::new("alias.txt").unwrap(),
|
||||
offset: 0,
|
||||
limit: 10,
|
||||
max_bytes: 1024,
|
||||
}
|
||||
)
|
||||
.await,
|
||||
Err(WorkdirError::SymlinkOutOfScope { .. })
|
||||
));
|
||||
assert!(matches!(
|
||||
WorkdirSession::write(
|
||||
&workdir,
|
||||
WriteRequest {
|
||||
path: WorkdirPath::new("alias.txt").unwrap(),
|
||||
content: b"changed".to_vec(),
|
||||
expected_hash: None,
|
||||
}
|
||||
)
|
||||
.await,
|
||||
Err(WorkdirError::SymlinkOutOfScope { .. })
|
||||
));
|
||||
assert_eq!(fs::read_to_string(target).unwrap(), "secret");
|
||||
assert!(matches!(
|
||||
WorkdirSession::write(
|
||||
&workdir,
|
||||
WriteRequest {
|
||||
path: WorkdirPath::new("alias-dir/new.txt").unwrap(),
|
||||
content: b"new".to_vec(),
|
||||
expected_hash: None,
|
||||
}
|
||||
)
|
||||
.await,
|
||||
Err(WorkdirError::ReadOnly(_))
|
||||
));
|
||||
assert!(!outside.path().join("new.txt").exists());
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
#[tokio::test]
|
||||
async fn resolved_deny_blocks_missing_write_through_logical_alias() {
|
||||
use std::os::unix::fs::symlink;
|
||||
|
||||
let root = TempDir::new().unwrap();
|
||||
let outside = TempDir::new().unwrap();
|
||||
symlink(outside.path(), root.path().join("alias")).unwrap();
|
||||
let workdir = LocalWorkdirSession::new(
|
||||
Scope::from_config(&ScopeConfig {
|
||||
allow: vec![ScopeRule {
|
||||
target: root.path().to_path_buf(),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: SymlinkPolicy::Logical,
|
||||
}],
|
||||
deny: vec![ScopeRule {
|
||||
target: outside.path().join("blocked.txt"),
|
||||
permission: Permission::Read,
|
||||
recursive: false,
|
||||
symlink_policy: SymlinkPolicy::Logical,
|
||||
}],
|
||||
})
|
||||
.unwrap(),
|
||||
root.path().to_path_buf(),
|
||||
);
|
||||
|
||||
assert!(matches!(
|
||||
WorkdirSession::write(
|
||||
&workdir,
|
||||
WriteRequest {
|
||||
path: WorkdirPath::new("alias/blocked.txt").unwrap(),
|
||||
content: b"blocked".to_vec(),
|
||||
expected_hash: None,
|
||||
}
|
||||
)
|
||||
.await,
|
||||
Err(WorkdirError::ReadOnly(_))
|
||||
));
|
||||
assert!(!outside.path().join("blocked.txt").exists());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn capability_boundary_rejects_direct_unsupported_operation() {
|
||||
let dir = TempDir::new().unwrap();
|
||||
@@ -1688,7 +1875,7 @@ mod tests {
|
||||
|
||||
#[cfg(unix)]
|
||||
#[test]
|
||||
fn read_bytes_reports_symlink_target_outside_scope() {
|
||||
fn read_bytes_allows_logical_symlink_path_with_target_outside_scope() {
|
||||
use std::os::unix::fs::symlink;
|
||||
|
||||
let dir = TempDir::new().unwrap();
|
||||
@@ -1698,16 +1885,8 @@ mod tests {
|
||||
let link = dir.path().join("outside-repo.txt");
|
||||
symlink(&target, &link).unwrap();
|
||||
|
||||
let fs = make_fs(&dir);
|
||||
let err = fs.read_bytes(&link).unwrap_err();
|
||||
assert!(
|
||||
matches!(
|
||||
err,
|
||||
WorkdirError::SymlinkOutOfScope { ref path, target: ref err_target, required_permission: "read" }
|
||||
if path == &link && err_target == &target.canonicalize().unwrap()
|
||||
),
|
||||
"expected symlink out-of-scope diagnostic, got {err:?}"
|
||||
);
|
||||
let fs = make_logical_fs(&dir);
|
||||
assert_eq!(fs.read_bytes(&link).unwrap(), b"secret");
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
@@ -1799,7 +1978,7 @@ mod tests {
|
||||
|
||||
#[cfg(unix)]
|
||||
#[test]
|
||||
fn write_reports_symlink_target_outside_scope() {
|
||||
fn write_allows_logical_symlink_path_with_target_outside_scope() {
|
||||
use std::os::unix::fs::symlink;
|
||||
|
||||
let dir = TempDir::new().unwrap();
|
||||
@@ -1809,15 +1988,14 @@ mod tests {
|
||||
let link = dir.path().join("outside-repo.txt");
|
||||
symlink(&target, &link).unwrap();
|
||||
|
||||
let fs = make_fs(&dir);
|
||||
let err = fs.write(&link, b"new").unwrap_err();
|
||||
let fs = make_logical_fs(&dir);
|
||||
fs.write(&link, b"new").unwrap();
|
||||
assert_eq!(fs::read(&target).unwrap(), b"new");
|
||||
assert!(
|
||||
matches!(
|
||||
err,
|
||||
WorkdirError::SymlinkOutOfScope { ref path, target: ref err_target, required_permission: "write" }
|
||||
if path == &link && err_target == &target.canonicalize().unwrap()
|
||||
),
|
||||
"expected write symlink out-of-scope diagnostic, got {err:?}"
|
||||
fs::symlink_metadata(&link)
|
||||
.unwrap()
|
||||
.file_type()
|
||||
.is_symlink()
|
||||
);
|
||||
}
|
||||
|
||||
@@ -1840,11 +2018,13 @@ mod tests {
|
||||
target: dir.path().to_path_buf(),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}],
|
||||
deny: vec![ScopeRule {
|
||||
target: sub.clone(),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}],
|
||||
};
|
||||
let scope = Scope::from_config(&cfg).unwrap();
|
||||
@@ -1908,6 +2088,7 @@ mod tests {
|
||||
target: extra.path().to_path_buf(),
|
||||
permission: Permission::Read,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}])
|
||||
})
|
||||
.unwrap();
|
||||
@@ -1944,6 +2125,7 @@ mod tests {
|
||||
target: sub.clone(),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}])
|
||||
})
|
||||
.unwrap();
|
||||
@@ -1980,6 +2162,7 @@ mod tests {
|
||||
target: dir.path().to_path_buf(),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
}])
|
||||
})
|
||||
.unwrap();
|
||||
@@ -1995,6 +2178,83 @@ mod tests {
|
||||
));
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
#[tokio::test]
|
||||
async fn provider_uses_explicit_logical_policy_through_symlinked_directories() {
|
||||
use std::os::unix::fs::symlink;
|
||||
|
||||
let dir = TempDir::new().unwrap();
|
||||
let outside = TempDir::new().unwrap();
|
||||
std::fs::write(outside.path().join("worker.json"), "scope-needle\n").unwrap();
|
||||
symlink(outside.path(), dir.path().join("yoi.local")).unwrap();
|
||||
let workdir = make_logical_fs(&dir);
|
||||
|
||||
let read = WorkdirSession::read(
|
||||
&workdir,
|
||||
ReadRequest {
|
||||
path: WorkdirPath::new("yoi.local/worker.json").unwrap(),
|
||||
offset: 0,
|
||||
limit: 100,
|
||||
max_bytes: 1024,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(read.bytes, b"scope-needle\n");
|
||||
let list = WorkdirSession::list(
|
||||
&workdir,
|
||||
ListRequest {
|
||||
path: WorkdirPath::new("yoi.local").unwrap(),
|
||||
limit: 10,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
list.entries[0].path,
|
||||
WorkdirPath::new("yoi.local/worker.json").unwrap()
|
||||
);
|
||||
let glob = WorkdirSession::glob(
|
||||
&workdir,
|
||||
GlobRequest {
|
||||
pattern: "**/*.json".into(),
|
||||
path: WorkdirPath::new("yoi.local").unwrap(),
|
||||
limit: 10,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
glob.paths,
|
||||
[WorkdirPath::new("yoi.local/worker.json").unwrap()]
|
||||
);
|
||||
let grep = WorkdirSession::grep(
|
||||
&workdir,
|
||||
GrepRequest {
|
||||
pattern: "scope-needle".into(),
|
||||
path: WorkdirPath::new("yoi.local").unwrap(),
|
||||
glob: Some("*.json".into()),
|
||||
file_type: None,
|
||||
case_insensitive: false,
|
||||
before_context: 0,
|
||||
after_context: 0,
|
||||
multiline: false,
|
||||
output_mode: crate::GrepOutputMode::Content,
|
||||
limit: 10,
|
||||
offset: 0,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(grep.match_count, 1);
|
||||
assert!(grep.output.contains("yoi.local/worker.json"));
|
||||
assert!(
|
||||
!workdir
|
||||
.scope()
|
||||
.is_readable(&outside.path().join("worker.json"))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn provider_executes_glob_grep_and_command_at_the_materialization() {
|
||||
let dir = TempDir::new().unwrap();
|
||||
@@ -2043,6 +2303,7 @@ mod tests {
|
||||
command: "pwd && printf provider-command".into(),
|
||||
timeout_secs: 5,
|
||||
output_limit: 4096,
|
||||
cwd: None,
|
||||
spill_dir: None,
|
||||
tool_call_id: None,
|
||||
},
|
||||
@@ -2081,11 +2342,13 @@ mod tests {
|
||||
target: dir.path().to_path_buf(),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
},
|
||||
ScopeRule {
|
||||
target: spill.path().to_path_buf(),
|
||||
permission: Permission::Read,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
},
|
||||
],
|
||||
deny: Vec::new(),
|
||||
@@ -2141,6 +2404,7 @@ mod tests {
|
||||
command: "printf hidden".into(),
|
||||
timeout_secs: 5,
|
||||
output_limit: 1,
|
||||
cwd: None,
|
||||
spill_dir: Some(spill.path().to_path_buf()),
|
||||
tool_call_id: None,
|
||||
},
|
||||
@@ -2161,11 +2425,13 @@ mod tests {
|
||||
target: dir.path().to_path_buf(),
|
||||
permission: Permission::Write,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
},
|
||||
ScopeRule {
|
||||
target: spill.path().to_path_buf(),
|
||||
permission: Permission::Read,
|
||||
recursive: true,
|
||||
symlink_policy: Default::default(),
|
||||
},
|
||||
],
|
||||
deny: Vec::new(),
|
||||
@@ -2178,6 +2444,7 @@ mod tests {
|
||||
command: "i=0; while [ $i -lt 200 ]; do printf 'line-%03d\\n' \"$i\"; i=$((i+1)); done; printf 'FINAL-NEEDLE\\n'".into(),
|
||||
timeout_secs: 5,
|
||||
output_limit: 64,
|
||||
cwd: None,
|
||||
spill_dir: Some(spill.path().to_path_buf()),
|
||||
tool_call_id: None,
|
||||
},
|
||||
@@ -2224,6 +2491,7 @@ mod tests {
|
||||
command: "printf 'aéz'".into(),
|
||||
timeout_secs: 5,
|
||||
output_limit: 1024,
|
||||
cwd: None,
|
||||
spill_dir: None,
|
||||
tool_call_id: None,
|
||||
},
|
||||
@@ -2449,6 +2717,7 @@ mod tests {
|
||||
command: "printf ready; printf warning >&2; sleep 0.2; printf done".into(),
|
||||
timeout_secs: 5,
|
||||
output_limit: 1024,
|
||||
cwd: None,
|
||||
spill_dir: None,
|
||||
tool_call_id: Some("tool-7".into()),
|
||||
},
|
||||
@@ -2553,6 +2822,7 @@ mod tests {
|
||||
command: "sleep 30".into(),
|
||||
timeout_secs: 1,
|
||||
output_limit: 1024,
|
||||
cwd: None,
|
||||
spill_dir: None,
|
||||
tool_call_id: None,
|
||||
},
|
||||
@@ -2623,6 +2893,7 @@ mod tests {
|
||||
command: "sleep 30".into(),
|
||||
timeout_secs: 60,
|
||||
output_limit: 1024,
|
||||
cwd: None,
|
||||
spill_dir: None,
|
||||
tool_call_id: None,
|
||||
},
|
||||
|
||||
@@ -11,6 +11,10 @@ pub struct CommandRequest {
|
||||
pub command: String,
|
||||
pub timeout_secs: u64,
|
||||
pub output_limit: usize,
|
||||
/// Workdir-relative command directory. Providers validate it against the
|
||||
/// active session before process start.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cwd: Option<fs_operation::FsPath>,
|
||||
/// Provider-local directory where complete output is retained when the
|
||||
/// inline result exceeds `output_limit`.
|
||||
pub spill_dir: Option<PathBuf>,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+12
-261
@@ -6,7 +6,12 @@
|
||||
//! [`crate::http`].
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::fmt;
|
||||
|
||||
pub use workspace_api::{
|
||||
RuntimeWorkingDirectoryCleanupTarget, RuntimeWorkingDirectorySummary,
|
||||
WorkingDirectoryCleanupTarget, WorkingDirectoryMaterializerKind as MaterializerKind,
|
||||
WorkingDirectoryOccupancy, WorkingDirectoryStatusKind, WorkingDirectorySummary,
|
||||
};
|
||||
|
||||
/// Stable Workspace identity for a Worker hosted by a Runtime.
|
||||
#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
|
||||
@@ -26,83 +31,6 @@ impl RuntimeWorkerRef {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum MaterializerKind {
|
||||
#[default]
|
||||
RuntimeGitCache,
|
||||
/// Legacy persisted value from the pre-cache local `git worktree` materializer.
|
||||
LocalGitWorktree,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum WorkingDirectoryStatusKind {
|
||||
Active,
|
||||
CleanupPending,
|
||||
Corrupted,
|
||||
NotFound,
|
||||
Unknown,
|
||||
}
|
||||
|
||||
impl WorkingDirectoryStatusKind {
|
||||
pub const fn as_str(&self) -> &'static str {
|
||||
match self {
|
||||
Self::Active => "active",
|
||||
Self::CleanupPending => "cleanup_pending",
|
||||
Self::Corrupted => "corrupted",
|
||||
Self::NotFound => "not_found",
|
||||
Self::Unknown => "unknown",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for WorkingDirectoryStatusKind {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter.write_str(self.as_str())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct WorkingDirectoryCleanupTarget {
|
||||
pub kind: String,
|
||||
pub working_directory_id: String,
|
||||
pub repository_id: String,
|
||||
}
|
||||
|
||||
/// Durable Workspace occupancy projection for one Workdir.
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
|
||||
pub struct WorkingDirectoryOccupancy {
|
||||
#[serde(flatten)]
|
||||
pub worker: RuntimeWorkerRef,
|
||||
pub display_name: String,
|
||||
pub linked_at: String,
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for WorkingDirectoryOccupancy {
|
||||
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
|
||||
where
|
||||
D: serde::Deserializer<'de>,
|
||||
{
|
||||
#[derive(Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct Wire {
|
||||
runtime_id: String,
|
||||
worker_id: String,
|
||||
display_name: String,
|
||||
linked_at: String,
|
||||
}
|
||||
|
||||
let wire = Wire::deserialize(deserializer)?;
|
||||
Ok(Self {
|
||||
worker: RuntimeWorkerRef::new(wire.runtime_id, wire.worker_id),
|
||||
display_name: wire.display_name,
|
||||
linked_at: wire.linked_at,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Immutable materialization provenance retained by Workspace inventory.
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
@@ -139,100 +67,6 @@ pub struct WorkingDirectoryCurrentObservation {
|
||||
pub occupied_by: Option<WorkingDirectoryOccupancy>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct WorkingDirectorySummary {
|
||||
pub working_directory_id: String,
|
||||
pub repository_id: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub creation_selector: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub creation_ref: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub creation_tree: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub current_selector: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub current_ref: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub current_tree: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub observed_at_epoch_seconds: Option<u64>,
|
||||
pub materializer_kind: MaterializerKind,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cleanup_target: Option<WorkingDirectoryCleanupTarget>,
|
||||
pub status: WorkingDirectoryStatusKind,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cleanliness: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub primary_worker_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub occupied_by: Option<WorkingDirectoryOccupancy>,
|
||||
}
|
||||
|
||||
impl WorkingDirectorySummary {
|
||||
/// Workspace-managed inventory rows carry explicit cleanup authority.
|
||||
pub fn is_workspace_managed(&self) -> bool {
|
||||
self.cleanup_target.is_some()
|
||||
}
|
||||
|
||||
pub fn provenance(&self) -> WorkingDirectoryProvenance {
|
||||
WorkingDirectoryProvenance {
|
||||
creation_selector: self.creation_selector.clone(),
|
||||
creation_ref: self.creation_ref.clone(),
|
||||
creation_tree: self.creation_tree.clone(),
|
||||
materializer_kind: self.materializer_kind.clone(),
|
||||
cleanup_target: self.cleanup_target.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn current_observation(&self) -> WorkingDirectoryCurrentObservation {
|
||||
WorkingDirectoryCurrentObservation {
|
||||
current_selector: self.current_selector.clone(),
|
||||
current_ref: self.current_ref.clone(),
|
||||
current_tree: self.current_tree.clone(),
|
||||
observed_at_epoch_seconds: self.observed_at_epoch_seconds,
|
||||
status: self.status.clone(),
|
||||
cleanliness: self.cleanliness.clone(),
|
||||
primary_worker_id: self.primary_worker_id.clone(),
|
||||
occupied_by: self.occupied_by.clone(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum WorkingDirectoryDiagnosticSeverity {
|
||||
Info,
|
||||
Warning,
|
||||
Error,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct WorkingDirectoryDiagnostic {
|
||||
pub code: String,
|
||||
pub severity: WorkingDirectoryDiagnosticSeverity,
|
||||
pub message: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct WorkingDirectoryListResponse {
|
||||
pub workspace_id: String,
|
||||
pub items: Vec<WorkingDirectorySummary>,
|
||||
pub diagnostics: Vec<WorkingDirectoryDiagnostic>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct WorkingDirectoryDetailResponse {
|
||||
pub workspace_id: String,
|
||||
pub runtime_id: String,
|
||||
pub item: WorkingDirectorySummary,
|
||||
pub diagnostics: Vec<WorkingDirectoryDiagnostic>,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
@@ -255,103 +89,20 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn occupied_and_free_list_response_round_trips() {
|
||||
let response = WorkingDirectoryListResponse {
|
||||
workspace_id: "workspace".to_string(),
|
||||
items: vec![
|
||||
WorkingDirectorySummary {
|
||||
working_directory_id: "occupied".to_string(),
|
||||
repository_id: "repo".to_string(),
|
||||
creation_selector: Some("develop".to_string()),
|
||||
creation_ref: Some("abc123".to_string()),
|
||||
creation_tree: Some("tree123".to_string()),
|
||||
current_selector: Some("work/ticket".to_string()),
|
||||
current_ref: Some("def456".to_string()),
|
||||
current_tree: Some("tree456".to_string()),
|
||||
observed_at_epoch_seconds: Some(1_777_777_777),
|
||||
materializer_kind: MaterializerKind::LocalGitWorktree,
|
||||
cleanup_target: Some(WorkingDirectoryCleanupTarget {
|
||||
kind: "git_worktree".to_string(),
|
||||
working_directory_id: "occupied".to_string(),
|
||||
repository_id: "repo".to_string(),
|
||||
}),
|
||||
status: WorkingDirectoryStatusKind::Active,
|
||||
cleanliness: Some("clean".to_string()),
|
||||
primary_worker_id: None,
|
||||
occupied_by: Some(WorkingDirectoryOccupancy {
|
||||
worker: RuntimeWorkerRef::new("arcadia", "worker-opaque-64"),
|
||||
display_name: "Coder".to_string(),
|
||||
linked_at: "2026-08-12T00:00:00Z".to_string(),
|
||||
}),
|
||||
},
|
||||
WorkingDirectorySummary {
|
||||
working_directory_id: "free".to_string(),
|
||||
repository_id: "repo".to_string(),
|
||||
creation_selector: None,
|
||||
creation_ref: None,
|
||||
creation_tree: None,
|
||||
current_selector: None,
|
||||
current_ref: Some("987fed".to_string()),
|
||||
current_tree: None,
|
||||
observed_at_epoch_seconds: None,
|
||||
materializer_kind: MaterializerKind::LocalGitWorktree,
|
||||
cleanup_target: None,
|
||||
status: WorkingDirectoryStatusKind::Active,
|
||||
cleanliness: Some("unknown".to_string()),
|
||||
primary_worker_id: None,
|
||||
occupied_by: None,
|
||||
},
|
||||
],
|
||||
diagnostics: vec![WorkingDirectoryDiagnostic {
|
||||
code: "observed".to_string(),
|
||||
severity: WorkingDirectoryDiagnosticSeverity::Info,
|
||||
message: "inventory observed".to_string(),
|
||||
}],
|
||||
};
|
||||
|
||||
let encoded = serde_json::to_value(&response).unwrap();
|
||||
fn workspace_workdir_projection_reexports_workspace_api_authority() {
|
||||
assert_eq!(
|
||||
encoded["items"][0]["occupied_by"]["worker_id"],
|
||||
"worker-opaque-64"
|
||||
std::any::TypeId::of::<WorkingDirectorySummary>(),
|
||||
std::any::TypeId::of::<workspace_api::WorkingDirectorySummary>()
|
||||
);
|
||||
assert!(
|
||||
encoded["items"][0]["occupied_by"]
|
||||
.get("runtime_worker_id")
|
||||
.is_none()
|
||||
assert_eq!(
|
||||
std::any::TypeId::of::<WorkingDirectoryOccupancy>(),
|
||||
std::any::TypeId::of::<workspace_api::WorkingDirectoryOccupancy>()
|
||||
);
|
||||
assert!(encoded["items"][1].get("occupied_by").is_none());
|
||||
|
||||
let mut stale = encoded.clone();
|
||||
stale["items"][0]["occupied_by"]["runtime_worker_id"] = serde_json::json!(64);
|
||||
assert!(serde_json::from_value::<WorkingDirectoryListResponse>(stale).is_err());
|
||||
|
||||
let decoded: WorkingDirectoryListResponse = serde_json::from_value(encoded).unwrap();
|
||||
assert_eq!(decoded, response);
|
||||
|
||||
let detail = WorkingDirectoryDetailResponse {
|
||||
workspace_id: decoded.workspace_id.clone(),
|
||||
runtime_id: "arcadia".to_string(),
|
||||
item: decoded.items[0].clone(),
|
||||
diagnostics: decoded.diagnostics.clone(),
|
||||
};
|
||||
let encoded = serde_json::to_value(&detail).unwrap();
|
||||
let decoded: WorkingDirectoryDetailResponse = serde_json::from_value(encoded).unwrap();
|
||||
assert_eq!(decoded, detail);
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct WorkspaceWorkdirSessionOperationRequest {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub expected_session_fence: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
pub delegations: Vec<crate::WorkdirDelegationRequest>,
|
||||
pub operation: crate::http::WorkdirSessionOperation,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct WorkspaceWorkdirSessionFence {
|
||||
pub value: String,
|
||||
}
|
||||
|
||||
@@ -39,7 +39,9 @@ reqwest = { version = "0.13", optional = true, default-features = false, feature
|
||||
ring.workspace = true
|
||||
tar.workspace = true
|
||||
thiserror = { workspace = true }
|
||||
tokio = { workspace = true, features = ["net", "rt", "sync", "time"] }
|
||||
tokio = { workspace = true, features = ["net", "process", "rt", "sync", "time"] }
|
||||
tracing.workspace = true
|
||||
tracing-subscriber.workspace = true
|
||||
toml.workspace = true
|
||||
url.workspace = true
|
||||
uuid = { workspace = true, features = ["v7"] }
|
||||
|
||||
@@ -2,6 +2,7 @@ use base64::Engine;
|
||||
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
|
||||
use ring::rand::{SecureRandom, SystemRandom};
|
||||
use ring::signature::{ED25519, Ed25519KeyPair, KeyPair, UnparsedPublicKey};
|
||||
use serde::de::DeserializeOwned;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use sha2::{Digest, Sha256};
|
||||
use std::fmt;
|
||||
@@ -9,8 +10,6 @@ use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
const PUBLIC_KEY_PREFIX: &str = "yoi-ed25519-pub:v1:";
|
||||
const PRIVATE_KEY_PREFIX: &str = "yoi-ed25519-pkcs8:v1:";
|
||||
const TOKEN_PREFIX: &str = "yoi-cap-v1";
|
||||
const SIGNING_INPUT_PREFIX: &str = "yoi-cap-v1.";
|
||||
pub const WORKER_MUTATION_SOURCE_PROOF_HEADER: &str = "x-yoi-worker-mutation-proof";
|
||||
const WORKER_MUTATION_SOURCE_PROOF_PREFIX: &str = "yoi-worker-source-v1";
|
||||
const WORKER_MUTATION_SOURCE_SIGNING_INPUT_PREFIX: &str = "yoi-worker-source-v1.";
|
||||
@@ -68,6 +67,74 @@ pub enum RuntimeAuthError {
|
||||
WrongMutationTarget,
|
||||
}
|
||||
|
||||
pub(crate) struct SignedJsonToken<T> {
|
||||
pub payload: String,
|
||||
pub signature: Vec<u8>,
|
||||
pub claims: T,
|
||||
}
|
||||
|
||||
pub(crate) fn sign_json_token<T: Serialize>(
|
||||
token_prefix: &str,
|
||||
signing_input_prefix: &str,
|
||||
signing_key: &Ed25519KeyPair,
|
||||
claims: &T,
|
||||
) -> Result<String, RuntimeAuthError> {
|
||||
let payload = URL_SAFE_NO_PAD.encode(serde_json::to_vec(claims)?);
|
||||
let signing_input = format!("{signing_input_prefix}{payload}");
|
||||
let signature = signing_key.sign(signing_input.as_bytes());
|
||||
Ok(format!(
|
||||
"{token_prefix}.{payload}.{}",
|
||||
URL_SAFE_NO_PAD.encode(signature.as_ref())
|
||||
))
|
||||
}
|
||||
|
||||
pub(crate) fn decode_signed_json_token<T: DeserializeOwned>(
|
||||
token: &str,
|
||||
expected_prefix: &str,
|
||||
) -> Result<SignedJsonToken<T>, RuntimeAuthError> {
|
||||
let (prefix, payload, signature) = split_three_part_token(token)?;
|
||||
if prefix != expected_prefix {
|
||||
return Err(RuntimeAuthError::InvalidTokenFormat);
|
||||
}
|
||||
let signature = URL_SAFE_NO_PAD.decode(signature)?;
|
||||
let claims = serde_json::from_slice(&URL_SAFE_NO_PAD.decode(payload)?)?;
|
||||
Ok(SignedJsonToken {
|
||||
payload: payload.to_string(),
|
||||
signature,
|
||||
claims,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn verify_signed_json_token(
|
||||
signing_input_prefix: &str,
|
||||
payload: &str,
|
||||
signature: &[u8],
|
||||
public_key: &str,
|
||||
) -> Result<(), RuntimeAuthError> {
|
||||
let public_key = decode_public_key(public_key)?;
|
||||
let signing_input = format!("{signing_input_prefix}{payload}");
|
||||
UnparsedPublicKey::new(&ED25519, public_key)
|
||||
.verify(signing_input.as_bytes(), signature)
|
||||
.map_err(|_| RuntimeAuthError::InvalidSignature)
|
||||
}
|
||||
|
||||
fn split_three_part_token(token: &str) -> Result<(&str, &str, &str), RuntimeAuthError> {
|
||||
let mut parts = token.split('.');
|
||||
let prefix = parts.next().unwrap_or_default();
|
||||
let payload = parts.next().unwrap_or_default();
|
||||
let signature = parts.next().unwrap_or_default();
|
||||
if prefix.is_empty() || payload.is_empty() || signature.is_empty() || parts.next().is_some() {
|
||||
return Err(RuntimeAuthError::InvalidTokenFormat);
|
||||
}
|
||||
Ok((prefix, payload, signature))
|
||||
}
|
||||
|
||||
pub(crate) fn is_request_body_digest(value: &str) -> bool {
|
||||
URL_SAFE_NO_PAD
|
||||
.decode(value)
|
||||
.is_ok_and(|decoded| decoded.len() == 32 && URL_SAFE_NO_PAD.encode(decoded) == value)
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct RuntimeIdentityMaterial {
|
||||
pub identity_id: String,
|
||||
@@ -95,21 +162,6 @@ impl RuntimeIdentityMaterial {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct TrustedServerKey {
|
||||
pub server_id: String,
|
||||
pub public_key: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub display_name: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct RuntimeHttpAuthConfig {
|
||||
pub runtime_id: String,
|
||||
#[serde(default)]
|
||||
pub trusted_servers: Vec<TrustedServerKey>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct RuntimeAuthContext {
|
||||
pub server_id: String,
|
||||
@@ -119,122 +171,6 @@ pub struct RuntimeAuthContext {
|
||||
pub expires_at: u64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct CapabilityClaims {
|
||||
pub iss: String,
|
||||
pub aud: String,
|
||||
pub workspace_id: String,
|
||||
pub permissions: Vec<String>,
|
||||
pub exp: u64,
|
||||
pub jti: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct CapabilityTokenSigner {
|
||||
server_id: String,
|
||||
private_key: String,
|
||||
}
|
||||
|
||||
impl CapabilityTokenSigner {
|
||||
pub fn new(server_id: impl Into<String>, private_key: impl Into<String>) -> Self {
|
||||
Self {
|
||||
server_id: server_id.into(),
|
||||
private_key: private_key.into(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn server_id(&self) -> &str {
|
||||
&self.server_id
|
||||
}
|
||||
|
||||
pub fn sign(&self, claims: &CapabilityClaims) -> Result<String, RuntimeAuthError> {
|
||||
if claims.iss != self.server_id {
|
||||
return Err(RuntimeAuthError::UnknownIssuer(claims.iss.clone()));
|
||||
}
|
||||
let private = decode_private_key(&self.private_key)?;
|
||||
let pair = Ed25519KeyPair::from_pkcs8(&private)
|
||||
.map_err(|_| RuntimeAuthError::InvalidPrivateKey)?;
|
||||
let payload = serde_json::to_vec(claims)?;
|
||||
let payload = URL_SAFE_NO_PAD.encode(payload);
|
||||
let signing_input = format!("{SIGNING_INPUT_PREFIX}{payload}");
|
||||
let signature = pair.sign(signing_input.as_bytes());
|
||||
Ok(format!(
|
||||
"{TOKEN_PREFIX}.{payload}.{}",
|
||||
URL_SAFE_NO_PAD.encode(signature.as_ref())
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
pub fn capability_claims(
|
||||
server_id: impl Into<String>,
|
||||
runtime_id: impl Into<String>,
|
||||
workspace_id: impl Into<String>,
|
||||
permissions: Vec<String>,
|
||||
ttl_seconds: u64,
|
||||
) -> Result<CapabilityClaims, RuntimeAuthError> {
|
||||
let exp = unix_now_seconds().saturating_add(ttl_seconds);
|
||||
Ok(CapabilityClaims {
|
||||
iss: server_id.into(),
|
||||
aud: runtime_id.into(),
|
||||
workspace_id: workspace_id.into(),
|
||||
permissions,
|
||||
exp,
|
||||
jti: new_token_id()?,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn verify_capability_token(
|
||||
config: &RuntimeHttpAuthConfig,
|
||||
token: &str,
|
||||
required_permission: Option<&str>,
|
||||
now_seconds: u64,
|
||||
) -> Result<RuntimeAuthContext, RuntimeAuthError> {
|
||||
let (payload, signature) = split_token(token)?;
|
||||
let claims_json = URL_SAFE_NO_PAD.decode(payload)?;
|
||||
let claims: CapabilityClaims = serde_json::from_slice(&claims_json)?;
|
||||
let Some(server) = config
|
||||
.trusted_servers
|
||||
.iter()
|
||||
.find(|server| server.server_id == claims.iss)
|
||||
else {
|
||||
return Err(RuntimeAuthError::UnknownIssuer(claims.iss));
|
||||
};
|
||||
let public_key = decode_public_key(&server.public_key)?;
|
||||
let signing_input = format!("{SIGNING_INPUT_PREFIX}{payload}");
|
||||
UnparsedPublicKey::new(&ED25519, public_key)
|
||||
.verify(signing_input.as_bytes(), &signature)
|
||||
.map_err(|_| RuntimeAuthError::InvalidSignature)?;
|
||||
|
||||
if claims.aud != config.runtime_id {
|
||||
return Err(RuntimeAuthError::WrongAudience {
|
||||
expected: config.runtime_id.clone(),
|
||||
actual: claims.aud,
|
||||
});
|
||||
}
|
||||
if claims.exp < now_seconds {
|
||||
return Err(RuntimeAuthError::Expired);
|
||||
}
|
||||
if claims.workspace_id.trim().is_empty() {
|
||||
return Err(RuntimeAuthError::MissingWorkspaceScope);
|
||||
}
|
||||
if let Some(required) = required_permission {
|
||||
if !claims
|
||||
.permissions
|
||||
.iter()
|
||||
.any(|permission| permission == required)
|
||||
{
|
||||
return Err(RuntimeAuthError::MissingPermission(required.to_string()));
|
||||
}
|
||||
}
|
||||
Ok(RuntimeAuthContext {
|
||||
server_id: claims.iss,
|
||||
workspace_id: claims.workspace_id,
|
||||
permissions: claims.permissions,
|
||||
token_id: claims.jti,
|
||||
expires_at: claims.exp,
|
||||
})
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct RuntimeRequestSourceClaims {
|
||||
pub iss: String,
|
||||
@@ -323,28 +259,22 @@ impl RuntimeRequestSourceSigner {
|
||||
exp: now_unix.saturating_add(ttl_seconds),
|
||||
jti: new_token_id()?,
|
||||
};
|
||||
let payload = serde_json::to_vec(&claims)?;
|
||||
let payload = URL_SAFE_NO_PAD.encode(payload);
|
||||
let signing_input = format!("{RUNTIME_REQUEST_SOURCE_SIGNING_INPUT_PREFIX}{payload}");
|
||||
let private = decode_private_key(&self.private_key)?;
|
||||
let key_pair = Ed25519KeyPair::from_pkcs8(&private)
|
||||
.map_err(|_| RuntimeAuthError::InvalidPrivateKey)?;
|
||||
let signature = URL_SAFE_NO_PAD.encode(key_pair.sign(signing_input.as_bytes()).as_ref());
|
||||
Ok(format!(
|
||||
"{RUNTIME_REQUEST_SOURCE_PROOF_PREFIX}.{payload}.{signature}"
|
||||
))
|
||||
sign_json_token(
|
||||
RUNTIME_REQUEST_SOURCE_PROOF_PREFIX,
|
||||
RUNTIME_REQUEST_SOURCE_SIGNING_INPUT_PREFIX,
|
||||
&key_pair,
|
||||
&claims,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
pub fn decode_runtime_request_source_claims(
|
||||
proof: &str,
|
||||
) -> Result<RuntimeRequestSourceClaims, RuntimeAuthError> {
|
||||
let (prefix, payload, _signature) = split_runtime_request_source_proof(proof)?;
|
||||
if prefix != RUNTIME_REQUEST_SOURCE_PROOF_PREFIX {
|
||||
return Err(RuntimeAuthError::InvalidTokenFormat);
|
||||
}
|
||||
let payload = URL_SAFE_NO_PAD.decode(payload)?;
|
||||
serde_json::from_slice(&payload).map_err(RuntimeAuthError::from)
|
||||
Ok(decode_signed_json_token(proof, RUNTIME_REQUEST_SOURCE_PROOF_PREFIX)?.claims)
|
||||
}
|
||||
|
||||
pub fn verify_runtime_request_source(
|
||||
@@ -352,17 +282,17 @@ pub fn verify_runtime_request_source(
|
||||
public_key: &str,
|
||||
expected: &RuntimeRequestSourceExpectation<'_>,
|
||||
) -> Result<RuntimeRequestSourceClaims, RuntimeAuthError> {
|
||||
let (prefix, payload, signature) = split_runtime_request_source_proof(proof)?;
|
||||
if prefix != RUNTIME_REQUEST_SOURCE_PROOF_PREFIX {
|
||||
return Err(RuntimeAuthError::InvalidTokenFormat);
|
||||
}
|
||||
let signature = URL_SAFE_NO_PAD.decode(signature)?;
|
||||
let signing_input = format!("{RUNTIME_REQUEST_SOURCE_SIGNING_INPUT_PREFIX}{payload}");
|
||||
let public_key = decode_public_key(public_key)?;
|
||||
UnparsedPublicKey::new(&ED25519, public_key)
|
||||
.verify(signing_input.as_bytes(), &signature)
|
||||
.map_err(|_| RuntimeAuthError::InvalidSignature)?;
|
||||
let claims = decode_runtime_request_source_claims(proof)?;
|
||||
let signed = decode_signed_json_token::<RuntimeRequestSourceClaims>(
|
||||
proof,
|
||||
RUNTIME_REQUEST_SOURCE_PROOF_PREFIX,
|
||||
)?;
|
||||
verify_signed_json_token(
|
||||
RUNTIME_REQUEST_SOURCE_SIGNING_INPUT_PREFIX,
|
||||
&signed.payload,
|
||||
&signed.signature,
|
||||
public_key,
|
||||
)?;
|
||||
let claims = signed.claims;
|
||||
if claims.iss != expected.identity_id
|
||||
|| claims.aud != expected.audience
|
||||
|| claims.workspace_id != expected.workspace_id
|
||||
@@ -380,17 +310,6 @@ pub fn verify_runtime_request_source(
|
||||
Ok(claims)
|
||||
}
|
||||
|
||||
fn split_runtime_request_source_proof(proof: &str) -> Result<(&str, &str, &str), RuntimeAuthError> {
|
||||
let mut parts = proof.split('.');
|
||||
let prefix = parts.next().unwrap_or_default();
|
||||
let payload = parts.next().unwrap_or_default();
|
||||
let signature = parts.next().unwrap_or_default();
|
||||
if prefix.is_empty() || payload.is_empty() || signature.is_empty() || parts.next().is_some() {
|
||||
return Err(RuntimeAuthError::InvalidTokenFormat);
|
||||
}
|
||||
Ok((prefix, payload, signature))
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct WorkerMutationSourceClaims {
|
||||
pub iss: String,
|
||||
@@ -586,16 +505,6 @@ fn split_worker_mutation_source_proof(token: &str) -> Result<(&str, Vec<u8>), Ru
|
||||
}
|
||||
}
|
||||
|
||||
fn split_token(token: &str) -> Result<(&str, Vec<u8>), RuntimeAuthError> {
|
||||
let mut parts = token.split('.');
|
||||
match (parts.next(), parts.next(), parts.next(), parts.next()) {
|
||||
(Some(prefix), Some(payload), Some(signature), None) if prefix == TOKEN_PREFIX => {
|
||||
Ok((payload, URL_SAFE_NO_PAD.decode(signature)?))
|
||||
}
|
||||
_ => Err(RuntimeAuthError::InvalidTokenFormat),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn encode_public_key(bytes: &[u8]) -> String {
|
||||
format!("{PUBLIC_KEY_PREFIX}{}", URL_SAFE_NO_PAD.encode(bytes))
|
||||
}
|
||||
@@ -851,46 +760,4 @@ mod tests {
|
||||
Err(RuntimeAuthError::Expired)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn capability_token_verifies_signature_audience_expiry_and_permission() {
|
||||
let server = RuntimeIdentityMaterial::generate("server-main").unwrap();
|
||||
let signer = CapabilityTokenSigner::new(&server.identity_id, &server.private_key);
|
||||
let claims = CapabilityClaims {
|
||||
iss: "server-main".to_string(),
|
||||
aud: "runtime-main".to_string(),
|
||||
workspace_id: "workspace-a".to_string(),
|
||||
permissions: vec!["workers:list".to_string()],
|
||||
exp: 100,
|
||||
jti: "token-1".to_string(),
|
||||
};
|
||||
let token = signer.sign(&claims).unwrap();
|
||||
let auth = RuntimeHttpAuthConfig {
|
||||
runtime_id: "runtime-main".to_string(),
|
||||
trusted_servers: vec![TrustedServerKey {
|
||||
server_id: "server-main".to_string(),
|
||||
public_key: server.public_key.clone(),
|
||||
display_name: None,
|
||||
}],
|
||||
};
|
||||
|
||||
let context = verify_capability_token(&auth, &token, Some("workers:list"), 99).unwrap();
|
||||
assert_eq!(context.workspace_id, "workspace-a");
|
||||
assert!(matches!(
|
||||
verify_capability_token(&auth, &token, Some("workers:create"), 99),
|
||||
Err(RuntimeAuthError::MissingPermission(permission)) if permission == "workers:create"
|
||||
));
|
||||
assert!(matches!(
|
||||
verify_capability_token(&auth, &token, Some("workers:list"), 101),
|
||||
Err(RuntimeAuthError::Expired)
|
||||
));
|
||||
let wrong_audience = RuntimeHttpAuthConfig {
|
||||
runtime_id: "other-runtime".to_string(),
|
||||
trusted_servers: auth.trusted_servers.clone(),
|
||||
};
|
||||
assert!(matches!(
|
||||
verify_capability_token(&wrong_audience, &token, Some("workers:list"), 99),
|
||||
Err(RuntimeAuthError::WrongAudience { .. })
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -15,32 +15,22 @@ pub enum ProfileSelector {
|
||||
Named(String),
|
||||
}
|
||||
|
||||
/// Runtime fetch/caching metadata for a Backend-authored Decodal profile source archive.
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct ProfileSourceArchiveHttpRef {
|
||||
pub url: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub etag: Option<String>,
|
||||
pub archive: ProfileSourceArchiveRef,
|
||||
}
|
||||
|
||||
/// Profile source material available to a Runtime during Worker creation.
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(tag = "kind", rename_all = "snake_case")]
|
||||
pub enum ProfileSourceArchiveSource {
|
||||
/// Backend-internal embedded runtimes may receive already-built archive bytes.
|
||||
Embedded { archive: ProfileSourceArchive },
|
||||
/// Standalone runtimes fetch/cache the tar archive over HTTP.
|
||||
Http {
|
||||
location: ProfileSourceArchiveHttpRef,
|
||||
},
|
||||
/// Standalone runtimes resolve this immutable archive from the latest
|
||||
/// Workspace Config bundle before creating the Worker.
|
||||
WorkspaceConfig { archive: ProfileSourceArchiveRef },
|
||||
}
|
||||
|
||||
impl ProfileSourceArchiveSource {
|
||||
pub fn reference(&self) -> ProfileSourceArchiveRef {
|
||||
match self {
|
||||
Self::Embedded { archive } => archive.reference.clone(),
|
||||
Self::Http { location } => location.archive.clone(),
|
||||
Self::WorkspaceConfig { archive } => archive.clone(),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -92,9 +82,9 @@ pub struct WorkingDirectoryRepository {
|
||||
}
|
||||
|
||||
pub use workdir::workspace::{
|
||||
MaterializerKind, WorkingDirectoryCleanupTarget, WorkingDirectoryCurrentObservation,
|
||||
MaterializerKind, RuntimeWorkingDirectoryCleanupTarget as WorkingDirectoryCleanupTarget,
|
||||
RuntimeWorkingDirectorySummary as WorkingDirectorySummary, WorkingDirectoryCurrentObservation,
|
||||
WorkingDirectoryOccupancy, WorkingDirectoryProvenance, WorkingDirectoryStatusKind,
|
||||
WorkingDirectorySummary,
|
||||
};
|
||||
|
||||
#[derive(Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
@@ -129,9 +119,16 @@ impl std::fmt::Debug for SensitiveString {
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct RepositorySshMaterializationAccess {
|
||||
pub struct RepositorySshCredentialCandidate {
|
||||
pub credential_id: String,
|
||||
pub credential_revision: u64,
|
||||
#[serde(skip, default)]
|
||||
pub private_key: SensitiveString,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct RepositorySshMaterializationAccess {
|
||||
pub credential_candidates: Vec<RepositorySshCredentialCandidate>,
|
||||
pub host_trust_id: String,
|
||||
pub host_trust_revision: u64,
|
||||
pub access: workspace_api::RepositoryAccessMode,
|
||||
@@ -141,8 +138,6 @@ pub struct RepositorySshMaterializationAccess {
|
||||
pub repository_uri: String,
|
||||
pub secret_resource: crate::resource::BackendResourceHandle,
|
||||
#[serde(skip, default)]
|
||||
pub private_key: SensitiveString,
|
||||
#[serde(skip, default)]
|
||||
pub known_hosts_entry: SensitiveString,
|
||||
}
|
||||
|
||||
@@ -153,8 +148,6 @@ pub struct RepositoryMaterializationContext {
|
||||
pub operation_id: String,
|
||||
pub config_revision: u64,
|
||||
pub config_projection_digest: String,
|
||||
#[serde(default)]
|
||||
pub cache_generation: u64,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub ssh: Option<RepositorySshMaterializationAccess>,
|
||||
}
|
||||
@@ -179,6 +172,30 @@ pub struct WorkingDirectoryRequest {
|
||||
pub materialization: Option<RepositoryMaterializationContext>,
|
||||
}
|
||||
|
||||
/// Backend-authorized request to freshly resolve one Repository provider ref.
|
||||
///
|
||||
/// Runtime executes this against the registered source itself rather than a Workdir
|
||||
/// or Runtime cache. Secret material is fetched through `materialization` and never
|
||||
/// appears in the result.
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct RepositoryRefObservationRequest {
|
||||
pub repository: WorkingDirectoryRepository,
|
||||
pub selector: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub materialization: Option<RepositoryMaterializationContext>,
|
||||
}
|
||||
|
||||
/// Provider-neutral proof of one freshly observed Repository ref.
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct RepositoryRefObservation {
|
||||
pub repository_id: String,
|
||||
pub source_revision: u64,
|
||||
pub source_fingerprint: String,
|
||||
pub selector: String,
|
||||
pub revision_ref: String,
|
||||
pub observed_at_epoch_seconds: u64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct WorkingDirectoryClaim {
|
||||
pub working_directory_id: String,
|
||||
@@ -250,6 +267,10 @@ pub struct CreateWorkerRequest {
|
||||
}
|
||||
|
||||
/// Worker lifecycle status for the in-memory embedded runtime.
|
||||
///
|
||||
/// Run termination details are carried separately by the Worker protocol. In
|
||||
/// particular, cancellation returns a Worker to `Idle`; it is not a lifecycle
|
||||
/// state of its own.
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum WorkerStatus {
|
||||
@@ -257,7 +278,6 @@ pub enum WorkerStatus {
|
||||
Running,
|
||||
Paused,
|
||||
Stopped,
|
||||
Cancelled,
|
||||
}
|
||||
|
||||
impl WorkerStatus {
|
||||
@@ -266,6 +286,13 @@ impl WorkerStatus {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub(crate) enum WorkerRestoreIntent {
|
||||
Automatic,
|
||||
Explicit,
|
||||
}
|
||||
|
||||
/// Lightweight catalog row.
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct WorkerSummary {
|
||||
@@ -273,6 +300,8 @@ pub struct WorkerSummary {
|
||||
pub worker_id: WorkerId,
|
||||
pub status: WorkerStatus,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub worker_state: Option<protocol::WorkerStateSnapshot>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub workspace_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub working_directory: Option<WorkingDirectoryStatus>,
|
||||
@@ -291,6 +320,8 @@ pub struct WorkerDetail {
|
||||
pub worker_id: WorkerId,
|
||||
pub status: WorkerStatus,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub worker_state: Option<protocol::WorkerStateSnapshot>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub workspace_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub working_directory: Option<WorkingDirectoryStatus>,
|
||||
@@ -307,6 +338,8 @@ pub struct WorkerDetail {
|
||||
pub struct WorkerLifecycleAck {
|
||||
pub worker_ref: WorkerRef,
|
||||
pub status: WorkerStatus,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub worker_state: Option<protocol::WorkerStateSnapshot>,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
||||
@@ -9,6 +9,11 @@ use serde::{Deserialize, Serialize};
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
pub const CONFIG_BUNDLE_DIGEST_ALGORITHM: &str = "sha256";
|
||||
pub const WORKSPACE_CONFIG_ETAG_PREFIX: &str = "workspace-config:";
|
||||
|
||||
pub fn workspace_config_etag(digest: &str) -> String {
|
||||
format!("\"{WORKSPACE_CONFIG_ETAG_PREFIX}{digest}\"")
|
||||
}
|
||||
|
||||
/// Backend-synced Profile/config bundle stored by a Runtime.
|
||||
///
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
use crate::catalog::{
|
||||
ConfigBundleRef, ProfileSelector, RepositoryRefObservation, RepositoryRefObservationRequest,
|
||||
WorkingDirectoryRepositoryAccessRequest, WorkingDirectoryRequest, WorkingDirectoryStatus,
|
||||
WorkspaceApiRef,
|
||||
};
|
||||
use crate::config_bundle::ConfigBundle;
|
||||
use crate::error::RuntimeError;
|
||||
@@ -8,24 +10,12 @@ use crate::interaction::WorkerInput;
|
||||
#[cfg(feature = "ws-server")]
|
||||
use crate::observation::WorkerObservationEvent;
|
||||
use crate::working_directory::{WorkingDirectoryBinding, WorkingDirectoryDiagnostic};
|
||||
use protocol::Method;
|
||||
use protocol::{Method, UploadedFileRef};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::fmt;
|
||||
use std::sync::Arc;
|
||||
use workdir::WorkdirSessionHandle;
|
||||
|
||||
/// Current execution-side run state for a Worker.
|
||||
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum WorkerExecutionRunState {
|
||||
#[default]
|
||||
Stopped,
|
||||
Idle,
|
||||
Busy,
|
||||
Rejected,
|
||||
Errored,
|
||||
}
|
||||
|
||||
/// Execution operation that produced a result.
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
@@ -33,19 +23,19 @@ pub enum WorkerExecutionOperation {
|
||||
Spawn,
|
||||
Restore,
|
||||
Input,
|
||||
UploadFile,
|
||||
DeleteUploadedFile,
|
||||
ProtocolMethod,
|
||||
Stop,
|
||||
Cancel,
|
||||
}
|
||||
|
||||
/// Evidence that a user input reached the durable Worker session boundary.
|
||||
///
|
||||
/// This is intentionally distinct from accepting a method on the Worker's
|
||||
/// in-memory channel. For Flow submissions, the committed UserInput entry also
|
||||
/// carries the initial Flow runtime-state extension.
|
||||
/// Evidence that a Submit request reached the durable Worker session boundary.
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct WorkerInputCommitAck {
|
||||
pub struct WorkerSubmissionAck {
|
||||
pub submission_request_id: String,
|
||||
pub submission_id: String,
|
||||
pub disposition: protocol::SubmissionDisposition,
|
||||
}
|
||||
|
||||
/// Typed execution result class. Results are transient operation outcomes and
|
||||
@@ -54,11 +44,12 @@ pub struct WorkerInputCommitAck {
|
||||
pub struct WorkerExecutionResult {
|
||||
pub operation: WorkerExecutionOperation,
|
||||
pub outcome: WorkerExecutionOutcome,
|
||||
pub run_state: WorkerExecutionRunState,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub worker_state: Option<protocol::WorkerStateSnapshot>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub message: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_commit: Option<WorkerInputCommitAck>,
|
||||
pub submission: Option<WorkerSubmissionAck>,
|
||||
}
|
||||
|
||||
/// Backend result class for a Worker execution operation.
|
||||
@@ -73,31 +64,36 @@ pub enum WorkerExecutionOutcome {
|
||||
}
|
||||
|
||||
impl WorkerExecutionResult {
|
||||
pub fn accepted(
|
||||
operation: WorkerExecutionOperation,
|
||||
run_state: WorkerExecutionRunState,
|
||||
) -> Self {
|
||||
pub fn accepted(operation: WorkerExecutionOperation) -> Self {
|
||||
Self {
|
||||
operation,
|
||||
outcome: WorkerExecutionOutcome::Accepted,
|
||||
run_state,
|
||||
worker_state: None,
|
||||
message: None,
|
||||
input_commit: None,
|
||||
submission: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn accepted_input_committed(
|
||||
pub fn with_worker_state(mut self, worker_state: protocol::WorkerStateSnapshot) -> Self {
|
||||
self.worker_state = Some(worker_state);
|
||||
self
|
||||
}
|
||||
|
||||
pub fn accepted_submission(
|
||||
operation: WorkerExecutionOperation,
|
||||
run_state: WorkerExecutionRunState,
|
||||
submission_request_id: impl Into<String>,
|
||||
submission_id: impl Into<String>,
|
||||
disposition: protocol::SubmissionDisposition,
|
||||
) -> Self {
|
||||
Self {
|
||||
operation,
|
||||
outcome: WorkerExecutionOutcome::Accepted,
|
||||
run_state,
|
||||
worker_state: None,
|
||||
message: None,
|
||||
input_commit: Some(WorkerInputCommitAck {
|
||||
submission: Some(WorkerSubmissionAck {
|
||||
submission_request_id: submission_request_id.into(),
|
||||
submission_id: submission_id.into(),
|
||||
disposition,
|
||||
}),
|
||||
}
|
||||
}
|
||||
@@ -106,9 +102,9 @@ impl WorkerExecutionResult {
|
||||
Self {
|
||||
operation,
|
||||
outcome: WorkerExecutionOutcome::Busy,
|
||||
run_state: WorkerExecutionRunState::Busy,
|
||||
worker_state: None,
|
||||
message: Some(message.into()),
|
||||
input_commit: None,
|
||||
submission: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -116,9 +112,9 @@ impl WorkerExecutionResult {
|
||||
Self {
|
||||
operation,
|
||||
outcome: WorkerExecutionOutcome::Rejected,
|
||||
run_state: WorkerExecutionRunState::Stopped,
|
||||
worker_state: None,
|
||||
message: Some(message.into()),
|
||||
input_commit: None,
|
||||
submission: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -126,9 +122,9 @@ impl WorkerExecutionResult {
|
||||
Self {
|
||||
operation,
|
||||
outcome: WorkerExecutionOutcome::Errored,
|
||||
run_state: WorkerExecutionRunState::Errored,
|
||||
worker_state: None,
|
||||
message: Some(message.into()),
|
||||
input_commit: None,
|
||||
submission: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -136,9 +132,9 @@ impl WorkerExecutionResult {
|
||||
Self {
|
||||
operation,
|
||||
outcome: WorkerExecutionOutcome::Unsupported,
|
||||
run_state: WorkerExecutionRunState::Stopped,
|
||||
worker_state: None,
|
||||
message: Some(message.into()),
|
||||
input_commit: None,
|
||||
submission: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -270,12 +266,28 @@ pub struct WorkerExecutionRestoreRequest {
|
||||
pub config_bundle: Option<ConfigBundle>,
|
||||
}
|
||||
|
||||
/// Runtime-side request to refresh the latest Workspace Config before Worker creation.
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct WorkspaceConfigFetchRequest {
|
||||
pub workspace_api: WorkspaceApiRef,
|
||||
pub profile: ProfileSelector,
|
||||
pub expected: ConfigBundleRef,
|
||||
pub cached: Option<ConfigBundleRef>,
|
||||
}
|
||||
|
||||
/// Result of a conditional Workspace Config fetch.
|
||||
#[derive(Clone, Debug)]
|
||||
pub enum WorkspaceConfigFetchResult {
|
||||
NotModified,
|
||||
Modified(ConfigBundle),
|
||||
}
|
||||
|
||||
/// Backend outcome for Worker spawn/restore operations.
|
||||
#[derive(Clone, Debug)]
|
||||
pub enum WorkerExecutionSpawnResult {
|
||||
Connected {
|
||||
handle: WorkerExecutionHandle,
|
||||
run_state: WorkerExecutionRunState,
|
||||
worker_state: protocol::WorkerStateSnapshot,
|
||||
working_directory: Option<WorkingDirectoryStatus>,
|
||||
},
|
||||
Rejected(WorkerExecutionResult),
|
||||
@@ -285,12 +297,12 @@ pub enum WorkerExecutionSpawnResult {
|
||||
impl WorkerExecutionSpawnResult {
|
||||
pub fn connected(
|
||||
handle: WorkerExecutionHandle,
|
||||
run_state: WorkerExecutionRunState,
|
||||
worker_state: protocol::WorkerStateSnapshot,
|
||||
working_directory: Option<WorkingDirectoryStatus>,
|
||||
) -> Self {
|
||||
Self::Connected {
|
||||
handle,
|
||||
run_state,
|
||||
worker_state,
|
||||
working_directory,
|
||||
}
|
||||
}
|
||||
@@ -299,6 +311,13 @@ impl WorkerExecutionSpawnResult {
|
||||
pub trait WorkerExecutionBackend: Send + Sync + 'static {
|
||||
fn backend_id(&self) -> &str;
|
||||
|
||||
fn fetch_workspace_config(
|
||||
&self,
|
||||
_request: WorkspaceConfigFetchRequest,
|
||||
) -> Result<WorkspaceConfigFetchResult, String> {
|
||||
Err("execution backend does not support Workspace Config fetching".to_string())
|
||||
}
|
||||
|
||||
fn spawn_worker(&self, request: WorkerExecutionSpawnRequest) -> WorkerExecutionSpawnResult;
|
||||
|
||||
fn restore_worker(
|
||||
@@ -331,6 +350,16 @@ pub trait WorkerExecutionBackend: Send + Sync + 'static {
|
||||
))
|
||||
}
|
||||
|
||||
fn observe_repository_ref(
|
||||
&self,
|
||||
_request: &RepositoryRefObservationRequest,
|
||||
) -> Result<RepositoryRefObservation, WorkingDirectoryDiagnostic> {
|
||||
Err(WorkingDirectoryDiagnostic::rejected(
|
||||
"repository_ref_provider_unavailable",
|
||||
"Worker execution backend does not support Repository ref observation",
|
||||
))
|
||||
}
|
||||
|
||||
fn list_working_directories(&self) -> Vec<WorkingDirectoryStatus> {
|
||||
Vec::new()
|
||||
}
|
||||
@@ -385,6 +414,31 @@ pub trait WorkerExecutionBackend: Send + Sync + 'static {
|
||||
input: WorkerInput,
|
||||
) -> WorkerExecutionResult;
|
||||
|
||||
fn upload_file(
|
||||
&self,
|
||||
_handle: &WorkerExecutionHandle,
|
||||
_file_name: &str,
|
||||
_media_type: &str,
|
||||
_content: &[u8],
|
||||
_context: Option<&session_store::UploadedFileUploadContext>,
|
||||
) -> Result<UploadedFileRef, WorkerExecutionResult> {
|
||||
Err(WorkerExecutionResult::unsupported(
|
||||
WorkerExecutionOperation::UploadFile,
|
||||
"execution backend does not support file upload",
|
||||
))
|
||||
}
|
||||
|
||||
fn delete_uploaded_file(
|
||||
&self,
|
||||
_handle: &WorkerExecutionHandle,
|
||||
_artifact_id: &str,
|
||||
) -> WorkerExecutionResult {
|
||||
WorkerExecutionResult::unsupported(
|
||||
WorkerExecutionOperation::DeleteUploadedFile,
|
||||
"execution backend does not support uploaded-file deletion",
|
||||
)
|
||||
}
|
||||
|
||||
fn dispatch_method(
|
||||
&self,
|
||||
_handle: &WorkerExecutionHandle,
|
||||
@@ -445,6 +499,13 @@ impl WorkerExecutionBackendRef {
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn fetch_workspace_config(
|
||||
&self,
|
||||
request: WorkspaceConfigFetchRequest,
|
||||
) -> Result<WorkspaceConfigFetchResult, String> {
|
||||
self.backend.fetch_workspace_config(request)
|
||||
}
|
||||
|
||||
pub(crate) fn spawn_worker(
|
||||
&self,
|
||||
request: WorkerExecutionSpawnRequest,
|
||||
@@ -474,6 +535,13 @@ impl WorkerExecutionBackendRef {
|
||||
.authorize_working_directory_repository_access(request)
|
||||
}
|
||||
|
||||
pub(crate) fn observe_repository_ref(
|
||||
&self,
|
||||
request: &RepositoryRefObservationRequest,
|
||||
) -> Result<RepositoryRefObservation, WorkingDirectoryDiagnostic> {
|
||||
self.backend.observe_repository_ref(request)
|
||||
}
|
||||
|
||||
pub(crate) fn list_working_directories(&self) -> Vec<WorkingDirectoryStatus> {
|
||||
self.backend.list_working_directories()
|
||||
}
|
||||
@@ -514,6 +582,26 @@ impl WorkerExecutionBackendRef {
|
||||
self.backend.dispatch_input(handle, input)
|
||||
}
|
||||
|
||||
pub(crate) fn upload_file(
|
||||
&self,
|
||||
handle: &WorkerExecutionHandle,
|
||||
file_name: &str,
|
||||
media_type: &str,
|
||||
content: &[u8],
|
||||
context: Option<&session_store::UploadedFileUploadContext>,
|
||||
) -> Result<UploadedFileRef, WorkerExecutionResult> {
|
||||
self.backend
|
||||
.upload_file(handle, file_name, media_type, content, context)
|
||||
}
|
||||
|
||||
pub(crate) fn delete_uploaded_file(
|
||||
&self,
|
||||
handle: &WorkerExecutionHandle,
|
||||
artifact_id: &str,
|
||||
) -> WorkerExecutionResult {
|
||||
self.backend.delete_uploaded_file(handle, artifact_id)
|
||||
}
|
||||
|
||||
pub(crate) fn dispatch_method(
|
||||
&self,
|
||||
handle: &WorkerExecutionHandle,
|
||||
@@ -553,14 +641,16 @@ mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn input_commit_ack_survives_json_round_trip() {
|
||||
let result = WorkerExecutionResult::accepted_input_committed(
|
||||
fn submission_ack_survives_json_round_trip() {
|
||||
let result = WorkerExecutionResult::accepted_submission(
|
||||
WorkerExecutionOperation::Input,
|
||||
WorkerExecutionRunState::Busy,
|
||||
"request-1",
|
||||
"submission-1",
|
||||
protocol::SubmissionDisposition::Started,
|
||||
);
|
||||
|
||||
let json = serde_json::to_string(&result).unwrap();
|
||||
assert!(json.contains("\"submission_request_id\":\"request-1\""));
|
||||
assert!(json.contains("\"submission_id\":\"submission-1\""));
|
||||
assert_eq!(
|
||||
serde_json::from_str::<WorkerExecutionResult>(&json).unwrap(),
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
use crate::catalog::{CreateWorkerRequest, WorkingDirectoryStatus};
|
||||
use crate::catalog::{
|
||||
CreateWorkerRequest, WorkerRestoreIntent, WorkerStatus, WorkingDirectoryStatus,
|
||||
};
|
||||
use crate::config_bundle::ConfigBundle;
|
||||
use crate::diagnostics::{DiagnosticSeverity, RuntimeDiagnostic};
|
||||
use crate::error::RuntimeError;
|
||||
@@ -13,7 +15,10 @@ use std::io::{BufReader, Write};
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
|
||||
const SCHEMA_VERSION: u32 = 3;
|
||||
const SCHEMA_VERSION: u32 = 6;
|
||||
const PREVIOUS_SCHEMA_VERSION: u32 = 5;
|
||||
const EXECUTION_SCHEMA_VERSION: u32 = 4;
|
||||
const PRE_EXECUTION_SCHEMA_VERSION: u32 = 3;
|
||||
const RUNTIME_FILE: &str = "runtime.json";
|
||||
const WORKERS_DIR: &str = "workers";
|
||||
const WORKER_FILE: &str = "worker.json";
|
||||
@@ -274,13 +279,25 @@ pub(crate) struct PersistedRuntimeState {
|
||||
pub(crate) diagnostics: Vec<RuntimeDiagnostic>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub(crate) struct PersistedWorkerExecutionBinding {
|
||||
pub(crate) run_generation: u64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub(crate) struct PersistedWorkerExecution {
|
||||
pub(crate) last_run_generation: u64,
|
||||
pub(crate) binding: Option<PersistedWorkerExecutionBinding>,
|
||||
pub(crate) restore_intent: WorkerRestoreIntent,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(crate) struct PersistedWorkerRecord {
|
||||
pub(crate) worker_ref: WorkerRef,
|
||||
pub(crate) worker_id: WorkerId,
|
||||
pub(crate) request: CreateWorkerRequest,
|
||||
/// Last generation durably reserved for this Worker's execution.
|
||||
pub(crate) run_generation: u64,
|
||||
pub(crate) status: WorkerStatus,
|
||||
pub(crate) execution: PersistedWorkerExecution,
|
||||
pub(crate) workspace_id: Option<String>,
|
||||
pub(crate) working_directory: Option<WorkingDirectoryStatus>,
|
||||
}
|
||||
@@ -357,8 +374,8 @@ fn plan_runtime_store_migration(
|
||||
format!("Runtime store schema version {schema_version} is out of range"),
|
||||
)
|
||||
})?;
|
||||
let staging = migration_sibling(root, "schema-v3-staging")?;
|
||||
let backup = migration_sibling(root, "pre-schema-v3-backup")?;
|
||||
let staging = migration_sibling(root, "schema-v6-staging")?;
|
||||
let backup = migration_sibling(root, "pre-schema-v6-backup")?;
|
||||
if staging.exists() || backup.exists() {
|
||||
return Err(runtime_store_corrupt(
|
||||
root,
|
||||
@@ -384,11 +401,14 @@ fn plan_runtime_store_migration(
|
||||
};
|
||||
return Ok((plan, Vec::new()));
|
||||
}
|
||||
if !matches!(current_schema_version, 1 | 2) {
|
||||
if !matches!(
|
||||
current_schema_version,
|
||||
PRE_EXECUTION_SCHEMA_VERSION | EXECUTION_SCHEMA_VERSION | PREVIOUS_SCHEMA_VERSION
|
||||
) {
|
||||
return Err(runtime_store_corrupt(
|
||||
&runtime_path,
|
||||
format!(
|
||||
"unsupported Runtime store schema version {schema_version}; expected 1, 2, or {SCHEMA_VERSION}"
|
||||
"unsupported Runtime store schema version {schema_version}; expected {PRE_EXECUTION_SCHEMA_VERSION}, {EXECUTION_SCHEMA_VERSION}, {PREVIOUS_SCHEMA_VERSION}, or {SCHEMA_VERSION}"
|
||||
),
|
||||
));
|
||||
}
|
||||
@@ -415,6 +435,16 @@ fn plan_runtime_store_migration(
|
||||
runtime_store_corrupt(&source_dir, "Worker directory is not UTF-8".to_string())
|
||||
})?;
|
||||
let snapshot_path = source_dir.join(WORKER_FILE);
|
||||
if !snapshot_path
|
||||
.try_exists()
|
||||
.map_err(|source| RuntimeError::StoreIo {
|
||||
operation: "inspect Worker snapshot",
|
||||
path: snapshot_path.clone(),
|
||||
source,
|
||||
})?
|
||||
{
|
||||
continue;
|
||||
}
|
||||
let snapshot: serde_json::Value = read_json(&snapshot_path, "read Worker snapshot")?;
|
||||
let (worker_id, workspace_id, legacy_mapping) = if current_schema_version == 1 {
|
||||
let legacy_worker_id = name.parse::<u64>().map_err(|_| {
|
||||
@@ -448,7 +478,7 @@ fn plan_runtime_store_migration(
|
||||
let worker_id = name.parse::<WorkerId>().map_err(|_| {
|
||||
runtime_store_corrupt(
|
||||
&source_dir,
|
||||
format!("schema-v2 Worker directory name must be a UUIDv7, found {name}"),
|
||||
format!("pre-v4 Worker directory name must be a UUIDv7, found {name}"),
|
||||
)
|
||||
})?;
|
||||
(worker_id, None, None)
|
||||
@@ -603,6 +633,38 @@ fn migrate_v1_worker_document(
|
||||
Ok(snapshot)
|
||||
}
|
||||
|
||||
fn max_persisted_run_generation(snapshot_path: &Path) -> Result<u64, RuntimeError> {
|
||||
let worker_dir = snapshot_path.parent().ok_or_else(|| {
|
||||
runtime_store_corrupt(
|
||||
snapshot_path,
|
||||
"Worker snapshot path is missing its aggregate directory".to_string(),
|
||||
)
|
||||
})?;
|
||||
let runs_dir = worker_dir.join("runs");
|
||||
if !runs_dir
|
||||
.try_exists()
|
||||
.map_err(|source| runtime_io_error("inspect Worker runs", &runs_dir, source))?
|
||||
{
|
||||
return Ok(0);
|
||||
}
|
||||
let entries = fs::read_dir(&runs_dir)
|
||||
.map_err(|source| runtime_io_error("read Worker runs", &runs_dir, source))?;
|
||||
let mut max_generation = 0;
|
||||
for entry in entries {
|
||||
let entry =
|
||||
entry.map_err(|source| runtime_io_error("read Worker runs", &runs_dir, source))?;
|
||||
let Some(generation) = entry
|
||||
.file_name()
|
||||
.to_str()
|
||||
.and_then(|name| name.parse::<u64>().ok())
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
max_generation = max_generation.max(generation);
|
||||
}
|
||||
Ok(max_generation)
|
||||
}
|
||||
|
||||
fn migrate_worker_document(
|
||||
mut document: serde_json::Value,
|
||||
source_schema_version: u32,
|
||||
@@ -610,7 +672,7 @@ fn migrate_worker_document(
|
||||
snapshot_path: &Path,
|
||||
) -> Result<serde_json::Value, RuntimeError> {
|
||||
if source_schema_version == 1 {
|
||||
return migrate_v1_worker_document(
|
||||
document = migrate_v1_worker_document(
|
||||
document,
|
||||
mapping.ok_or_else(|| {
|
||||
runtime_store_corrupt(
|
||||
@@ -619,7 +681,7 @@ fn migrate_worker_document(
|
||||
)
|
||||
})?,
|
||||
snapshot_path,
|
||||
);
|
||||
)?;
|
||||
}
|
||||
let object = document.as_object_mut().ok_or_else(|| {
|
||||
runtime_store_corrupt(
|
||||
@@ -627,10 +689,117 @@ fn migrate_worker_document(
|
||||
"Worker snapshot must be an object".to_string(),
|
||||
)
|
||||
})?;
|
||||
let declared_run_generation = object
|
||||
.remove("run_generation")
|
||||
.map(|value| {
|
||||
value.as_u64().ok_or_else(|| {
|
||||
runtime_store_corrupt(
|
||||
snapshot_path,
|
||||
"Worker snapshot run_generation must be an unsigned integer".to_string(),
|
||||
)
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let legacy_execution = object.remove("execution");
|
||||
let execution = legacy_execution
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object);
|
||||
let persisted_last_run_generation = execution
|
||||
.and_then(|execution| execution.get("last_run_generation"))
|
||||
.map(|value| {
|
||||
value.as_u64().ok_or_else(|| {
|
||||
runtime_store_corrupt(
|
||||
snapshot_path,
|
||||
"Worker execution last_run_generation must be an unsigned integer".to_string(),
|
||||
)
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let binding_run_generation = execution
|
||||
.and_then(|execution| execution.get("binding"))
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|binding| binding.get("run_generation"))
|
||||
.map(|value| {
|
||||
value.as_u64().ok_or_else(|| {
|
||||
runtime_store_corrupt(
|
||||
snapshot_path,
|
||||
"Worker execution binding run_generation must be an unsigned integer"
|
||||
.to_string(),
|
||||
)
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let run_generation = declared_run_generation
|
||||
.into_iter()
|
||||
.chain(persisted_last_run_generation)
|
||||
.chain(binding_run_generation)
|
||||
.chain(std::iter::once(max_persisted_run_generation(
|
||||
snapshot_path,
|
||||
)?))
|
||||
.max()
|
||||
.unwrap_or(0);
|
||||
if !object.contains_key("working_directory") {
|
||||
if let Some(working_directory) = legacy_execution
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|execution| execution.get("working_directory"))
|
||||
.cloned()
|
||||
{
|
||||
object.insert("working_directory".to_string(), working_directory);
|
||||
}
|
||||
}
|
||||
let legacy_materialization = object
|
||||
.get("working_directory")
|
||||
.and_then(|working_directory| working_directory.get("summary"))
|
||||
.and_then(|summary| summary.get("materializer_kind"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.is_some_and(|kind| matches!(kind, "runtime_git_cache" | "local_git_worktree"));
|
||||
if legacy_materialization {
|
||||
object.insert("working_directory".to_string(), serde_json::Value::Null);
|
||||
}
|
||||
if let Some(profile_source) = object
|
||||
.get_mut("request")
|
||||
.and_then(serde_json::Value::as_object_mut)
|
||||
.and_then(|request| request.get_mut("profile_source"))
|
||||
.and_then(serde_json::Value::as_object_mut)
|
||||
&& profile_source
|
||||
.get("kind")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
== Some("http")
|
||||
{
|
||||
let archive = profile_source
|
||||
.get_mut("location")
|
||||
.and_then(serde_json::Value::as_object_mut)
|
||||
.and_then(|location| location.remove("archive"))
|
||||
.ok_or_else(|| {
|
||||
runtime_store_corrupt(
|
||||
snapshot_path,
|
||||
"legacy HTTP profile source is missing its archive".to_string(),
|
||||
)
|
||||
})?;
|
||||
profile_source.clear();
|
||||
profile_source.insert(
|
||||
"kind".to_string(),
|
||||
serde_json::Value::String("workspace_config".to_string()),
|
||||
);
|
||||
profile_source.insert("archive".to_string(), archive);
|
||||
}
|
||||
object.insert(
|
||||
"schema_version".to_string(),
|
||||
serde_json::Value::from(SCHEMA_VERSION),
|
||||
);
|
||||
object.insert(
|
||||
"status".to_string(),
|
||||
serde_json::Value::String("stopped".to_string()),
|
||||
);
|
||||
object.insert(
|
||||
"execution".to_string(),
|
||||
serde_json::json!({
|
||||
"last_run_generation": run_generation,
|
||||
"binding": null,
|
||||
"restore_intent": "explicit",
|
||||
}),
|
||||
);
|
||||
Ok(document)
|
||||
}
|
||||
|
||||
@@ -710,8 +879,8 @@ fn migrate_worker_aggregate_document(
|
||||
.get_mut("resolved_manifest_snapshot")
|
||||
.filter(|snapshot| !snapshot.is_null())
|
||||
{
|
||||
let manifest: manifest::WorkerManifest =
|
||||
serde_json::from_value(snapshot.clone()).map_err(|error| {
|
||||
let mut manifest = manifest::read_persisted_worker_manifest_snapshot(snapshot.clone())
|
||||
.map_err(|error| {
|
||||
runtime_store_corrupt(
|
||||
metadata_path,
|
||||
format!("decode Worker aggregate resolved manifest snapshot: {error}"),
|
||||
@@ -726,20 +895,14 @@ fn migrate_worker_aggregate_document(
|
||||
),
|
||||
));
|
||||
}
|
||||
snapshot
|
||||
.as_object_mut()
|
||||
.and_then(|manifest| manifest.get_mut("worker"))
|
||||
.and_then(serde_json::Value::as_object_mut)
|
||||
.ok_or_else(|| {
|
||||
manifest.worker.name = expected_name.clone();
|
||||
*snapshot =
|
||||
manifest::write_persisted_worker_manifest_snapshot(&manifest).map_err(|error| {
|
||||
runtime_store_corrupt(
|
||||
metadata_path,
|
||||
"Worker aggregate resolved manifest is missing worker metadata".to_string(),
|
||||
format!("encode migrated Worker aggregate resolved manifest: {error}"),
|
||||
)
|
||||
})?
|
||||
.insert(
|
||||
"name".to_string(),
|
||||
serde_json::Value::String(expected_name.clone()),
|
||||
);
|
||||
})?;
|
||||
}
|
||||
metadata.insert(
|
||||
"worker_name".to_string(),
|
||||
@@ -760,8 +923,8 @@ fn migrate_worker_aggregate_document(
|
||||
));
|
||||
}
|
||||
if let Some(snapshot) = metadata.resolved_manifest_snapshot {
|
||||
let manifest: manifest::WorkerManifest =
|
||||
serde_json::from_value(snapshot).map_err(|error| {
|
||||
let manifest =
|
||||
manifest::read_persisted_worker_manifest_snapshot(snapshot).map_err(|error| {
|
||||
runtime_store_corrupt(
|
||||
metadata_path,
|
||||
format!("decode migrated Worker aggregate resolved manifest: {error}"),
|
||||
@@ -1005,8 +1168,8 @@ fn migrate_runtime_store(
|
||||
if !plan.migration_required {
|
||||
return Ok(plan);
|
||||
}
|
||||
let staging = migration_sibling(root, "schema-v3-staging")?;
|
||||
let backup = migration_sibling(root, "pre-schema-v3-backup")?;
|
||||
let staging = migration_sibling(root, "schema-v6-staging")?;
|
||||
let backup = migration_sibling(root, "pre-schema-v6-backup")?;
|
||||
if staging.exists() || backup.exists() {
|
||||
return Err(runtime_store_corrupt(
|
||||
root,
|
||||
@@ -1236,22 +1399,12 @@ struct WorkerSnapshot {
|
||||
worker_ref: WorkerRef,
|
||||
worker_id: WorkerId,
|
||||
request: CreateWorkerRequest,
|
||||
#[serde(default)]
|
||||
run_generation: u64,
|
||||
status: WorkerStatus,
|
||||
execution: PersistedWorkerExecution,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
workspace_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
working_directory: Option<WorkingDirectoryStatus>,
|
||||
/// One-way migration input for schema-v1 snapshots. New snapshots never
|
||||
/// write the removed execution projection.
|
||||
#[serde(default, rename = "execution", skip_serializing)]
|
||||
legacy_execution: Option<LegacyWorkerExecutionProjection>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
struct LegacyWorkerExecutionProjection {
|
||||
#[serde(default)]
|
||||
working_directory: Option<WorkingDirectoryStatus>,
|
||||
}
|
||||
|
||||
impl WorkerSnapshot {
|
||||
@@ -1261,10 +1414,10 @@ impl WorkerSnapshot {
|
||||
worker_ref: worker.worker_ref.clone(),
|
||||
worker_id: worker.worker_id.clone(),
|
||||
request: worker.request.clone(),
|
||||
run_generation: worker.run_generation,
|
||||
status: worker.status,
|
||||
execution: worker.execution.clone(),
|
||||
workspace_id: worker.workspace_id.clone(),
|
||||
working_directory: worker.working_directory.clone(),
|
||||
legacy_execution: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1289,6 +1442,63 @@ impl WorkerSnapshot {
|
||||
),
|
||||
});
|
||||
}
|
||||
if let Some(binding) = self.execution.binding.as_ref()
|
||||
&& binding.run_generation != self.execution.last_run_generation
|
||||
{
|
||||
return Err(RuntimeError::StoreCorrupt {
|
||||
operation: "read worker snapshot",
|
||||
path: path.to_path_buf(),
|
||||
message: format!(
|
||||
"execution binding run_generation {} does not match last_run_generation {}",
|
||||
binding.run_generation, self.execution.last_run_generation
|
||||
),
|
||||
});
|
||||
}
|
||||
match (self.status, self.execution.restore_intent) {
|
||||
(status, WorkerRestoreIntent::Automatic) if status.is_active() => {
|
||||
let Some(binding) = self.execution.binding.as_ref() else {
|
||||
return Err(RuntimeError::StoreCorrupt {
|
||||
operation: "read worker snapshot",
|
||||
path: path.to_path_buf(),
|
||||
message: "automatic restore intent requires an execution binding"
|
||||
.to_string(),
|
||||
});
|
||||
};
|
||||
if binding.run_generation == 0 {
|
||||
return Err(RuntimeError::StoreCorrupt {
|
||||
operation: "read worker snapshot",
|
||||
path: path.to_path_buf(),
|
||||
message: "execution binding run_generation must be greater than zero"
|
||||
.to_string(),
|
||||
});
|
||||
}
|
||||
}
|
||||
(WorkerStatus::Stopped, WorkerRestoreIntent::Explicit) => {
|
||||
if self
|
||||
.execution
|
||||
.binding
|
||||
.as_ref()
|
||||
.is_some_and(|binding| binding.run_generation == 0)
|
||||
{
|
||||
return Err(RuntimeError::StoreCorrupt {
|
||||
operation: "read worker snapshot",
|
||||
path: path.to_path_buf(),
|
||||
message: "execution binding run_generation must be greater than zero"
|
||||
.to_string(),
|
||||
});
|
||||
}
|
||||
}
|
||||
_ => {
|
||||
return Err(RuntimeError::StoreCorrupt {
|
||||
operation: "read worker snapshot",
|
||||
path: path.to_path_buf(),
|
||||
message: format!(
|
||||
"worker status {:?} conflicts with restore intent {:?}",
|
||||
self.status, self.execution.restore_intent
|
||||
),
|
||||
});
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -1303,12 +1513,10 @@ impl WorkerSnapshot {
|
||||
worker_ref: self.worker_ref,
|
||||
worker_id: self.worker_id,
|
||||
request: self.request,
|
||||
run_generation: self.run_generation,
|
||||
status: self.status,
|
||||
execution: self.execution,
|
||||
workspace_id,
|
||||
working_directory: self.working_directory.or_else(|| {
|
||||
self.legacy_execution
|
||||
.and_then(|execution| execution.working_directory)
|
||||
}),
|
||||
working_directory: self.working_directory,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1415,3 +1623,138 @@ fn sync_directory(path: &Path, operation: &'static str) -> Result<(), RuntimeErr
|
||||
source,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn schema_v4_migration_plan_ignores_orphan_worker_directories() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
fs::write(
|
||||
root.path().join(RUNTIME_FILE),
|
||||
serde_json::to_vec_pretty(&serde_json::json!({
|
||||
"schema_version": PREVIOUS_SCHEMA_VERSION,
|
||||
"display_name": null,
|
||||
"backend": "fs_store",
|
||||
"status": "running",
|
||||
"next_diagnostic_id": 1,
|
||||
"config_bundles": {},
|
||||
"workspace_owners": {},
|
||||
"diagnostics": []
|
||||
}))
|
||||
.unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
fs::create_dir_all(root.path().join(WORKERS_DIR).join("orphan").join("session")).unwrap();
|
||||
fs::write(
|
||||
root.path()
|
||||
.join(WORKERS_DIR)
|
||||
.join("orphan")
|
||||
.join("session")
|
||||
.join("history.json"),
|
||||
b"[]",
|
||||
)
|
||||
.unwrap();
|
||||
let (plan, _) = plan_runtime_store_migration(root.path(), "runtime-test").unwrap();
|
||||
|
||||
assert!(plan.migration_required);
|
||||
assert_eq!(plan.current_schema_version, PREVIOUS_SCHEMA_VERSION);
|
||||
assert_eq!(plan.target_schema_version, SCHEMA_VERSION);
|
||||
assert_eq!(plan.worker_count, 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn schema_v5_worker_migration_recovers_last_generation_from_run_aggregates() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let worker_dir = root.path().join("worker-a");
|
||||
fs::create_dir_all(worker_dir.join("runs/1")).unwrap();
|
||||
fs::create_dir_all(worker_dir.join("runs/7")).unwrap();
|
||||
fs::create_dir_all(worker_dir.join("runs/incomplete")).unwrap();
|
||||
let path = worker_dir.join(WORKER_FILE);
|
||||
let source = serde_json::json!({
|
||||
"schema_version": 5,
|
||||
"execution": {
|
||||
"binding": null,
|
||||
"restore_intent": "explicit"
|
||||
}
|
||||
});
|
||||
|
||||
let migrated =
|
||||
migrate_worker_document(source, PREVIOUS_SCHEMA_VERSION, None, &path).unwrap();
|
||||
|
||||
assert_eq!(
|
||||
migrated["execution"]["last_run_generation"],
|
||||
serde_json::json!(7)
|
||||
);
|
||||
assert_eq!(migrated["execution"]["binding"], serde_json::Value::Null);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn schema_v4_worker_migration_discards_unsupported_linked_worktree_binding() {
|
||||
let source = serde_json::json!({
|
||||
"schema_version": 4,
|
||||
"request": {
|
||||
"profile_source": {
|
||||
"kind": "http",
|
||||
"location": {
|
||||
"url": "https://workspace.example.test/archive",
|
||||
"etag": "profile-source:test",
|
||||
"archive": {
|
||||
"id": "profiles-v1",
|
||||
"digest": "sha256:test",
|
||||
"size_bytes": 1,
|
||||
"source_graph": {
|
||||
"source_count": 1,
|
||||
"total_source_bytes": 1,
|
||||
"entrypoints": {},
|
||||
"import_count": 0
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"working_directory": {
|
||||
"summary": {
|
||||
"materializer_kind": "runtime_git_cache"
|
||||
}
|
||||
}
|
||||
});
|
||||
let path = Path::new("worker.json");
|
||||
|
||||
let migrated =
|
||||
migrate_worker_document(source, EXECUTION_SCHEMA_VERSION, None, path).unwrap();
|
||||
|
||||
assert_eq!(migrated["schema_version"], SCHEMA_VERSION);
|
||||
assert_eq!(migrated["status"], "stopped");
|
||||
assert_eq!(migrated["working_directory"], serde_json::Value::Null);
|
||||
assert_eq!(
|
||||
migrated["request"]["profile_source"]["kind"],
|
||||
"workspace_config"
|
||||
);
|
||||
assert_eq!(
|
||||
migrated["request"]["profile_source"]["archive"]["id"],
|
||||
"profiles-v1"
|
||||
);
|
||||
assert_eq!(migrated["execution"]["restore_intent"], "explicit");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn schema_v4_worker_migration_preserves_runtime_clone_observation() {
|
||||
let source = serde_json::json!({
|
||||
"schema_version": 4,
|
||||
"working_directory": {
|
||||
"summary": {
|
||||
"materializer_kind": "runtime_git_clone"
|
||||
}
|
||||
}
|
||||
});
|
||||
let expected = source["working_directory"].clone();
|
||||
let path = Path::new("worker.json");
|
||||
|
||||
let migrated =
|
||||
migrate_worker_document(source, EXECUTION_SCHEMA_VERSION, None, path).unwrap();
|
||||
|
||||
assert_eq!(migrated["working_directory"], expected);
|
||||
}
|
||||
}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user