Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
91 changes: 57 additions & 34 deletions examples/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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,
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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")
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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()
Expand Down
22 changes: 21 additions & 1 deletion tests/test_example_server_errors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"}}

Expand Down Expand Up @@ -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)
Expand Down
Loading