66import time
77
88write_lock = threading .Lock ()
9+ log_lock = threading .Lock ()
10+ state_lock = threading .Lock ()
911next_session = 1
1012supports_fork = "--no-fork" not in sys .argv
1113supports_models = "--models" in sys .argv
14+ fail_first_close = "--fail-first-close" in sys .argv
15+ failed_close = False
1216selected_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" )
1332model_ids = ["mock/default" , "mock/requested" ]
1433
1534if "--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+
2989def 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+
97189for 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 ()
0 commit comments