diff --git a/examples/server.py b/examples/server.py index 96c68a20d..84864a803 100644 --- a/examples/server.py +++ b/examples/server.py @@ -71,6 +71,12 @@ MAX_TOTAL_OPTIONS = getattr(_laya_serve, "MAX_TOTAL_OPTIONS", 512) +# The sentence `laya.serve._state_length` answers a state it cannot measure with. Named because +# `_limit_violation` hands this answer back as data for `/gui` rather than reading it off the +# exception (`py/stack-trace-exposure`): one literal, two readers. +_STATE_NOT_SERIALIZABLE = "'state' must be JSON-serializable" + + def _fallback_state_length(state: Any) -> int: """`laya.serve._state_length` for a laya that predates it, 400 included. @@ -83,7 +89,7 @@ def _fallback_state_length(state: Any) -> int: try: return len(json.dumps(state, ensure_ascii=False)) except (TypeError, ValueError, RecursionError): - raise HTTPException(status_code=400, detail="'state' must be JSON-serializable") + raise HTTPException(status_code=400, detail=_STATE_NOT_SERIALIZABLE) # Borrowed, not restated: the state length has to be measured on the text the tokenizer sees, @@ -333,8 +339,10 @@ def presets() -> Dict[str, Any]: return PRESETS -def _check_request_limits(state: Any, questions: Dict[str, Any]) -> None: - """Refuse an oversized request, as `laya.serve._check_request_limits` does. +def _limit_violation(state: Any, questions: Dict[str, Any]) -> Optional[tuple]: + """The first limit this request breaks, as `(status, message)`, or None when it fits. + + Refuses an oversized request, as `laya.serve._check_request_limits` does. Laya encodes the state once per question, so cost is questions x state size, collated into one tensor, and a choice or score question adds one sequence per @@ -351,18 +359,23 @@ def _check_request_limits(state: Any, questions: Dict[str, Any]) -> None: laya.serve answers 413, and FastAPI's validation-error response includes the offending `input`, so rejecting a 5 MB state would echo all 5 MB back to the caller -- turning a size limit into an amplifier. + + Returned rather than raised, and the message built here from the request and the constants: + `/gui` puts this line on a page, and taking it back off the exception is what code scanning + flags as stack-trace exposure (`py/stack-trace-exposure`). One place builds the words, so the + API and the page cannot disagree about them -- `_check_request_limits` raises the same pair + for the callers that answer in JSON. """ if len(questions) > MAX_QUESTIONS: - raise HTTPException( - status_code=413, - detail="too many questions (%d > %d)" % (len(questions), MAX_QUESTIONS), - ) - size = _state_length(state) + return 413, "too many questions (%d > %d)" % (len(questions), MAX_QUESTIONS) + try: + size = _state_length(state) + except HTTPException: + # `_state_length` answers a state it cannot measure with the fixed 400 named above. As + # data, for the same reason: the page shows this sentence, it does not read it. + return 400, _STATE_NOT_SERIALIZABLE if size > MAX_STATE_CHARS: - raise HTTPException( - status_code=413, - detail="state too large (%d > %d chars)" % (size, MAX_STATE_CHARS), - ) + return 413, "state too large (%d > %d chars)" % (size, MAX_STATE_CHARS) # Counted exactly as laya.serve counts them, and refused for the same reason. The # increment belongs inside the two branches, as it does there: a `noul` question carries # false/true criteria, which are option *texts* rather than answer options, so a total @@ -382,24 +395,27 @@ def _check_request_limits(state: Any, questions: Dict[str, Any]) -> None: count = len(crit) total_options += count if count > MAX_CHOICE_OPTIONS: - raise HTTPException( - status_code=413, - detail="too many choice options for %r (%d > %d)" % (qid, count, MAX_CHOICE_OPTIONS), - ) + return 413, "too many choice options for %r (%d > %d)" % (qid, count, MAX_CHOICE_OPTIONS) elif qtype == "score" and isinstance(crit, list): count = len(crit) total_options += count if count > MAX_SCORE_LEVELS: - raise HTTPException( - status_code=413, - detail="too many score levels for %r (%d > %d)" % (qid, count, MAX_SCORE_LEVELS), - ) + return 413, "too many score levels for %r (%d > %d)" % (qid, count, MAX_SCORE_LEVELS) if total_options > MAX_TOTAL_OPTIONS: - raise HTTPException( - status_code=413, - detail="too many answer options across questions (%d > %d)" - % (total_options, MAX_TOTAL_OPTIONS), - ) + return 413, ("too many answer options across questions (%d > %d)" + % (total_options, MAX_TOTAL_OPTIONS)) + return None + + +def _check_request_limits(state: Any, questions: Dict[str, Any]) -> None: + """Refuse an oversized request: the 413 (or 400) `laya.serve` answers with. + + The JSON surfaces raise and let FastAPI write the response; the sentence is the one + `_limit_violation` built, so `/predict` and the `/gui` page answer with the same words. + """ + violation = _limit_violation(state, questions) + if violation: + raise HTTPException(status_code=violation[0], detail=violation[1]) @app.post("/predict") @@ -3487,7 +3503,12 @@ def _error_lines(exc: Exception) -> List[str]: return lines if isinstance(exc, json.JSONDecodeError): return [f"Invalid JSON: {exc.msg} at line {exc.lineno}, column {exc.colno}"] - return [str(exc)] + # An unclassified failure: the caller gets a fixed line, the cause goes to the log. Its text + # names paths and libraries -- it describes the deployment, not the request -- and reading it + # back out is what code scanning flags (py/stack-trace-exposure). The branches above carry the + # caller's own field names and a JSON position, addressed to whoever fixes the request. + _log.warning("unreadable request", exc_info=exc) + return ["The request could not be read."] def _gui_error(title: str, lines: List[str], hint: str) -> HTMLResponse: @@ -3519,16 +3540,18 @@ async def gui_predict(request: Request) -> HTMLResponse: except (json.JSONDecodeError, PydanticValidationError, ValueError) as exc: return _gui_error("Bad request", _error_lines(exc), "Fix the request and send it again.") - try: - _check_request_limits(req.state, req.questions) - except HTTPException as exc: - # Branch on the status: this check answers 400 as well as 413 now, and a size hint under - # the heading "Request too large" would describe the wrong problem. - if exc.status_code == 413: - return _gui_error("Request too large", [str(exc.detail)], + violation = _limit_violation(req.state, req.questions) + if violation: + # Branch on the status: this check answers 400 as well as 413, and a size hint under the + # heading "Request too large" would describe the wrong problem. The line is the one the + # check built from the request -- not text read back off the exception it would have + # raised, which is what code scanning flags (py/stack-trace-exposure). + status, message = violation + if status == 413: + return _gui_error("Request too large", [message], f"The server takes up to {MAX_QUESTIONS} questions and a state of up to " f"{MAX_STATE_CHARS:,} characters. Split the request or shorten the state.") - return _gui_error("Bad request", [str(exc.detail)], "Fix the request and send it again.") + return _gui_error("Bad request", [message], "Fix the request and send it again.") questions = _questions(req.questions) try: started = time.perf_counter() diff --git a/tests/test_example_server_errors.py b/tests/test_example_server_errors.py index 6dfc99f7f..f9672a9e3 100644 --- a/tests/test_example_server_errors.py +++ b/tests/test_example_server_errors.py @@ -90,7 +90,9 @@ def main(): handler = Collect() logger = logging.getLogger("laya.example-server") logger.addHandler(handler) - logger.setLevel(logging.ERROR) + # WARNING, not ERROR: the parse-path fallback answers the caller and logs at warning level -- + # a client's bad request is not a server error, but its cause still belongs in the log. + logger.setLevel(logging.WARNING) one = {"a": {"type": "noul", "instructions": "x"}} @@ -129,6 +131,24 @@ def main(): ok("a validation error is still 422 with its reason", r.status_code == 422 and "bogus" in r.text, "%s %s" % (r.status_code, r.text[:200])) + # --- a limit the /gui page reports is built from the request ------------- + # The page used to read its line back off the raised HTTPException. It names the limit and + # the number that broke it either way -- what changed is where the words come from. + one_q = {"a": {"type": "noul", "instructions": "x"}} + r = client.post("/gui", json={"state": "x" * (demo.MAX_STATE_CHARS + 1), "questions": one_q}) + ok("an oversized /gui names the limit", + r.status_code == 200 and "state too large" in r.text, r.text[:300]) + ok("the limit page carries no exception text", + not any(s in r.text for s in LEAKS), r.text[:300]) + + # --- an unclassified failure on the parse path keeps its cause in the log - + before = len(handler.records) + lines = demo._error_lines(ValueError(SECRET)) + ok("an unreadable request answers a fixed line", + lines == ["The request could not be read."], repr(lines)) + logged = [rec.exc_info[1] for rec in handler.records[before:] if rec.exc_info] + ok("and its cause still reaches the log", [str(c) for c in logged] == [SECRET], str(logged)[:200]) + print("\n%d passed, %d failed" % (len(PASS), len(FAIL))) for f in FAIL: print(" FAIL " + f)