Skip to content

Commit ceb727e

Browse files
authored
fix(subagents): harden roster lifecycle (#56)
1 parent 48b9e65 commit ceb727e

8 files changed

Lines changed: 2392 additions & 1425 deletions

File tree

Cargo.lock

Lines changed: 2 additions & 1 deletion
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

Cargo.toml

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
[package]
22
name = "kit"
3-
version = "0.1.109"
3+
version = "0.1.110"
44
edition = "2024"
55
rust-version = "1.94.0"
66
publish = false
@@ -66,6 +66,7 @@ tower = { version = "=0.5.3", features = ["util"] }
6666
tracing = "=0.1.44"
6767
tracing-opentelemetry = { version = "=0.33.0", default-features = false }
6868
tracing-subscriber = { version = "=0.3.23", default-features = false, features = ["registry", "std"] }
69+
unicode-segmentation = "=1.13.3"
6970
unicode-width = "=0.2.2"
7071
url = "=2.5.8"
7172
uuid = { version = "=1.26.0", features = ["v5"] }

fixtures/mock-acp.py

Lines changed: 98 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -6,10 +6,29 @@
66
import time
77

88
write_lock = threading.Lock()
9+
log_lock = threading.Lock()
10+
state_lock = threading.Lock()
911
next_session = 1
1012
supports_fork = "--no-fork" not in sys.argv
1113
supports_models = "--models" in sys.argv
14+
fail_first_close = "--fail-first-close" in sys.argv
15+
failed_close = False
1216
selected_models = {}
17+
18+
19+
def option(name):
20+
prefix = name + "="
21+
return next((arg[len(prefix):] for arg in sys.argv if arg.startswith(prefix)), None)
22+
23+
24+
request_log = option("--request-log")
25+
new_release = option("--new-release")
26+
fork_release = option("--fork-release")
27+
prompt_release = option("--prompt-release")
28+
prompt_release_text = option("--prompt-release-text")
29+
close_release = option("--close-release")
30+
close_release_session = option("--close-release-session")
31+
fail_close_session = option("--fail-close-session")
1332
model_ids = ["mock/default", "mock/requested"]
1433

1534
if "--fail-start" in sys.argv:
@@ -26,11 +45,58 @@ def respond(request_id, result):
2645
send({"jsonrpc": "2.0", "id": request_id, "result": result})
2746

2847

48+
def log_request(request):
49+
if request_log is None:
50+
return
51+
params = request.get("params", {})
52+
entry = {"method": request.get("method")}
53+
if "sessionId" in params:
54+
entry["sessionId"] = params["sessionId"]
55+
if request.get("method") == "session/prompt":
56+
entry["text"] = params["prompt"][0]["text"]
57+
with log_lock:
58+
with open(request_log, "a", encoding="utf-8") as log:
59+
log.write(json.dumps(entry, separators=(",", ":")) + "\n")
60+
log.flush()
61+
62+
63+
def fork(request):
64+
global next_session
65+
if "--fail-fork" in sys.argv:
66+
send({
67+
"jsonrpc": "2.0",
68+
"id": request["id"],
69+
"error": {"code": -32000, "message": "fork failed"},
70+
})
71+
return
72+
if not supports_fork:
73+
send({
74+
"jsonrpc": "2.0",
75+
"id": request["id"],
76+
"error": {"code": -32601, "message": "Method not found"},
77+
})
78+
return
79+
source_id = request["params"]["sessionId"]
80+
with state_lock:
81+
session_id = f"branch-{next_session}"
82+
next_session += 1
83+
selected_models[session_id] = selected_models.get(source_id, model_ids[0])
84+
while fork_release is not None and not os.path.exists(fork_release):
85+
time.sleep(0.01)
86+
respond(request["id"], {"sessionId": session_id})
87+
88+
2989
def prompt(request):
3090
params = request["params"]
3191
session_id = params["sessionId"]
3292
text = params["prompt"][0]["text"]
33-
time.sleep(0.40)
93+
should_gate = prompt_release is not None and (
94+
prompt_release_text is None or prompt_release_text == text
95+
)
96+
while should_gate and not os.path.exists(prompt_release):
97+
time.sleep(0.01)
98+
if prompt_release is None:
99+
time.sleep(0.40)
34100
if "MOCK_SELECTED_MODEL" in text:
35101
text = selected_models.get(session_id, model_ids[0])
36102
if "MOCK_STRUCTURED_OUTPUT" in text:
@@ -94,9 +160,36 @@ def prompt(request):
94160
os._exit(0)
95161

96162

163+
def close(request):
164+
global failed_close
165+
if "--slow-close" in sys.argv:
166+
time.sleep(0.40)
167+
session_id = request["params"]["sessionId"]
168+
should_gate = close_release is not None and (
169+
close_release_session is None or close_release_session == session_id
170+
)
171+
while should_gate and not os.path.exists(close_release):
172+
time.sleep(0.01)
173+
with state_lock:
174+
should_fail = (fail_close_session == session_id and not failed_close) or (
175+
fail_first_close and not failed_close
176+
)
177+
if should_fail:
178+
failed_close = True
179+
if should_fail:
180+
send({
181+
"jsonrpc": "2.0",
182+
"id": request["id"],
183+
"error": {"code": -32000, "message": "close failed"},
184+
})
185+
else:
186+
respond(request["id"], {})
187+
188+
97189
for line in sys.stdin:
98190
request = json.loads(line)
99191
method = request.get("method")
192+
log_request(request)
100193
if method == "initialize":
101194
respond(request["id"], {
102195
"protocolVersion": 1,
@@ -107,6 +200,8 @@ def prompt(request):
107200
},
108201
})
109202
elif method == "session/new":
203+
while new_release is not None and not os.path.exists(new_release):
204+
time.sleep(0.01)
110205
selected_models["base"] = model_ids[0]
111206
result = {"sessionId": "base"}
112207
if supports_models:
@@ -122,25 +217,7 @@ def prompt(request):
122217
}]
123218
respond(request["id"], result)
124219
elif method == "session/fork":
125-
if "--fail-fork" in sys.argv:
126-
send({
127-
"jsonrpc": "2.0",
128-
"id": request["id"],
129-
"error": {"code": -32000, "message": "fork failed"},
130-
})
131-
continue
132-
if not supports_fork:
133-
send({
134-
"jsonrpc": "2.0",
135-
"id": request["id"],
136-
"error": {"code": -32601, "message": "Method not found"},
137-
})
138-
continue
139-
session_id = f"branch-{next_session}"
140-
next_session += 1
141-
source_id = request["params"]["sessionId"]
142-
selected_models[session_id] = selected_models.get(source_id, model_ids[0])
143-
respond(request["id"], {"sessionId": session_id})
220+
threading.Thread(target=fork, args=(request,), daemon=True).start()
144221
elif method == "session/prompt":
145222
threading.Thread(target=prompt, args=(request,), daemon=True).start()
146223
elif method == "session/set_config_option":
@@ -158,6 +235,4 @@ def prompt(request):
158235
elif method == "session/cancel":
159236
pass
160237
elif method == "session/close":
161-
if "--slow-close" in sys.argv:
162-
time.sleep(0.40)
163-
respond(request["id"], {})
238+
threading.Thread(target=close, args=(request,), daemon=True).start()

