@@ -240,11 +240,15 @@ def _watch_context(
240240 context ,
241241 cancel_event : threading .Event ,
242242 session_id : str ,
243+ completed_event : threading .Event | None = None ,
243244 ) -> None :
244245 callback = getattr (context , "add_done_callback" , None )
245246 if callback is not None :
246247 def done (done_context ) -> None :
247- if done_context .cancelled ():
248+ logically_completed = (
249+ completed_event is not None and completed_event .is_set ()
250+ )
251+ if done_context .cancelled () and not logically_completed :
248252 cancel_event .set ()
249253 self ._store .remove_session_if_present (
250254 session_id , reason = "client_cancelled" ,
@@ -305,7 +309,13 @@ async def AppendTokens( # noqa: N802 — gRPC-generated method casing
305309 f"session { request .session_id !r} already has an active operation" ,
306310 )
307311 cancel_event = threading .Event ()
308- self ._watch_context (context , cancel_event , request .session_id )
312+ completed_event = threading .Event ()
313+ self ._watch_context (
314+ context ,
315+ cancel_event ,
316+ request .session_id ,
317+ completed_event ,
318+ )
309319 try :
310320 kwargs = {
311321 "session_id" : request .session_id ,
@@ -316,6 +326,7 @@ async def AppendTokens( # noqa: N802 — gRPC-generated method casing
316326 new_history_length = await asyncio .to_thread (
317327 self ._append .append_tokens , ** kwargs ,
318328 )
329+ completed_event .set ()
319330 except asyncio .CancelledError :
320331 cancel_event .set ()
321332 self ._store .remove_session_if_present (
@@ -390,7 +401,13 @@ async def Generate( # noqa: N802 — gRPC-generated method casing
390401 f"session { request .session_id !r} already has an active operation" ,
391402 )
392403 cancel_event = threading .Event ()
393- self ._watch_context (context , cancel_event , request .session_id )
404+ completed_event = threading .Event ()
405+ self ._watch_context (
406+ context ,
407+ cancel_event ,
408+ request .session_id ,
409+ completed_event ,
410+ )
394411 kwargs = {
395412 "session_id" : request .session_id ,
396413 "max_tokens" : request .max_tokens ,
@@ -413,6 +430,7 @@ async def Generate( # noqa: N802 — gRPC-generated method casing
413430 if error is not None :
414431 raise error
415432 if event is _STREAM_END :
433+ completed_event .set ()
416434 break
417435 if context .cancelled ():
418436 threaded_stream .cancel ()
@@ -436,6 +454,7 @@ async def Generate( # noqa: N802 — gRPC-generated method casing
436454 # DoneEvent — the only remaining event type per
437455 # the GenerateEvent union.
438456 assert isinstance (event , DoneEvent )
457+ completed_event .set ()
439458 yield runtime_pb2 .GenerateResponse (
440459 done = runtime_pb2 .GenerateDone (
441460 stop_reason = _STOP_REASON_TO_PROTO [
@@ -456,9 +475,10 @@ async def Generate( # noqa: N802 — gRPC-generated method casing
456475 await context .abort (grpc .StatusCode .ABORTED , str (exc ))
457476 except asyncio .CancelledError :
458477 threaded_stream .cancel ()
459- self ._store .remove_session_if_present (
460- request .session_id , reason = "client_cancelled" ,
461- )
478+ if not completed_event .is_set ():
479+ self ._store .remove_session_if_present (
480+ request .session_id , reason = "client_cancelled" ,
481+ )
462482 raise
463483 finally :
464484 threaded_stream .cancel ()
0 commit comments