fix(backend): API 文本输入加 max_length + 请求体大小中间件(CR-H9)
This commit is contained in:
109
apps/api/tests/test_input_size_limits.py
Normal file
109
apps/api/tests/test_input_size_limits.py
Normal file
@@ -0,0 +1,109 @@
|
||||
"""CR-H9:用户请求文本字段 max_length 上界 + 应用级请求体大小中间件(防 DoS/成本)。
|
||||
|
||||
- 单元:各请求 schema 的字符串字段在 cap 处可构造、cap+1 触发 ValidationError(含样本列表长度
|
||||
与列表项两处上界——最大 DoS 面)。
|
||||
- 集成(无 pg):超大请求体经中间件直接 413(不入路由/DB)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from cryptography.fernet import Fernet
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
# create_app() 于 import ww_api.main 时构建,需 CREDENTIAL_ENC_KEY——先于任何 main 导入置好。
|
||||
os.environ.setdefault("CREDENTIAL_ENC_KEY", Fernet.generate_key().decode())
|
||||
|
||||
from ww_api.schemas.foreshadow import ForeshadowRegisterRequest
|
||||
from ww_api.schemas.generation import CharacterGenerateRequest, WorldGenerateRequest
|
||||
from ww_api.schemas.projects import (
|
||||
AcceptRequest,
|
||||
DraftSaveRequest,
|
||||
ProjectCreateRequest,
|
||||
ReviewRequest,
|
||||
)
|
||||
from ww_api.schemas.providers import (
|
||||
ProviderCredentialInput,
|
||||
TierRoutingInput,
|
||||
)
|
||||
from ww_api.schemas.providers import (
|
||||
TestConnectionRequest as _TestConnectionRequest, # 别名避免 pytest 误采集 Test* 类
|
||||
)
|
||||
from ww_api.schemas.rules import RuleCreateRequest
|
||||
from ww_api.schemas.style import RefineRequest, StyleLearnRequest
|
||||
from ww_api.schemas.templates import TemplateCreateRequest
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _StrCase:
|
||||
"""一个字符串字段的上界用例:cap 处可构造、cap+1 触发 ValidationError。"""
|
||||
|
||||
schema: type[BaseModel]
|
||||
field_name: str
|
||||
cap: int
|
||||
case_id: str
|
||||
base: dict[str, Any] = field(default_factory=dict) # 该 schema 其余必填字段
|
||||
|
||||
|
||||
_STR_CASES: list[_StrCase] = [
|
||||
_StrCase(ProviderCredentialInput, "api_key", 512, "api_key<=512", {"provider": "p"}),
|
||||
_StrCase(ProviderCredentialInput, "provider", 64, "provider<=64", {"api_key": "k"}),
|
||||
_StrCase(TierRoutingInput, "model", 128, "model<=128", {"tier": "writer", "provider": "p"}),
|
||||
_StrCase(
|
||||
TierRoutingInput, "provider", 64, "routing-provider<=64", {"tier": "writer", "model": "m"}
|
||||
),
|
||||
_StrCase(_TestConnectionRequest, "provider", 64, "test-provider<=64"),
|
||||
_StrCase(WorldGenerateRequest, "brief", 10_000, "world-brief<=10000"),
|
||||
_StrCase(CharacterGenerateRequest, "brief", 10_000, "char-brief<=10000"),
|
||||
_StrCase(ProjectCreateRequest, "title", 200, "title<=200"),
|
||||
_StrCase(AcceptRequest, "final_text", 200_000, "final_text<=200000"),
|
||||
_StrCase(DraftSaveRequest, "text", 200_000, "draft-text<=200000"),
|
||||
_StrCase(ReviewRequest, "draft", 200_000, "review-draft<=200000"),
|
||||
_StrCase(RefineRequest, "segment", 20_000, "segment<=20000"),
|
||||
_StrCase(ForeshadowRegisterRequest, "code", 100, "code<=100", {"title": "t"}),
|
||||
_StrCase(ForeshadowRegisterRequest, "title", 500, "foreshadow-title<=500", {"code": "c"}),
|
||||
_StrCase(RuleCreateRequest, "content", 10_000, "rule-content<=10000", {"level": "project"}),
|
||||
_StrCase(TemplateCreateRequest, "title", 200, "template-title<=200", {"body": "b"}),
|
||||
_StrCase(TemplateCreateRequest, "body", 20_000, "template-body<=20000", {"title": "t"}),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("case", _STR_CASES, ids=[c.case_id for c in _STR_CASES])
|
||||
def test_request_str_fields_reject_over_max_length(case: _StrCase) -> None:
|
||||
# Arrange / Act — 长度 == cap:构造成功。
|
||||
case.schema(**{**case.base, case.field_name: "x" * case.cap})
|
||||
# Assert — 长度 == cap+1:超界即 ValidationError。
|
||||
with pytest.raises(ValidationError):
|
||||
case.schema(**{**case.base, case.field_name: "x" * (case.cap + 1)})
|
||||
|
||||
|
||||
def test_style_learn_samples_item_and_list_caps() -> None:
|
||||
# 单条样本 == 20_000 字符 OK;20_001 → ValidationError(最大 DoS 面:每条正文都设上界)。
|
||||
StyleLearnRequest(samples=["x" * 20_000])
|
||||
with pytest.raises(ValidationError):
|
||||
StyleLearnRequest(samples=["x" * 20_001])
|
||||
# 列表 == 20 条 OK;21 条 → ValidationError(同时约束列表长度)。
|
||||
StyleLearnRequest(samples=["s"] * 20)
|
||||
with pytest.raises(ValidationError):
|
||||
StyleLearnRequest(samples=["s"] * 21)
|
||||
|
||||
|
||||
def test_oversized_request_body_rejected_413() -> None:
|
||||
"""超大请求体经 body-size 中间件直接 413(不入路由/DB)。
|
||||
|
||||
不以 context manager 进入 TestClient → 不跑 lifespan/seed → 不需 Postgres。中间件在
|
||||
路由前拦截,故请求根本不触达任何 DB。
|
||||
"""
|
||||
from fastapi.testclient import TestClient
|
||||
from ww_api.main import create_app
|
||||
from ww_api.middleware import MAX_BODY_BYTES
|
||||
|
||||
client = TestClient(create_app())
|
||||
resp = client.post("/projects", content=b"x" * (MAX_BODY_BYTES + 1))
|
||||
|
||||
assert resp.status_code == 413
|
||||
assert resp.json()["error"]["code"] == "PAYLOAD_TOO_LARGE"
|
||||
@@ -14,7 +14,11 @@ from ww_db import get_sessionmaker
|
||||
from ww_shared import AppError, ErrorBody, ErrorCode, ErrorEnvelope
|
||||
|
||||
from ww_api.logging_config import configure_logging, get_logger
|
||||
from ww_api.middleware import REQUEST_ID_HEADER, request_id_middleware
|
||||
from ww_api.middleware import (
|
||||
REQUEST_ID_HEADER,
|
||||
body_size_limit_middleware,
|
||||
request_id_middleware,
|
||||
)
|
||||
from ww_api.routers import (
|
||||
chain,
|
||||
foreshadow,
|
||||
@@ -74,6 +78,8 @@ def create_app() -> FastAPI:
|
||||
allow_headers=["content-type", "authorization", REQUEST_ID_HEADER],
|
||||
)
|
||||
app.middleware("http")(request_id_middleware)
|
||||
# 请求体大小守卫(CR-H9):在路由前拦截超大 body,直接 413。
|
||||
app.middleware("http")(body_size_limit_middleware)
|
||||
|
||||
@app.exception_handler(AppError)
|
||||
async def _app_error_handler(request: Request, exc: AppError) -> JSONResponse:
|
||||
|
||||
@@ -12,10 +12,14 @@ from collections.abc import Awaitable, Callable
|
||||
|
||||
import structlog
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
from starlette.responses import JSONResponse, Response
|
||||
from ww_shared import ErrorBody, ErrorCode, ErrorEnvelope
|
||||
|
||||
REQUEST_ID_HEADER = "x-request-id"
|
||||
|
||||
# 请求体大小上限(2 MiB):超过即 413,防超大 body 撑爆内存/成本(CR-H9)。
|
||||
MAX_BODY_BYTES = 2 * 1024 * 1024
|
||||
|
||||
# 放行安全字符集(字母数字 . _ -,1–128 位):覆盖 uuid4().hex、常见 trace id 与客户端
|
||||
# 自定义短 id;含空白/控制字符/换行/注入序列/超长的一律丢弃改生成新 uuid——防止把未经
|
||||
# 校验的客户端值原样写进日志/响应头(日志注入防护)。
|
||||
@@ -38,3 +42,29 @@ async def request_id_middleware(
|
||||
response = await call_next(request)
|
||||
response.headers[REQUEST_ID_HEADER] = request_id
|
||||
return response
|
||||
|
||||
|
||||
async def body_size_limit_middleware(
|
||||
request: Request, call_next: Callable[[Request], Awaitable[Response]]
|
||||
) -> Response:
|
||||
"""按 `Content-Length` 拦截超大请求体,超限直接 413(CR-H9)。
|
||||
|
||||
用户中间件跑在 FastAPI 的 AppError 处理器**之外**,故必须自建错误信封响应(不能靠抛
|
||||
AppError)。分块编码(无 `Content-Length`)无法在此拦截——单用户原型可接受的限制。
|
||||
"""
|
||||
content_length = request.headers.get("content-length")
|
||||
if content_length is not None:
|
||||
try:
|
||||
size = int(content_length)
|
||||
except ValueError:
|
||||
size = 0 # 非法头防御性视作 0,交由下游正常解析/报错。
|
||||
if size > MAX_BODY_BYTES:
|
||||
envelope = ErrorEnvelope(
|
||||
error=ErrorBody(
|
||||
code=ErrorCode.PAYLOAD_TOO_LARGE,
|
||||
message="请求体过大",
|
||||
request_id=_sanitize_request_id(request.headers.get(REQUEST_ID_HEADER)),
|
||||
)
|
||||
)
|
||||
return JSONResponse(status_code=413, content=envelope.model_dump())
|
||||
return await call_next(request)
|
||||
|
||||
@@ -11,12 +11,18 @@ from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
# 请求字符串字段长度上界(CR-H9:防超长入参 DoS/成本)。代号=短标识;标题=一句话。
|
||||
_CODE_MAX = 100
|
||||
_TITLE_MAX = 500
|
||||
|
||||
|
||||
class ForeshadowRegisterRequest(BaseModel):
|
||||
"""POST /projects/:id/foreshadow:作者显式登记一条伏笔(status 落 OPEN)。"""
|
||||
|
||||
code: str = Field(min_length=1, description="伏笔代号,`(project_id, code)` 唯一")
|
||||
title: str = Field(min_length=1, description="伏笔标题/一句话描述")
|
||||
code: str = Field(
|
||||
min_length=1, max_length=_CODE_MAX, description="伏笔代号,`(project_id, code)` 唯一"
|
||||
)
|
||||
title: str = Field(min_length=1, max_length=_TITLE_MAX, description="伏笔标题/一句话描述")
|
||||
planted_at: int | None = Field(default=None, description="埋设章号")
|
||||
content: str | None = Field(default=None, description="伏笔正文/线索")
|
||||
expected_close_from: int | None = Field(default=None, description="预期回收窗口起始章")
|
||||
|
||||
@@ -19,6 +19,9 @@ from ww_api.schemas.rules import RuleView
|
||||
# 批量生成数量上限(防一次性巨量调用;原型保守值)。
|
||||
MAX_GENERATE_COUNT = 12
|
||||
|
||||
# 一句话需求文本上界(CR-H9:防超长入参撑爆上下文/成本)。
|
||||
_BRIEF_MAX = 10_000
|
||||
|
||||
|
||||
# ---- 世界观生成(预览)----
|
||||
|
||||
@@ -26,7 +29,9 @@ MAX_GENERATE_COUNT = 12
|
||||
class WorldGenerateRequest(BaseModel):
|
||||
"""POST /projects/:id/world/generate:据需求生成世界观实体(预览)。"""
|
||||
|
||||
brief: str = Field(min_length=1, description="世界观生成需求(题材/设定方向/约束)")
|
||||
brief: str = Field(
|
||||
min_length=1, max_length=_BRIEF_MAX, description="世界观生成需求(题材/设定方向/约束)"
|
||||
)
|
||||
|
||||
|
||||
class WorldEntityCardView(BaseModel):
|
||||
@@ -54,7 +59,7 @@ class CharacterGenerateRequest(BaseModel):
|
||||
单一 `role` + `count` 语义。
|
||||
"""
|
||||
|
||||
brief: str = Field(min_length=1, description="角色生成需求")
|
||||
brief: str = Field(min_length=1, max_length=_BRIEF_MAX, description="角色生成需求")
|
||||
count: int = Field(
|
||||
default=1,
|
||||
ge=1,
|
||||
|
||||
@@ -11,11 +11,15 @@ from typing import Any, Literal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
# 请求文本长度上界(CR-H9:防超长入参 DoS/成本)。标题=标签级;章正文=整章级。
|
||||
_TITLE_MAX = 200
|
||||
_CHAPTER_TEXT_MAX = 200_000
|
||||
|
||||
|
||||
class ProjectCreateRequest(BaseModel):
|
||||
"""POST /projects:立项向导字段(owner_id 由后端补 stub,不入参)。"""
|
||||
|
||||
title: str = Field(min_length=1)
|
||||
title: str = Field(min_length=1, max_length=_TITLE_MAX)
|
||||
genre: str | None = None
|
||||
logline: str | None = None
|
||||
premise: str | None = None
|
||||
@@ -120,9 +124,9 @@ class DraftStreamRequest(BaseModel):
|
||||
directive: str | None = None
|
||||
|
||||
|
||||
# 整章重写输入上界(守 DoS/成本,CR-H9 方向):意见短、整章草稿大。
|
||||
# 整章重写输入上界(守 DoS/成本,CR-H9 方向):意见短、整章草稿大(复用章正文上界,DRY)。
|
||||
_REWRITE_FEEDBACK_MAX = 4000
|
||||
_REWRITE_DRAFT_MAX = 200_000
|
||||
_REWRITE_DRAFT_MAX = _CHAPTER_TEXT_MAX
|
||||
|
||||
|
||||
class RewriteStreamRequest(BaseModel):
|
||||
@@ -140,7 +144,7 @@ class RewriteStreamRequest(BaseModel):
|
||||
class DraftSaveRequest(BaseModel):
|
||||
"""PUT /projects/:id/chapters/:no/draft:自动保存草稿正文。"""
|
||||
|
||||
text: str
|
||||
text: str = Field(max_length=_CHAPTER_TEXT_MAX)
|
||||
|
||||
|
||||
class DraftResponse(BaseModel):
|
||||
@@ -179,7 +183,7 @@ class ReviewRequest(BaseModel):
|
||||
不传 `draft` 时端点回退到已保存的草稿(chapter_repo)。
|
||||
"""
|
||||
|
||||
draft: str | None = None
|
||||
draft: str | None = Field(default=None, max_length=_CHAPTER_TEXT_MAX)
|
||||
|
||||
|
||||
class ReviewConflictView(BaseModel):
|
||||
@@ -237,7 +241,9 @@ class ConflictDecision(BaseModel):
|
||||
class AcceptRequest(BaseModel):
|
||||
"""POST /projects/:id/chapters/:no/accept:裁决清单 + 可能改过的终稿。"""
|
||||
|
||||
final_text: str = Field(min_length=1, description="作者裁决/改稿后的最终验收文本")
|
||||
final_text: str = Field(
|
||||
min_length=1, max_length=_CHAPTER_TEXT_MAX, description="作者裁决/改稿后的最终验收文本"
|
||||
)
|
||||
decisions: list[ConflictDecision] = Field(
|
||||
default_factory=list, description="对最近审稿每个冲突的裁决(每冲突必有其一,R5)"
|
||||
)
|
||||
|
||||
@@ -8,6 +8,11 @@ from __future__ import annotations
|
||||
from pydantic import BaseModel, Field
|
||||
from ww_llm_gateway.types import Tier
|
||||
|
||||
# 请求字符串字段长度上界(CR-H9:防超长入参 DoS/成本)。
|
||||
_PROVIDER_MAX = 64
|
||||
_MODEL_MAX = 128
|
||||
_API_KEY_MAX = 512
|
||||
|
||||
|
||||
class ProviderView(BaseModel):
|
||||
"""已配置提供商(掩码视图)。"""
|
||||
@@ -35,8 +40,8 @@ class ProvidersResponse(BaseModel):
|
||||
class ProviderCredentialInput(BaseModel):
|
||||
"""单条提供商凭据写入。"""
|
||||
|
||||
provider: str = Field(min_length=1)
|
||||
api_key: str = Field(min_length=1) # 明文入站,加密入库,绝不回显
|
||||
provider: str = Field(min_length=1, max_length=_PROVIDER_MAX)
|
||||
api_key: str = Field(min_length=1, max_length=_API_KEY_MAX) # 明文入站,加密入库,绝不回显
|
||||
|
||||
|
||||
class TierRoutingInput(BaseModel):
|
||||
@@ -44,8 +49,8 @@ class TierRoutingInput(BaseModel):
|
||||
|
||||
# tier 限定已知档位 writer/analyst/light;未知档位 → 422(QA MEDIUM:原接受任意字符串)。
|
||||
tier: Tier
|
||||
provider: str = Field(min_length=1)
|
||||
model: str = Field(min_length=1)
|
||||
provider: str = Field(min_length=1, max_length=_PROVIDER_MAX)
|
||||
model: str = Field(min_length=1, max_length=_MODEL_MAX)
|
||||
fallback: list[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
@@ -59,7 +64,7 @@ class ProvidersUpsertRequest(BaseModel):
|
||||
class TestConnectionRequest(BaseModel):
|
||||
"""POST /test:最小探测请求。"""
|
||||
|
||||
provider: str = Field(min_length=1)
|
||||
provider: str = Field(min_length=1, max_length=_PROVIDER_MAX)
|
||||
|
||||
|
||||
class CapabilitiesView(BaseModel):
|
||||
|
||||
@@ -14,8 +14,13 @@ from pydantic import BaseModel, Field, StringConstraints
|
||||
|
||||
RuleLevel = Literal["global", "genre", "style", "project"]
|
||||
|
||||
# 规则正文长度上界(CR-H9:防超长入参 DoS/成本)。
|
||||
_RULE_CONTENT_MAX = 10_000
|
||||
|
||||
# 先 strip 再校验长度:纯空白内容(" ")strip 后为空 → 422(QA MEDIUM:此前被接受入库)。
|
||||
RuleContent = Annotated[str, StringConstraints(strip_whitespace=True, min_length=1)]
|
||||
RuleContent = Annotated[
|
||||
str, StringConstraints(strip_whitespace=True, min_length=1, max_length=_RULE_CONTENT_MAX)
|
||||
]
|
||||
|
||||
|
||||
class RuleCreateRequest(BaseModel):
|
||||
|
||||
@@ -7,15 +7,26 @@ snake_case(命名契约,见 memory/gotchas)。前端经 OpenAPI→TS 客
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from typing import Literal
|
||||
from typing import Annotated, Literal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic import BaseModel, Field, StringConstraints
|
||||
|
||||
# 请求文本长度上界(CR-H9:防超长入参空跑 LLM/DoS)。选段/样本 20k 字符,样本列表 ≤20 条。
|
||||
_SEGMENT_MAX = 20_000
|
||||
_SAMPLE_MAX = 20_000
|
||||
_SAMPLES_LIST_MAX = 20
|
||||
# 预检澄清沿用相同选段上界(DRY);指令短。
|
||||
_MAX_CLARIFY_SEGMENT_LEN = _SEGMENT_MAX
|
||||
_MAX_CLARIFY_INSTRUCTION_LEN = 2000
|
||||
|
||||
|
||||
class StyleLearnRequest(BaseModel):
|
||||
"""学文风:上传样本正文(前端可读文件转文本)+ 模式(首学 / 更新)。"""
|
||||
|
||||
samples: list[str] = Field(min_length=1)
|
||||
# 同时约束**列表条数**与**每条样本长度**(最大 DoS 面:多条超长正文)。
|
||||
samples: list[Annotated[str, StringConstraints(max_length=_SAMPLE_MAX)]] = Field(
|
||||
min_length=1, max_length=_SAMPLES_LIST_MAX
|
||||
)
|
||||
mode: Literal["create", "update"] = "create"
|
||||
|
||||
|
||||
@@ -47,7 +58,7 @@ class StyleFingerprintResponse(BaseModel):
|
||||
class RefineRequest(BaseModel):
|
||||
"""回炉:重写选中段(可选改写指令)。"""
|
||||
|
||||
segment: str = Field(min_length=1)
|
||||
segment: str = Field(min_length=1, max_length=_SEGMENT_MAX)
|
||||
instruction: str | None = None
|
||||
|
||||
|
||||
@@ -58,11 +69,6 @@ class RefineResponse(BaseModel):
|
||||
refined: str
|
||||
|
||||
|
||||
# 上界:选段/意见文本长度守卫(防超长入参空跑 LLM;预检只看本段,无需整章)。
|
||||
_MAX_CLARIFY_SEGMENT_LEN = 20000
|
||||
_MAX_CLARIFY_INSTRUCTION_LEN = 2000
|
||||
|
||||
|
||||
class RefineClarifyRequest(BaseModel):
|
||||
"""润色预检澄清:选段 + 再沟通意见(WFW-9 M1 路线A 两阶段之「问题」阶段)。
|
||||
|
||||
|
||||
@@ -12,15 +12,24 @@ from typing import Annotated
|
||||
|
||||
from pydantic import BaseModel, Field, StringConstraints
|
||||
|
||||
# 长度上界(CR-H9:防超长入参 DoS/成本)。标题=标签级;正文=段落级。
|
||||
_TITLE_MAX = 200
|
||||
_BODY_MAX = 20_000
|
||||
|
||||
# 先去首尾空白再校验长度:纯空白(" ")strip 后为空 → min_length=1 不满足 → 422。
|
||||
NonBlankStr = Annotated[str, StringConstraints(strip_whitespace=True, min_length=1)]
|
||||
NonBlankTitle = Annotated[
|
||||
str, StringConstraints(strip_whitespace=True, min_length=1, max_length=_TITLE_MAX)
|
||||
]
|
||||
NonBlankBody = Annotated[
|
||||
str, StringConstraints(strip_whitespace=True, min_length=1, max_length=_BODY_MAX)
|
||||
]
|
||||
|
||||
|
||||
class TemplateCreateRequest(BaseModel):
|
||||
"""POST /templates:新建一条提示词模板。"""
|
||||
|
||||
title: NonBlankStr = Field(description="模板标题(非空,纯空白 → 422)")
|
||||
body: NonBlankStr = Field(description="模板正文(非空,一键填入生成器的 brief/text)")
|
||||
title: NonBlankTitle = Field(description="模板标题(非空,纯空白 → 422)")
|
||||
body: NonBlankBody = Field(description="模板正文(非空,一键填入生成器的 brief/text)")
|
||||
category: str | None = Field(default=None, description="可选分类")
|
||||
tool_key: str | None = Field(default=None, description="可选关联生成器(NULL=通用)")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user