src/acp_child.rs

Lines changed: 13 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -275,8 +275,7 @@ impl AcpHarnesses {
275275
command
276276
.arg("--subagent-parent-id")
277277
.arg(parent_id)
278-
.arg("--subagent-parent-name")
279-
.arg(parent_name);
278+
.arg(format!("--subagent-parent-name={parent_name}"));
280279
}
281280
if resume {
282281
command.arg("--resume");
@@ -2000,10 +1999,18 @@ mod tests {
20001999
args.windows(2)
20012000
.any(|pair| pair == ["--subagent-parent-id", "s-parent"])
20022001
);
2003-
assert!(
2004-
args.windows(2)
2005-
.any(|pair| pair == ["--subagent-parent-name", "偵察 🦀"])
2006-
);
2002+
assert!(args.contains(&"--subagent-parent-name=偵察 🦀".into()));
2003+
}
2004+
2005+
#[test]
2006+
fn kit_child_parent_name_is_safe_when_it_starts_with_a_hyphen() {
2007+
let config = config(Some("s-parent"), Some("--reviewer"));
2008+
let command = config
2009+
.harnesses
2010+
.spawn(BUILTIN_HARNESS, &config, Some(("session", false)), 1)
2011+
.unwrap();
2012+
2013+
assert!(args(&command).contains(&"--subagent-parent-name=--reviewer".into()));
20072014
}
20082015

20092016
#[test]

0 commit comments

Comments
 (0)