Fix: abort handling for dispatched requests after client disconnect (#35255)

Signed-off-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com>
Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
Co-authored-by: cctry <cctry@fb.com>
This commit is contained in:
Shijin Zhang
2026-09-03 21:51:57 -07:00
committed by GitHub
co-authored by Xinyuan Tong cctry
parent e787de5478
commit f478b2bb2d
6 changed files with 297 additions and 48 deletions
+5
View File
@@ -3349,6 +3349,11 @@ class Scheduler(
if req.finished() or not req.kv.holds_kv: if req.finished() or not req.kv.holds_kv:
self._pending_chunked_abort_req = None self._pending_chunked_abort_req = None
return return
# The request moved to another scheduler queue after abort_request
# deferred it, so retry against its current location.
self._pending_chunked_abort_req = None
self.abort_request(AbortReq(rid=req.rid))
return
prepare_abort(req, "Aborted") prepare_abort(req, "Aborted")
req.time_stats.trace_ctx.abort(abort_info={"reason": "Aborted"}) req.time_stats.trace_ctx.abort(abort_info={"reason": "Aborted"})
+58 -24
View File
@@ -233,6 +233,9 @@ class ReqState:
last_completion_tokens: int = 1 last_completion_tokens: int = 1
ttft_observed: bool = False ttft_observed: bool = False
dispatched: bool = False
abort_sent: bool = False
# For streaming output # For streaming output
last_output_offset: int = 0 last_output_offset: int = 0
@@ -799,6 +802,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
) )
self._init_req_state(obj, request) self._init_req_state(obj, request)
request_rids = {obj.rid} if obj.is_single else set(obj.rid)
try: try:
if get_disagg().language_only: if get_disagg().language_only:
self._handle_epd_disaggregation_encode_request(obj) self._handle_epd_disaggregation_encode_request(obj)
@@ -822,17 +826,19 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
async for response in self._wait_one_response(obj, request): async for response in self._wait_one_response(obj, request):
yield response yield response
else: else:
async for response in self._handle_batch_request(obj, request): async for response in self._handle_batch_request(
obj, request, request_rids
):
yield response yield response
except BaseException: except BaseException:
# _init_req_state created a rid_to_state entry per (sub-)request up # _init_req_state created a rid_to_state entry per (sub-)request up
# front. The normal remover is the scheduler-response path # front. The normal remover is the scheduler-response path
# (_handle_batch_output), so a failure *before* a request reaches the # (_handle_batch_output), so a failure *before* a request reaches the
# scheduler -- e.g. input-length validation rejecting an over-context # scheduler -- e.g. input-length validation rejecting an over-context
# request -- would otherwise leak those entries forever. Drop any that # request -- would otherwise leak those entries forever. Drop
# are still pending; entries already removed on the normal completion # undelivered states, but abort dispatched requests for scheduler-side
# path are left untouched (pop is a no-op). # cleanup.
self._discard_pending_req_states(obj) self._release_req_states_on_failure(request_rids)
raise raise
def _detect_input_format( def _detect_input_format(
@@ -1571,6 +1577,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
time_stats = tokenized_obj.time_stats time_stats = tokenized_obj.time_stats
tokenized_obj.wrap_pickle_fields() tokenized_obj.wrap_pickle_fields()
self._dispatch_to_scheduler(tokenized_obj) self._dispatch_to_scheduler(tokenized_obj)
self._mark_state_dispatched(tokenized_obj.rid)
dispatched = True dispatched = True
tokenized_obj.time_stats = time_stats tokenized_obj.time_stats = time_stats
tokenized_obj.time_stats.set_api_server_dispatch_finish_time() tokenized_obj.time_stats.set_api_server_dispatch_finish_time()
@@ -1578,6 +1585,16 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if not dispatched: if not dispatched:
self.cuda_vmm_feature_transport.cancel_for_dispatch(prepared_mm_items) self.cuda_vmm_feature_transport.cancel_for_dispatch(prepared_mm_items)
def _mark_state_dispatched(self, rid: str):
"""Record that *rid* reached the scheduler.
Only dispatched requests are aborted (not discarded) by the
handler-failure cleanup; see _release_req_states_on_failure.
"""
state = self.rid_to_state.get(rid)
if state is not None:
state.dispatched = True
async def _send_batch_request( async def _send_batch_request(
self, self,
tokenized_objs: List[ tokenized_objs: List[
@@ -1605,6 +1622,8 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs) batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs)
self._dispatch_to_scheduler(batch_req) self._dispatch_to_scheduler(batch_req)
for tokenized_obj in tokenized_objs:
self._mark_state_dispatched(tokenized_obj.rid)
dispatched = True dispatched = True
for tokenized_obj, time_stat in zip(tokenized_objs, time_stats): for tokenized_obj, time_stat in zip(tokenized_objs, time_stats):
tokenized_obj.time_stats = time_stat tokenized_obj.time_stats = time_stat
@@ -1819,7 +1838,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self, self,
obj: Union[GenerateReqInput, EmbeddingReqInput], obj: Union[GenerateReqInput, EmbeddingReqInput],
request: Optional[fastapi.Request] = None, request: Optional[fastapi.Request] = None,
request_rids: Optional[set[str]] = None,
): ):
if request_rids is None:
request_rids = set(obj.rid)
batch_size = obj.batch_size batch_size = obj.batch_size
generators = [] generators = []
@@ -1885,6 +1907,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
tokenized_obj.sampling_params.max_new_tokens = 0 tokenized_obj.sampling_params.max_new_tokens = 0
tokenized_obj.stream = False tokenized_obj.stream = False
self._init_req_state(tmp_obj) self._init_req_state(tmp_obj)
request_rids.add(tmp_obj.rid)
await self._send_one_request(tokenized_obj) await self._send_one_request(tokenized_obj)
await self._wait_one_response(tmp_obj, request).__anext__() await self._wait_one_response(tmp_obj, request).__anext__()
@@ -1901,6 +1924,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
] ]
tokenized_obj.rid = tmp_obj.regenerate_rid() tokenized_obj.rid = tmp_obj.regenerate_rid()
self._init_req_state(tmp_obj) self._init_req_state(tmp_obj)
request_rids.add(tmp_obj.rid)
state = self.rid_to_state[tmp_obj.rid] state = self.rid_to_state[tmp_obj.rid]
tokenized_obj.time_stats = state.time_stats tokenized_obj.time_stats = state.time_stats
if tmp_obj.return_prompt_token_ids: if tmp_obj.return_prompt_token_ids:
@@ -1970,14 +1994,21 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if not abort_all and not rid: if not abort_all and not rid:
logger.warning("Ignore abort_request with empty rid and abort_all=False") logger.warning("Ignore abort_request with empty rid and abort_all=False")
return return
if ( state = None if abort_all else self.rid_to_state.get(rid)
not abort_all if not abort_all:
and get_serving().tokenizer_worker_num == 1 if state is not None:
and rid not in self.rid_to_state if state.abort_sent:
): return
state.abort_sent = True
elif get_serving().tokenizer_worker_num == 1:
return return
req = AbortReq(rid=rid, abort_all=abort_all) req = AbortReq(rid=rid, abort_all=abort_all)
try:
self._dispatch_to_scheduler(req) self._dispatch_to_scheduler(req)
except BaseException:
if state is not None:
state.abort_sent = False
raise
if self.enable_metrics: if self.enable_metrics:
# TODO: also use custom_labels from the request # TODO: also use custom_labels from the request
self.metrics_collector.observe_one_aborted_request( self.metrics_collector.observe_one_aborted_request(
@@ -2140,10 +2171,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Abort the request if the client is disconnected. # Abort the request if the client is disconnected.
async def abort_request(): async def abort_request():
await asyncio.sleep(2) await asyncio.sleep(2)
if obj.is_single: rids = [obj.rid] if obj.is_single else obj.rid
self.abort_request(obj.rid) for rid in rids:
else: if rid in self.rid_to_state:
for rid in obj.rid:
self.abort_request(rid) self.abort_request(rid)
background_tasks = BackgroundTasks() background_tasks = BackgroundTasks()
@@ -3455,19 +3485,23 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
time_stats.init_trace_ctx(rid, bootstrap_room, external_trace_header) time_stats.init_trace_ctx(rid, bootstrap_room, external_trace_header)
time_stats.set_created_time(created_time) time_stats.set_created_time(created_time)
def _discard_pending_req_states(self, obj): def _release_req_states_on_failure(self, rids: Iterable[str]):
"""Drop rid_to_state entries created by _init_req_state for *obj*. """Release rid_to_state entries created for a failed handler.
Safe to call after a partial/failed dispatch: only entries still present Undelivered states are removed locally. Dispatched requests are aborted
are removed, and the scheduler-response path looks up state with and retained until the scheduler response removes them.
``.get(...)`` so a later output for a discarded rid is ignored, not fatal.
""" """
if not hasattr(obj, "is_single") or obj.is_single:
rids = [obj.rid]
else:
rids = obj.rid
for rid in rids: for rid in rids:
self.rid_to_state.pop(rid, None) state = self.rid_to_state.get(rid)
if state is None:
continue
if state.dispatched:
try:
self.abort_request(rid)
except Exception:
logger.exception("Failed to abort request %s during cleanup", rid)
else:
del self.rid_to_state[rid]
def _should_dispatch_to_encoder( def _should_dispatch_to_encoder(
self, obj: Union[GenerateReqInput, EmbeddingReqInput] self, obj: Union[GenerateReqInput, EmbeddingReqInput]
@@ -0,0 +1,69 @@
"""Tests for deferred chunked-prefill aborts."""
import unittest
from types import SimpleNamespace
from unittest.mock import Mock
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
maybe_stub_sgl_kernel()
from sglang.srt.managers.scheduler import Scheduler # noqa: E402
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
class _FakeReq:
"""Minimal stand-in for Req: only the fields the abort paths touch."""
def __init__(self, rid: str):
self.rid = rid
# Mirrors Req.kv; the abort paths read only these two predicates.
self.kv = SimpleNamespace(holds_kv=True, holds_mamba=False)
self.to_finish = None
self._finished = False
def finished(self):
return self._finished
def _make_scheduler(pending_req, *, chunked_req, running_reqs) -> Scheduler:
sched = Scheduler.__new__(Scheduler)
sched.chunked_req = chunked_req
sched._pending_chunked_abort_req = pending_req
sched.waiting_queue = []
sched.dllm_config = None
sched.grammar_manager = Mock()
sched.disaggregation_mode = None
sched.enable_hicache_storage = False
sched.mm_receiver = None
sched.ps = SimpleNamespace(pp_size=1)
sched.running_batch = SimpleNamespace(reqs=running_reqs)
sched.last_batch = None
return sched
class TestPendingChunkedAbortRace(CustomTestCase):
def test_req_left_chunked_slot_is_aborted(self):
req = _FakeReq("zombie_rid")
sched = _make_scheduler(req, chunked_req=None, running_reqs=[req])
sched.process_pending_chunked_abort()
self.assertIsNotNone(req.to_finish, "recorded abort was never applied")
self.assertIsNone(sched._pending_chunked_abort_req)
def test_finished_req_only_clears_marker(self):
req = _FakeReq("done_rid")
req._finished = True
sched = _make_scheduler(req, chunked_req=None, running_reqs=[])
sched.process_pending_chunked_abort()
self.assertIsNone(req.to_finish)
self.assertIsNone(sched._pending_chunked_abort_req)
if __name__ == "__main__":
unittest.main(verbosity=2)
@@ -10,11 +10,12 @@ Covers:
- _handle_batch_output cleans up rid_to_state on finished requests - _handle_batch_output cleans up rid_to_state on finished requests
- _init_req_state rejects duplicate rids - _init_req_state rejects duplicate rids
- Resubmission succeeds after cleanup - Resubmission succeeds after cleanup
- Handler failures clean up pending and dispatched requests
""" """
import asyncio import asyncio
import unittest import unittest
from unittest.mock import AsyncMock, MagicMock, Mock from unittest.mock import AsyncMock, MagicMock, Mock, patch
import msgspec import msgspec
@@ -454,41 +455,88 @@ def _make_generate_obj(rid, is_single):
return obj return obj
class TestDiscardPendingReqStates(CustomTestCase): class TestReleaseReqStatesOnFailure(CustomTestCase):
"""Direct tests for _discard_pending_req_states.""" """Direct tests for _release_req_states_on_failure."""
def test_discard_single(self): def test_undelivered_single_is_dropped(self):
tm = _make_tokenizer_manager(self) tm = _make_tokenizer_manager(self)
rid = "d_single" rid = "d_single"
tm.rid_to_state[rid] = _make_req_state(rid) tm.rid_to_state[rid] = _make_req_state(rid)
obj = Mock(spec=GenerateReqInput) tm._release_req_states_on_failure([rid])
obj.is_single = True
obj.rid = rid
tm._discard_pending_req_states(obj)
self.assertNotIn(rid, tm.rid_to_state) self.assertNotIn(rid, tm.rid_to_state)
def test_discard_batch_removes_all(self): def test_undelivered_batch_removes_all(self):
tm = _make_tokenizer_manager(self) tm = _make_tokenizer_manager(self)
rids = ["d0", "d1", "d2"] rids = ["d0", "d1", "d2"]
for r in rids: for r in rids:
tm.rid_to_state[r] = _make_req_state(r) tm.rid_to_state[r] = _make_req_state(r)
obj = Mock(spec=GenerateReqInput) tm._release_req_states_on_failure(rids)
obj.is_single = False
obj.rid = list(rids)
tm._discard_pending_req_states(obj)
for r in rids: for r in rids:
self.assertNotIn(r, tm.rid_to_state) self.assertNotIn(r, tm.rid_to_state)
def test_discard_ignores_already_removed(self): def test_ignores_already_removed(self):
"""Popping a rid that is no longer present must not raise.""" """A rid that is no longer present must not raise."""
tm = _make_tokenizer_manager(self) tm = _make_tokenizer_manager(self)
tm.rid_to_state["p1"] = _make_req_state("p1") tm.rid_to_state["p1"] = _make_req_state("p1")
obj = Mock(spec=GenerateReqInput) tm._release_req_states_on_failure(["p1", "already_gone"])
obj.is_single = False
obj.rid = ["p1", "already_gone"]
tm._discard_pending_req_states(obj) # must not raise
self.assertNotIn("p1", tm.rid_to_state) self.assertNotIn("p1", tm.rid_to_state)
def test_dispatched_single_is_aborted_and_state_kept(self):
tm = _make_tokenizer_manager(self)
tm.server_args.tokenizer_worker_num = 1
tm._dispatch_to_scheduler = Mock()
tm.enable_metrics = True
tm.metrics_collector = MagicMock()
rid = "d_live"
state = _make_req_state(rid)
state.dispatched = True
tm.rid_to_state[rid] = state
tm._release_req_states_on_failure([rid])
tm._release_req_states_on_failure([rid])
sent = [c.args[0] for c in tm._dispatch_to_scheduler.call_args_list]
self.assertEqual(
[type(m) for m in sent], [AbortReq], "expected exactly one AbortReq"
)
self.assertEqual(sent[0].rid, rid)
self.assertIn(rid, tm.rid_to_state)
self.assertTrue(state.abort_sent)
tm.metrics_collector.observe_one_aborted_request.assert_called_once()
def test_dispatched_batch_aborts_delivered_and_drops_rest(self):
tm = _make_tokenizer_manager(self)
tm.server_args.tokenizer_worker_num = 1
tm._dispatch_to_scheduler = Mock()
delivered, undelivered = "d_delivered", "d_undelivered"
live = _make_req_state(delivered)
live.dispatched = True
tm.rid_to_state[delivered] = live
tm.rid_to_state[undelivered] = _make_req_state(undelivered)
tm._release_req_states_on_failure([delivered, undelivered])
sent = [c.args[0] for c in tm._dispatch_to_scheduler.call_args_list]
self.assertEqual([type(m) for m in sent], [AbortReq])
self.assertEqual(sent[0].rid, delivered)
self.assertIn(delivered, tm.rid_to_state)
self.assertNotIn(undelivered, tm.rid_to_state)
def test_abort_failure_does_not_stop_cleanup(self):
tm = _make_tokenizer_manager(self)
tm.server_args.tokenizer_worker_num = 1
tm._dispatch_to_scheduler = Mock(side_effect=RuntimeError("send failed"))
delivered, undelivered = "live", "pending"
live = _make_req_state(delivered)
live.dispatched = True
tm.rid_to_state[delivered] = live
tm.rid_to_state[undelivered] = _make_req_state(undelivered)
with self.assertLogs(level="ERROR"):
tm._release_req_states_on_failure([delivered, undelivered])
self.assertIn(delivered, tm.rid_to_state)
self.assertFalse(live.abort_sent)
self.assertNotIn(undelivered, tm.rid_to_state)
class TestParallelStreamTaskCleanup(CustomTestCase): class TestParallelStreamTaskCleanup(CustomTestCase):
def test_failing_choice_cancels_and_closes_sibling_waiters(self): def test_failing_choice_cancels_and_closes_sibling_waiters(self):
@@ -595,6 +643,27 @@ class TestGenerateRequestCleanupOnDispatchFailure(CustomTestCase):
for r in rids: for r in rids:
self.assertNotIn(r, tm.rid_to_state) self.assertNotIn(r, tm.rid_to_state)
def test_parallel_sampling_failure_cleans_generated_rid(self):
tm = _make_tm_for_generate(self)
obj = GenerateReqInput(
text=["hello"],
rid=["base"],
sampling_params={"n": 2},
)
tokenized = MagicMock()
tokenized.mm_inputs = None
tokenized.sampling_params = MagicMock()
tm._tokenize_one_request = AsyncMock(return_value=tokenized)
tm._send_one_request = Mock(side_effect=RuntimeError("dispatch failed"))
async def drive():
await tm.generate_request(obj).__anext__()
with self.assertRaisesRegex(RuntimeError, "dispatch failed"):
asyncio.run(drive())
self.assertFalse(tm.rid_to_state)
def test_thinking_budget_rejects_runtime_without_strict_thinking(self): def test_thinking_budget_rejects_runtime_without_strict_thinking(self):
tm = _make_tm_for_generate(self) tm = _make_tm_for_generate(self)
obj = GenerateReqInput( obj = GenerateReqInput(
@@ -641,5 +710,52 @@ class TestWaitOneResponseAfterStateFreed(CustomTestCase):
self.assertEqual(out["text"], "hello") self.assertEqual(out["text"], "hello")
class TestDisconnectAfterDispatchAbortsRequest(CustomTestCase):
"""Cancellation after dispatch must stop the scheduler request."""
@patch(
"sglang.srt.managers.tokenizer_manager.wrap_shm_features",
side_effect=lambda obj: obj,
)
def test_cancel_after_dispatch_sends_abort_and_keeps_state(self, _wrap_shm):
tm = _make_tm_for_generate(self)
tm.cuda_vmm_feature_transport = Mock()
tm.cuda_vmm_feature_transport.prepare_for_dispatch_async = AsyncMock(
return_value=[]
)
tm._dispatch_to_scheduler = Mock()
rid = "disconnect_zombie"
obj = _make_generate_obj(rid, is_single=True)
obj.return_prompt_token_ids = False
tokenized = MagicMock()
tokenized.rid = rid
tokenized.mm_inputs = None
tm._tokenize_one_request = AsyncMock(return_value=tokenized)
async def drive():
task = asyncio.create_task(tm.generate_request(obj).__anext__())
for _ in range(100):
await asyncio.sleep(0)
if tm._dispatch_to_scheduler.called:
break
self.assertTrue(
tm._dispatch_to_scheduler.called, "request never dispatched"
)
state = tm.rid_to_state.get(rid)
self.assertIsNotNone(state)
self.assertTrue(state.dispatched)
task.cancel()
with self.assertRaises(asyncio.CancelledError):
await task
asyncio.run(drive())
sent = [c.args[0] for c in tm._dispatch_to_scheduler.call_args_list]
aborts = [m for m in sent if isinstance(m, AbortReq) and m.rid == rid]
self.assertTrue(aborts, "disconnect must send an AbortReq to the scheduler")
self.assertIn(rid, tm.rid_to_state)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main(verbosity=2) unittest.main(verbosity=2)
@@ -376,6 +376,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.srt.managers.tokenizer_manager import TokenizerManager
manager = object.__new__(TokenizerManager) manager = object.__new__(TokenizerManager)
manager.rid_to_state = {}
transport = MagicMock() transport = MagicMock()
transport.prepare_for_dispatch_async = AsyncMock(return_value=[]) transport.prepare_for_dispatch_async = AsyncMock(return_value=[])
manager.cuda_vmm_feature_transport = transport manager.cuda_vmm_feature_transport = transport
@@ -403,6 +404,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
) )
manager = object.__new__(tokenizer_manager.TokenizerManager) manager = object.__new__(tokenizer_manager.TokenizerManager)
manager.rid_to_state = {}
transport = MagicMock() transport = MagicMock()
manager._dispatch_to_scheduler = MagicMock( manager._dispatch_to_scheduler = MagicMock(
side_effect=RuntimeError("send failed") side_effect=RuntimeError("send failed")
@@ -437,6 +439,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
) )
manager = object.__new__(tokenizer_manager.TokenizerManager) manager = object.__new__(tokenizer_manager.TokenizerManager)
manager.rid_to_state = {}
transport = MagicMock() transport = MagicMock()
manager._dispatch_to_scheduler = MagicMock() manager._dispatch_to_scheduler = MagicMock()
time_stats = MagicMock() time_stats = MagicMock()
@@ -78,9 +78,25 @@ def _returned_field_names(function):
returned literal, assignments (annotated or not) to a returned name, a returned literal, assignments (annotated or not) to a returned name, a
literal-key subscript write on it, and `.update(field=...)` on it. A literal-key subscript write on it, and `.update(field=...)` on it. A
spelling this cannot see raises instead of skipping. spelling this cannot see raises instead of skipping.
``overrides[name]`` is also accepted when ``name`` comes from
``for name in ("a", "b", ...)`` -- the keys stay statically enumerable.
""" """
names = set() names = set()
returned = set() returned = set()
# for x in ("a", "b"): ... -> {"x": {"a", "b"}}
loop_keys = {
node.target.id: {elt.value for elt in node.iter.elts}
for node in ast.walk(function)
if isinstance(node, ast.For)
and isinstance(node.target, ast.Name)
and isinstance(node.iter, (ast.Tuple, ast.List))
and node.iter.elts
and all(
isinstance(elt, ast.Constant) and isinstance(elt.value, str)
for elt in node.iter.elts
)
}
def top_level_keys(mapping): def top_level_keys(mapping):
for key in mapping.keys: for key in mapping.keys:
@@ -88,6 +104,14 @@ def _returned_field_names(function):
raise AssertionError(f"non-literal key in {function.name}") raise AssertionError(f"non-literal key in {function.name}")
names.add(key.value) names.add(key.value)
def add_subscript_key(key):
if isinstance(key, ast.Constant):
names.add(key.value)
elif isinstance(key, ast.Name) and key.id in loop_keys:
names.update(loop_keys[key.id])
else:
raise AssertionError(f"non-literal key in {function.name}")
for node in ast.walk(function): for node in ast.walk(function):
if isinstance(node, ast.Return) and node.value is not None: if isinstance(node, ast.Return) and node.value is not None:
value = node.value value = node.value
@@ -114,9 +138,7 @@ def _returned_field_names(function):
if isinstance(target, ast.Subscript) and ( if isinstance(target, ast.Subscript) and (
isinstance(target.value, ast.Name) and target.value.id in returned isinstance(target.value, ast.Name) and target.value.id in returned
): ):
if not isinstance(target.slice, ast.Constant): add_subscript_key(target.slice)
raise AssertionError(f"non-literal key in {function.name}")
names.add(target.slice.value)
if ( if (
isinstance(node, ast.Call) isinstance(node, ast.Call)
and isinstance(node.func, ast.Attribute) and isinstance(node.func, ast.Attribute)