From 5478e3138d5eaaf251b3395cde0061ec9b928c84 Mon Sep 17 00:00:00 2001 From: pearseona Date: Fri, 28 Aug 2026 13:36:06 +0900 Subject: [PATCH 1/4] feat: limit request payload sizes to Dos and LLM cost blowup --- app/analysis/schemas.py | 17 ++++++++++++++++- app/chat/schemas.py | 16 ++++++++++++++-- app/core/config.py | 17 +++++++++++++++++ app/infrastructure/rabbitmq/schemas.py | 6 +++++- 4 files changed, 52 insertions(+), 4 deletions(-) diff --git a/app/analysis/schemas.py b/app/analysis/schemas.py index fdb0e30..bf1e963 100644 --- a/app/analysis/schemas.py +++ b/app/analysis/schemas.py @@ -1,12 +1,27 @@ from enum import Enum -from pydantic import BaseModel, ConfigDict, Field +from pydantic import BaseModel, ConfigDict, Field, field_validator + +from app.core.config import settings # Spring Boot Gateway에서 Python FastAPI로 검사를 요청할 때의 바디 규격 class SmishingAnalysisRequest(BaseModel): + # hide_input_in_errors: 검증 실패 시 원문(PII)이 에러에 담기지 않도록 한다. + model_config = ConfigDict(hide_input_in_errors=True) + text: str = Field(..., description="검사할 문자 메시지 본문 텍스트") + @field_validator("text") + @classmethod + def validate_text(cls, value: str) -> str: + if not value.strip(): + raise ValueError("text must not be blank") + # 분석 content와 동일한 상한(BE @Size(max=5000)와 정합)을 공유한다. + if len(value) > settings.MAX_ANALYSIS_CONTENT_LENGTH: + raise ValueError("text exceeds max length") + return value + # Python FastAPI가 Spring Boot로 최종 전달할 하이브리드 검사 결과 규칙 class UrlAnalysisResponse(BaseModel): diff --git a/app/chat/schemas.py b/app/chat/schemas.py index fac805f..cbdf5ba 100644 --- a/app/chat/schemas.py +++ b/app/chat/schemas.py @@ -3,6 +3,7 @@ from pydantic import BaseModel, ConfigDict, Field, field_validator from app.analysis.schemas import RiskGrade +from app.core.config import settings class ChatRole(str, Enum): @@ -41,7 +42,8 @@ def explanation_must_not_be_blank(cls, value: str) -> str: class ChatMessage(BaseModel): - model_config = ConfigDict(extra="forbid") + # hide_input_in_errors: 검증 실패 시 챗 원문(PII)이 에러에 담기지 않도록 한다. + model_config = ConfigDict(extra="forbid", hide_input_in_errors=True) role: ChatRole content: str @@ -51,15 +53,25 @@ class ChatMessage(BaseModel): def content_must_not_be_blank(cls, value: str) -> str: if not value.strip(): raise ValueError("content must not be blank") + # BE ChatMessage @Size(max=2000)와 정합. + if len(value) > settings.MAX_CHAT_CONTENT_LENGTH: + raise ValueError("content exceeds max length") return value class ChatRequest(BaseModel): - model_config = ConfigDict(extra="forbid") + model_config = ConfigDict(extra="forbid", hide_input_in_errors=True) analysisContext: AnalysisContext | None = None messages: list[ChatMessage] = Field(..., min_length=1) + @field_validator("messages") + @classmethod + def messages_within_limit(cls, value: list[ChatMessage]) -> list[ChatMessage]: + if len(value) > settings.MAX_CHAT_MESSAGES: + raise ValueError("messages exceeds max count") + return value + class ChatResponse(BaseModel): model_config = ConfigDict(extra="forbid") diff --git a/app/core/config.py b/app/core/config.py index b460577..e238e8f 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -69,6 +69,23 @@ class Settings(BaseSettings): le=1.0, ) + # 입력 크기 제한 (DoS·LLM 비용 방지, issue #120) + # SafeFam_BE가 계약상 게이트키퍼이므로 BE의 @Size 검증값과 정합을 맞춘다. + # - 분석 content: BE AnalysisRequest @Size(max = 5000) + # - 챗 content: BE ChatMessage @Size(max = 2000) + MAX_ANALYSIS_CONTENT_LENGTH: int = Field( + default=5_000, + ge=1, + ) + MAX_CHAT_CONTENT_LENGTH: int = Field( + default=2_000, + ge=1, + ) + MAX_CHAT_MESSAGES: int = Field( + default=40, + ge=1, + ) + # Stacking 모델 임계값 STACKING_NORMAL_PROBABILITY_MAX: float = Field( default=0.1, diff --git a/app/infrastructure/rabbitmq/schemas.py b/app/infrastructure/rabbitmq/schemas.py index 50bb9e8..09a45a5 100644 --- a/app/infrastructure/rabbitmq/schemas.py +++ b/app/infrastructure/rabbitmq/schemas.py @@ -13,6 +13,7 @@ from app.analysis.phishing_type import PhishingType from app.analysis.schemas import RiskGrade +from app.core.config import settings class AnalysisSource(str, Enum): @@ -25,7 +26,8 @@ class AnalysisSource(str, Enum): class AnalysisRequestedPayload(BaseModel): """AI 분석에 필요한 실제 문자 본문 데이터 스키마""" - model_config = ConfigDict(extra="forbid") + # hide_input_in_errors: 검증 실패 시 ValidationError에 원문(PII)이 담기지 않도록 한다. + model_config = ConfigDict(extra="forbid", hide_input_in_errors=True) sender: str | None = Field(default=None, max_length=100) content: str @@ -37,6 +39,8 @@ class AnalysisRequestedPayload(BaseModel): def validate_content(cls, value: str) -> str: if not value.strip(): raise ValueError("content must not be blank") + if len(value) > settings.MAX_ANALYSIS_CONTENT_LENGTH: + raise ValueError("content exceeds max length") return value From 27001eb96eb84bd5e3a5b6fb755c3ee2a3517305 Mon Sep 17 00:00:00 2001 From: pearseona Date: Fri, 28 Aug 2026 13:40:25 +0900 Subject: [PATCH 2/4] test: cover request payload size limimts and PII-safe validation errors --- tests/analysis/test_schemas.py | 43 ++++++++++++++++ tests/chat/test_schemas.py | 51 +++++++++++++++++++ tests/core/test_config.py | 27 ++++++++++ tests/infrastructure/rabbitmq/test_schemas.py | 37 ++++++++++++++ 4 files changed, 158 insertions(+) create mode 100644 tests/analysis/test_schemas.py diff --git a/tests/analysis/test_schemas.py b/tests/analysis/test_schemas.py new file mode 100644 index 0000000..ba0fd40 --- /dev/null +++ b/tests/analysis/test_schemas.py @@ -0,0 +1,43 @@ +import pytest +from pydantic import ValidationError + +from app.analysis.schemas import SmishingAnalysisRequest +from app.core.config import settings + + +def test_accepts_valid_text(): + request = SmishingAnalysisRequest(text="[국민은행] 계좌가 정지되었습니다.") + + assert request.text == "[국민은행] 계좌가 정지되었습니다." + + +@pytest.mark.parametrize("text", ["", " ", " ", "\t", "\n"]) +def test_rejects_blank_text(text: str): + with pytest.raises(ValidationError): + SmishingAnalysisRequest(text=text) + + +def test_accepts_text_at_max_length(): + request = SmishingAnalysisRequest( + text="가" * settings.MAX_ANALYSIS_CONTENT_LENGTH + ) + + assert len(request.text) == settings.MAX_ANALYSIS_CONTENT_LENGTH + + +def test_rejects_text_over_max_length(): + with pytest.raises(ValidationError): + SmishingAnalysisRequest( + text="가" * (settings.MAX_ANALYSIS_CONTENT_LENGTH + 1) + ) + + +def test_oversized_text_is_not_leaked_in_validation_error(): + """hide_input_in_errors=True 로 검증 실패 시 원문(PII)이 에러에 담기지 않아야 한다.""" + secret_marker = "01012345678-비밀번호-보이스피싱" + with pytest.raises(ValidationError) as exception_info: + SmishingAnalysisRequest( + text=secret_marker + "가" * settings.MAX_ANALYSIS_CONTENT_LENGTH + ) + + assert secret_marker not in str(exception_info.value) diff --git a/tests/chat/test_schemas.py b/tests/chat/test_schemas.py index 773ba92..67ccc77 100644 --- a/tests/chat/test_schemas.py +++ b/tests/chat/test_schemas.py @@ -2,6 +2,7 @@ from pydantic import ValidationError from app.chat.schemas import AnalysisContext, ChatMessage, ChatRequest +from app.core.config import settings def _analysis_context(**overrides) -> AnalysisContext: @@ -53,6 +54,56 @@ def test_chat_message_rejects_blank_content(): ChatMessage(role="user", content=" ") +def test_chat_message_accepts_content_at_max_length(): + message = ChatMessage( + role="user", content="가" * settings.MAX_CHAT_CONTENT_LENGTH + ) + + assert len(message.content) == settings.MAX_CHAT_CONTENT_LENGTH + + +def test_chat_message_rejects_content_over_max_length(): + with pytest.raises(ValidationError): + ChatMessage( + role="user", content="가" * (settings.MAX_CHAT_CONTENT_LENGTH + 1) + ) + + +def test_chat_message_oversized_content_is_not_leaked_in_error(): + """hide_input_in_errors=True 로 챗 원문(PII)이 에러에 담기지 않아야 한다.""" + secret_marker = "01012345678-비밀번호" + with pytest.raises(ValidationError) as exception_info: + ChatMessage( + role="user", + content=secret_marker + "가" * settings.MAX_CHAT_CONTENT_LENGTH, + ) + + assert secret_marker not in str(exception_info.value) + + +def test_chat_request_accepts_messages_at_max_count(): + request = ChatRequest( + analysisContext=_analysis_context(), + messages=[ + {"role": "user", "content": "질문"} + for _ in range(settings.MAX_CHAT_MESSAGES) + ], + ) + + assert len(request.messages) == settings.MAX_CHAT_MESSAGES + + +def test_chat_request_rejects_messages_over_max_count(): + with pytest.raises(ValidationError): + ChatRequest( + analysisContext=_analysis_context(), + messages=[ + {"role": "user", "content": "질문"} + for _ in range(settings.MAX_CHAT_MESSAGES + 1) + ], + ) + + def test_analysis_context_rejects_blank_explanation(): with pytest.raises(ValidationError): _analysis_context(explanation=" ") diff --git a/tests/core/test_config.py b/tests/core/test_config.py index 0ba3141..a4aa752 100644 --- a/tests/core/test_config.py +++ b/tests/core/test_config.py @@ -47,6 +47,33 @@ def test_external_api_timeout_defaults(): assert configured.URL_TRACE_TIMEOUT_SECONDS == 3.0 +def test_input_size_limit_defaults_match_backend_contract(): + """입력 크기 상한은 SafeFam_BE의 @Size 검증값과 정합을 맞춘다(issue #120). + + BE가 게이트키퍼이므로 값이 어긋나면 한쪽만 통과하는 불일치가 생긴다. + - 분석 content: BE AnalysisRequest @Size(max = 5000) + - 챗 content: BE ChatMessage @Size(max = 2000) + """ + configured = Settings(_env_file=None) + + assert configured.MAX_ANALYSIS_CONTENT_LENGTH == 5000 + assert configured.MAX_CHAT_CONTENT_LENGTH == 2000 + assert configured.MAX_CHAT_MESSAGES == 40 + + +@pytest.mark.parametrize( + "field_name", + [ + "MAX_ANALYSIS_CONTENT_LENGTH", + "MAX_CHAT_CONTENT_LENGTH", + "MAX_CHAT_MESSAGES", + ], +) +def test_input_size_limits_reject_non_positive(field_name: str): + with pytest.raises(ValidationError): + Settings(**{field_name: 0}, _env_file=None) + + def test_settings_accept_valid_stacking_probability_bounds(): configured = Settings( STACKING_NORMAL_PROBABILITY_MAX=0.2, diff --git a/tests/infrastructure/rabbitmq/test_schemas.py b/tests/infrastructure/rabbitmq/test_schemas.py index dd9c6e8..a1ef243 100644 --- a/tests/infrastructure/rabbitmq/test_schemas.py +++ b/tests/infrastructure/rabbitmq/test_schemas.py @@ -5,6 +5,7 @@ import pytest from pydantic import ValidationError +from app.core.config import settings from app.infrastructure.rabbitmq.schemas import ( AnalysisRequestedEvent, AnalysisSource, @@ -83,6 +84,42 @@ def test_rejects_blank_content( AnalysisRequestedEvent.model_validate(valid_event_data) +def test_accepts_content_at_max_length( + valid_event_data: dict, +): + valid_event_data["payload"]["content"] = "가" * settings.MAX_ANALYSIS_CONTENT_LENGTH + + event = AnalysisRequestedEvent.model_validate(valid_event_data) + + assert len(event.payload.content) == settings.MAX_ANALYSIS_CONTENT_LENGTH + + +def test_rejects_content_over_max_length( + valid_event_data: dict, +): + valid_event_data["payload"]["content"] = "가" * ( + settings.MAX_ANALYSIS_CONTENT_LENGTH + 1 + ) + + with pytest.raises(ValidationError): + AnalysisRequestedEvent.model_validate(valid_event_data) + + +def test_oversized_content_is_not_leaked_in_validation_error( + valid_event_data: dict, +): + """hide_input_in_errors=True 로 검증 실패 시 원문(PII)이 에러에 담기지 않아야 한다.""" + secret_marker = "01012345678-비밀번호-보이스피싱" + valid_event_data["payload"]["content"] = ( + secret_marker + "가" * settings.MAX_ANALYSIS_CONTENT_LENGTH + ) + + with pytest.raises(ValidationError) as exception_info: + AnalysisRequestedEvent.model_validate(valid_event_data) + + assert secret_marker not in str(exception_info.value) + + def test_rejects_unsupported_schema_version( valid_event_data: dict, ): From 8ba82ae4336bc07fb46fbe85282f148b7ac047f7 Mon Sep 17 00:00:00 2001 From: pearseona Date: Fri, 28 Aug 2026 13:48:25 +0900 Subject: [PATCH 3/4] feat: enforce request input size limits to prevent DoS and LLM cost blowup --- .env.example | 8 +++++ .env.prod.example | 9 +++++ app/chat/schemas.py | 30 ++++++++++++++-- app/core/config.py | 16 +++++++++ app/core/middleware.py | 36 +++++++++++++++++++ app/main.py | 7 ++++ tests/chat/test_schemas.py | 68 +++++++++++++++++++++++++++++++++++ tests/core/test_config.py | 18 ++++++++++ tests/core/test_middleware.py | 44 +++++++++++++++++++++++ 9 files changed, 234 insertions(+), 2 deletions(-) create mode 100644 app/core/middleware.py create mode 100644 tests/core/test_middleware.py diff --git a/.env.example b/.env.example index 17fec7b..2794b20 100644 --- a/.env.example +++ b/.env.example @@ -9,6 +9,14 @@ LLM_MAX_RETRIES=2 LLM_MAX_OUTPUT_TOKENS=1024 LLM_TEMPERATURE=0.1 +# 입력 크기 제한 (issue #120). 분석/챗 content는 SafeFam_BE @Size와 정합(5000/2000). +MAX_ANALYSIS_CONTENT_LENGTH=5000 +MAX_CHAT_CONTENT_LENGTH=2000 +MAX_CHAT_MESSAGES=40 +MAX_CHAT_CONTEXT_TEXT_LENGTH=2000 +MAX_CHAT_INDICATORS=20 +MAX_REQUEST_BODY_BYTES=1048576 + VIRUSTOTAL_API_KEY= GOOGLE_SAFE_BROWSING_API_KEY= MOCK_SECURITY_API=false diff --git a/.env.prod.example b/.env.prod.example index 28e1bfb..dd0be9a 100644 --- a/.env.prod.example +++ b/.env.prod.example @@ -8,6 +8,15 @@ LLM_TIMEOUT_SECONDS=15 LLM_MAX_RETRIES=2 LLM_MAX_OUTPUT_TOKENS=1024 LLM_TEMPERATURE=0.1 + +# 입력 크기 제한 (issue #120). 분석/챗 content는 SafeFam_BE @Size와 정합(5000/2000). +MAX_ANALYSIS_CONTENT_LENGTH=5000 +MAX_CHAT_CONTENT_LENGTH=2000 +MAX_CHAT_MESSAGES=40 +MAX_CHAT_CONTEXT_TEXT_LENGTH=2000 +MAX_CHAT_INDICATORS=20 +MAX_REQUEST_BODY_BYTES=1048576 + VIRUSTOTAL_API_KEY= GOOGLE_SAFE_BROWSING_API_KEY= diff --git a/app/chat/schemas.py b/app/chat/schemas.py index cbdf5ba..e0ffb61 100644 --- a/app/chat/schemas.py +++ b/app/chat/schemas.py @@ -14,16 +14,25 @@ class ChatRole(str, Enum): class Indicator(BaseModel): """탐지 근거 하나. type/description 값 목록이 아직 확정되지 않아 자유 문자열로 수용""" - model_config = ConfigDict(extra="forbid") + # hide_input_in_errors: 검증 실패 시 원문이 에러에 담기지 않도록 한다. + model_config = ConfigDict(extra="forbid", hide_input_in_errors=True) type: str description: str + @field_validator("type", "description") + @classmethod + def within_context_text_limit(cls, value: str) -> str: + # LLM 프롬프트로 주입되므로 무제한 유입을 막는다(issue #120). + if len(value) > settings.MAX_CHAT_CONTEXT_TEXT_LENGTH: + raise ValueError("indicator field exceeds max length") + return value + class AnalysisContext(BaseModel): """Spring Boot가 /analyze 결과를 바탕으로 구성해 전달하는 분석 컨텍스트""" - model_config = ConfigDict(extra="forbid") + model_config = ConfigDict(extra="forbid", hide_input_in_errors=True) riskScore: int = Field(..., ge=0, le=100, description="최종 위험 점수 (0~100)") riskLevel: RiskGrade = Field(..., description="위험 등급 (HIGH/MEDIUM/LOW)") @@ -33,11 +42,28 @@ class AnalysisContext(BaseModel): default_factory=list, description="탐지 근거 목록" ) + @field_validator("category") + @classmethod + def category_within_limit(cls, value: str) -> str: + # LLM 프롬프트로 주입되므로 무제한 유입을 막는다(issue #120). + if len(value) > settings.MAX_CHAT_CONTEXT_TEXT_LENGTH: + raise ValueError("category exceeds max length") + return value + @field_validator("explanation") @classmethod def explanation_must_not_be_blank(cls, value: str) -> str: if not value.strip(): raise ValueError("explanation must not be blank") + if len(value) > settings.MAX_CHAT_CONTEXT_TEXT_LENGTH: + raise ValueError("explanation exceeds max length") + return value + + @field_validator("indicators") + @classmethod + def indicators_within_limit(cls, value: list[Indicator]) -> list[Indicator]: + if len(value) > settings.MAX_CHAT_INDICATORS: + raise ValueError("indicators exceeds max count") return value diff --git a/app/core/config.py b/app/core/config.py index e238e8f..32e86d7 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -85,6 +85,22 @@ class Settings(BaseSettings): default=40, ge=1, ) + # analysisContext는 BE가 생성해 전달하며 LLM 시스템 프롬프트로 주입되므로 + # (explanation/category/indicators) 무제한 유입을 막는 안전 상한을 둔다. + MAX_CHAT_CONTEXT_TEXT_LENGTH: int = Field( + default=2_000, + ge=1, + ) + MAX_CHAT_INDICATORS: int = Field( + default=20, + ge=1, + ) + # HTTP 바디 크기 상한(바이트). 스키마 문자 제한 이전에 초대형 바디를 파싱 전 차단. + # 스키마상 최대 유효 요청(챗 40메시지+컨텍스트)이 대략 0.5MiB라 1MiB로 여유를 둔다. + MAX_REQUEST_BODY_BYTES: int = Field( + default=1_048_576, + ge=1, + ) # Stacking 모델 임계값 STACKING_NORMAL_PROBABILITY_MAX: float = Field( diff --git a/app/core/middleware.py b/app/core/middleware.py new file mode 100644 index 0000000..26eb2a3 --- /dev/null +++ b/app/core/middleware.py @@ -0,0 +1,36 @@ +from starlette.datastructures import Headers +from starlette.responses import JSONResponse +from starlette.types import ASGIApp, Receive, Scope, Send + + +class BodySizeLimitMiddleware: + """Content-Length가 상한을 초과하는 요청을 파싱 전에 413으로 거부한다(issue #120). + + 스키마 레벨 문자 수 제한이 1차 방어이며, 이 미들웨어는 초대형 바디가 메모리에 + 적재/역직렬화되기 전에 차단하는 방어 심화(defense-in-depth) 계층이다. + 정상 JSON 요청(BE의 httpx, 테스트 클라이언트 등)은 Content-Length를 항상 보낸다. + """ + + def __init__(self, app: ASGIApp, max_body_bytes: int) -> None: + self.app = app + self.max_body_bytes = max_body_bytes + + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + if scope["type"] != "http": + await self.app(scope, receive, send) + return + + content_length = Headers(scope=scope).get("content-length") + if ( + content_length is not None + and content_length.isdigit() + and int(content_length) > self.max_body_bytes + ): + response = JSONResponse( + status_code=413, + content={"detail": "Request body too large"}, + ) + await response(scope, receive, send) + return + + await self.app(scope, receive, send) diff --git a/app/main.py b/app/main.py index c4b8a57..d80d8d4 100644 --- a/app/main.py +++ b/app/main.py @@ -9,6 +9,7 @@ from app.analysis.text.stacking_analyzer import is_stacking_model_loaded from app.chat import router as chat from app.core.config import settings +from app.core.middleware import BodySizeLimitMiddleware from app.infrastructure.rabbitmq.connection import ( RabbitMQConnection, ) @@ -118,6 +119,12 @@ def create_app( allow_headers=["*"], ) + # 초대형 바디를 파싱 전에 차단(가장 바깥에서 먼저 실행되도록 마지막에 추가). + application.add_middleware( + BodySizeLimitMiddleware, + max_body_bytes=settings.MAX_REQUEST_BODY_BYTES, + ) + application.include_router(analyze.router, prefix="/api") application.include_router(chat.router, prefix="/api") diff --git a/tests/chat/test_schemas.py b/tests/chat/test_schemas.py index 67ccc77..4b220fb 100644 --- a/tests/chat/test_schemas.py +++ b/tests/chat/test_schemas.py @@ -116,6 +116,74 @@ def test_analysis_context_indicator_requires_type_and_description(): _analysis_context(indicators=[{"description": "악성 이력이 확인된 URL입니다."}]) +def test_analysis_context_accepts_explanation_at_max_length(): + context = _analysis_context( + explanation="가" * settings.MAX_CHAT_CONTEXT_TEXT_LENGTH + ) + + assert len(context.explanation) == settings.MAX_CHAT_CONTEXT_TEXT_LENGTH + + +def test_analysis_context_rejects_explanation_over_max_length(): + with pytest.raises(ValidationError): + _analysis_context( + explanation="가" * (settings.MAX_CHAT_CONTEXT_TEXT_LENGTH + 1) + ) + + +def test_analysis_context_rejects_category_over_max_length(): + with pytest.raises(ValidationError): + _analysis_context( + category="A" * (settings.MAX_CHAT_CONTEXT_TEXT_LENGTH + 1) + ) + + +def test_analysis_context_rejects_indicator_description_over_max_length(): + with pytest.raises(ValidationError): + _analysis_context( + indicators=[ + { + "type": "MALICIOUS_URL", + "description": "가" + * (settings.MAX_CHAT_CONTEXT_TEXT_LENGTH + 1), + } + ] + ) + + +def test_analysis_context_accepts_indicators_at_max_count(): + context = _analysis_context( + indicators=[ + {"type": "MALICIOUS_URL", "description": "악성 URL"} + for _ in range(settings.MAX_CHAT_INDICATORS) + ] + ) + + assert len(context.indicators) == settings.MAX_CHAT_INDICATORS + + +def test_analysis_context_rejects_indicators_over_max_count(): + with pytest.raises(ValidationError): + _analysis_context( + indicators=[ + {"type": "MALICIOUS_URL", "description": "악성 URL"} + for _ in range(settings.MAX_CHAT_INDICATORS + 1) + ] + ) + + +def test_analysis_context_oversized_explanation_is_not_leaked_in_error(): + """hide_input_in_errors=True 로 컨텍스트 원문이 에러에 담기지 않아야 한다.""" + secret_marker = "01012345678-비밀번호" + with pytest.raises(ValidationError) as exception_info: + _analysis_context( + explanation=secret_marker + + "가" * settings.MAX_CHAT_CONTEXT_TEXT_LENGTH + ) + + assert secret_marker not in str(exception_info.value) + + def test_chat_request_rejects_unknown_fields(): with pytest.raises(ValidationError): ChatRequest( diff --git a/tests/core/test_config.py b/tests/core/test_config.py index a4aa752..9301f9e 100644 --- a/tests/core/test_config.py +++ b/tests/core/test_config.py @@ -61,12 +61,30 @@ def test_input_size_limit_defaults_match_backend_contract(): assert configured.MAX_CHAT_MESSAGES == 40 +def test_chat_context_size_limit_defaults(): + """analysisContext(LLM 프롬프트 주입)의 안전 상한 기본값(issue #120).""" + configured = Settings(_env_file=None) + + assert configured.MAX_CHAT_CONTEXT_TEXT_LENGTH == 2000 + assert configured.MAX_CHAT_INDICATORS == 20 + + +def test_request_body_size_limit_default(): + """HTTP 바디 크기 상한 기본값 1MiB(issue #120).""" + configured = Settings(_env_file=None) + + assert configured.MAX_REQUEST_BODY_BYTES == 1_048_576 + + @pytest.mark.parametrize( "field_name", [ "MAX_ANALYSIS_CONTENT_LENGTH", "MAX_CHAT_CONTENT_LENGTH", "MAX_CHAT_MESSAGES", + "MAX_CHAT_CONTEXT_TEXT_LENGTH", + "MAX_CHAT_INDICATORS", + "MAX_REQUEST_BODY_BYTES", ], ) def test_input_size_limits_reject_non_positive(field_name: str): diff --git a/tests/core/test_middleware.py b/tests/core/test_middleware.py new file mode 100644 index 0000000..0991b0d --- /dev/null +++ b/tests/core/test_middleware.py @@ -0,0 +1,44 @@ +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from app.core.middleware import BodySizeLimitMiddleware + + +def _build_client(max_body_bytes: int) -> TestClient: + application = FastAPI() + application.add_middleware( + BodySizeLimitMiddleware, + max_body_bytes=max_body_bytes, + ) + + @application.post("/echo") + async def echo(payload: dict) -> dict: + return {"ok": True} + + return TestClient(application) + + +def test_allows_body_within_limit(): + client = _build_client(max_body_bytes=1_000) + + response = client.post("/echo", json={"text": "hi"}) + + assert response.status_code == 200 + + +def test_rejects_body_over_limit_before_parsing(): + client = _build_client(max_body_bytes=50) + + response = client.post("/echo", json={"text": "x" * 500}) + + assert response.status_code == 413 + assert response.json() == {"detail": "Request body too large"} + + +def test_multibyte_body_within_limit_passes(): + """멀티바이트(한글) 본문도 상한 이내면 그대로 통과한다.""" + client = _build_client(max_body_bytes=2_000) + + response = client.post("/echo", json={"text": "가" * 100}) + + assert response.status_code == 200 From 2806a2c55c3a4204531058977d2f5d37b97e7160 Mon Sep 17 00:00:00 2001 From: pearseona Date: Fri, 28 Aug 2026 16:51:24 +0900 Subject: [PATCH 4/4] refactor: decompose analyze_pipeline into cohesive private methods --- app/analysis/schemas.py | 4 + app/analysis/service.py | 737 +++++++++++--------- app/chat/schemas.py | 18 + app/core/exception_handlers.py | 39 ++ app/core/middleware.py | 76 ++ app/infrastructure/rabbitmq/schemas.py | 1 + app/main.py | 14 + tests/core/test_middleware.py | 75 ++ tests/core/test_validation_error_handler.py | 88 +++ 9 files changed, 725 insertions(+), 327 deletions(-) create mode 100644 app/core/exception_handlers.py create mode 100644 tests/core/test_validation_error_handler.py diff --git a/app/analysis/schemas.py b/app/analysis/schemas.py index bf1e963..2f2fa0d 100644 --- a/app/analysis/schemas.py +++ b/app/analysis/schemas.py @@ -15,6 +15,10 @@ class SmishingAnalysisRequest(BaseModel): @field_validator("text") @classmethod def validate_text(cls, value: str) -> str: +<<<<<<< HEAD +======= + """공백이거나 상한을 초과하는 검사 요청 본문을 거부한다(issue #120).""" +>>>>>>> f57ae1c1b7e725f9df63b20d31f6be5505de2a1c if not value.strip(): raise ValueError("text must not be blank") # 분석 content와 동일한 상한(BE @Size(max=5000)와 정합)을 공유한다. diff --git a/app/analysis/service.py b/app/analysis/service.py index 0add472..4d87bef 100644 --- a/app/analysis/service.py +++ b/app/analysis/service.py @@ -2,6 +2,7 @@ import asyncio import logging from collections.abc import Callable +from dataclasses import dataclass from app.analysis.evidence import build_evidence from app.analysis.hybrid_policy import ( @@ -18,7 +19,6 @@ ContributionBreakdown, RiskGrade, SmishingAnalysisResponse, - UrlAnalysisResponse, ) from app.analysis.scoring import RiskScoringEngine from app.analysis.text.llm_analyzer import ( @@ -42,6 +42,31 @@ logger = logging.getLogger(__name__) + +@dataclass +class _TrackOutcome: + """텍스트/URL 트랙 실행 결과를 파이프라인 후속 단계로 전달하는 내부 상태. + + 응답에 직렬화되지 않으며 오케스트레이션 단계 간 값 전달 용도로만 사용한다. + """ + + text_analysis: dict + has_url: bool + original_url: str | None + traced_url: str | None + hybrid_url_result: dict + + +@dataclass +class _TrackAvailability: + """각 트랙의 사용 가능 여부와 스코어링 입력 점수를 담는 내부 상태.""" + + text_available: bool + scoring_text_score: int + url_available: bool + rules_available: bool + + class SmishingAnalysisService: """텍스트, URL, 규칙 분석 결과를 최종 응답으로 조립""" @@ -122,204 +147,28 @@ async def analyze_pipeline( """텍스트, URL, 규칙 분석을 실행하고 최종 점수를 계산""" try: - # 메시지에서 URL을 먼저 추출 - urls = extract_urls(text) - has_url = len(urls) > 0 - - rule_score_preview = self._preview_rule_score( - text - ) - force_llm = ( - rule_score_preview - >= RISK_MEDIUM_THRESHOLD - ) - - # 텍스트 분석은 항상 실행 - text_task = asyncio.create_task( - self.text_analyzer.analyze( - text, - force_llm=force_llm, - ) - ) - - url_task = None - original_url = None - - # URL이 있는 경우에만 URL 추적 및 보안 분석 실행 - if has_url: - original_url = urls[0] - - async def url_track(): - """단축 URL을 추적한 뒤 보안 엔진으로 검사""" - - traced = await trace_url( - original_url - ) - - analysis = ( - await self.url_analyzer.scan_url( - traced - ) - ) - - return traced, analysis - - url_task = asyncio.create_task( - url_track() - ) + tracks = await self._run_tracks(text) - # 텍스트와 URL 분석을 가능한 한 병렬로 실행 - if url_task is not None: - try: - ( - text_analysis, - ( - traced_url, - hybrid_url_result, - ), - ) = await asyncio.gather( - text_task, - url_task, - ) - except BaseException: - # 한 트랙이 실패하면 아직 실행 중인 형제 task를 - # 취소하고 두 결과를 모두 회수해 orphan task와 - # "Task exception was never retrieved"를 방지한다. - for task in (text_task, url_task): - if not task.done(): - task.cancel() - - await asyncio.gather( - text_task, - url_task, - return_exceptions=True, - ) - raise - else: - text_analysis = await text_task - - traced_url = None - - # URL이 없는 것은 URL 분석 실패가 아님 - # 단순히 URL 트랙이 적용되지 않은 상태 - hybrid_url_result = { - "is_malicious": False, - "url_risk_score": 0.0, - "source": ( - "Pre-Processing-Filter" - ), - "available": False, - "failed_providers": [], - "pending_providers": [], - "provider_error_codes": {}, - "error_message": None, - "is_gsb_confirmed": False, - "is_vt_confirmed": False, - } - - # HybridTextAnalyzer가 선택한 최종 텍스트 결과 - text_result = ( - text_analysis.get("result") or {} - ) - - # 하이브리드 분석기가 최종 선택한 텍스트 점수 - selected_text_score = ( - text_result.get("risk_score") - ) - - # 두 텍스트 엔진이 모두 실패한 경우 scoring engine이 fail-safe 점수를 적용할 수 있도록 0을 전달 - scoring_text_score = ( - int(selected_text_score) - if selected_text_score is not None - else 0 - ) - - # HybridTextAnalyzer가 선택한 단일 결과가 있을 때만 - # 텍스트 트랙을 사용 가능한 상태로 본다. - text_available = selected_text_score is not None - - # 로컬 규칙 분석 실행 - try: - rule_result = self.rule_analyzer( - text, - traced_url, - ) - except Exception as exception: - # 원문 메시지나 예외 메시지는 로그에 기록 X - logger.error( - "[Analysis Service] 규칙 분석 실패. " - "error_type=%s", - type(exception).__name__, - ) - - rule_result = { - "rule_score": 0, - "has_malicious_domain_pattern": ( - False - ), - "matched_rules": [], - "error_message": ( - "RULE_ANALYSIS_FAILED" - ), - } + rule_result = self._run_rules(text, tracks.traced_url) - rules_available = not bool( - rule_result.get("error_message") + availability = self._assess_track_availability( + text_analysis=tracks.text_analysis, + has_url=tracks.has_url, + hybrid_url_result=tracks.hybrid_url_result, + rule_result=rule_result, ) - # URL이 존재하면서 URL 보안 공급자 중 하나 이상이 - # 정상 결과를 제공했을 때만 URL 트랙을 available로 봄 - url_available = ( - has_url - and hybrid_url_result.get( - "available", - False, - ) - ) - - # 로컬 도메인 규칙에서 명확한 악성 패턴이 발견되면 - # URL 결과의 최소 위험도를 0.75로 올림 - if rule_result.get( - "has_malicious_domain_pattern", - False, - ): - hybrid_url_result[ - "is_malicious" - ] = True - - hybrid_url_result[ - "url_risk_score" - ] = max( - hybrid_url_result.get( - "url_risk_score", - 0.0, - ), - 0.75, - ) - - # 신뢰도가 높은 확정 악성 신호를 확인 - # 다음 신호 중 하나라도 있으면 scoring engine에서 최종 등급을 최소 HIGH로 보정 - is_confirmed_malicious = ( - hybrid_url_result.get( - "is_gsb_confirmed", - False, - ) - or hybrid_url_result.get( - "is_vt_confirmed", - False, - ) - or rule_result.get( - "has_malicious_domain_pattern", - False, - ) + is_confirmed_malicious = self._apply_confirmed_malicious_boost( + hybrid_url_result=tracks.hybrid_url_result, + rule_result=rule_result, ) # 텍스트, URL, 규칙이 전부 실패한 경우 # 정상 또는 LOW 응답을 반환하면 안됨 no_reliable_signal = ( - not text_available - and not url_available - and not rules_available + not availability.text_available + and not availability.url_available + and not availability.rules_available ) if no_reliable_signal: @@ -333,17 +182,17 @@ async def url_track(): breakdown, ) = RiskScoringEngine.calculate_score( # HybridTextAnalyzer가 선택한 최종 텍스트 점수 - llm_score=scoring_text_score, + llm_score=availability.scoring_text_score, is_url_malicious=( - hybrid_url_result.get( + tracks.hybrid_url_result.get( "is_malicious", False, ) ), url_risk_score=( - hybrid_url_result.get( + tracks.hybrid_url_result.get( "url_risk_score", 0.0, ) @@ -354,7 +203,7 @@ async def url_track(): 0, ), - has_url=has_url, + has_url=tracks.has_url, # HybridTextAnalyzer가 이미 하나의 최종 점수를 # 선택했으므로 자체 모델 점수를 다시 혼합하지 않는다. @@ -362,173 +211,407 @@ async def url_track(): # 인자명은 기존 호환성을 유지하지만 실제 의미는 # 선택된 텍스트 점수의 사용 가능 여부다. - llm_available=text_available, + llm_available=availability.text_available, is_confirmed_malicious=( is_confirmed_malicious ), - text_available=text_available, - url_available=url_available, - rules_available=rules_available, + text_available=availability.text_available, + url_available=availability.url_available, + rules_available=availability.rules_available, ) - # URL이 실제로 포함된 경우에만 - # URL 분석 상세 결과를 응답에 포함 - if has_url: - real_url_analysis = { - "has_url": True, + return self._assemble_success_response( + has_url=tracks.has_url, + original_url=tracks.original_url, + traced_url=tracks.traced_url, + hybrid_url_result=tracks.hybrid_url_result, + url_available=availability.url_available, + final_score=final_score, + risk_grade=risk_grade, + breakdown=breakdown, + rule_result=rule_result, + text_analysis=tracks.text_analysis, + ) - "is_shortened": ( - original_url != traced_url - ), + except Exception as exception: + # 예외 타입만 로그에 기록 + logger.error( + "[Analysis Service] 파이프라인 실패. " + "error_type=%s", + type(exception).__name__, + ) - "origin_url": traced_url, - "original_url": original_url, + return self._build_failure_response() - "is_url_malicious": ( - hybrid_url_result.get( - "is_malicious", - False, - ) - ), + async def _run_tracks( + self, + text: str, + ) -> _TrackOutcome: + """텍스트/URL 트랙을 병렬 실행하고 결과를 회수한다.""" - "url_risk_score": ( - hybrid_url_result.get( - "url_risk_score", - 0.0, - ) - ), + # 메시지에서 URL을 먼저 추출 + urls = extract_urls(text) + has_url = len(urls) > 0 - "engine_source": ( - hybrid_url_result.get( - "source", - "Hybrid-Engine", - ) - ), + rule_score_preview = self._preview_rule_score( + text + ) + force_llm = ( + rule_score_preview + >= RISK_MEDIUM_THRESHOLD + ) + + # 텍스트 분석은 항상 실행 + text_task = asyncio.create_task( + self.text_analyzer.analyze( + text, + force_llm=force_llm, + ) + ) - "available": url_available, + url_task = None + original_url = None - "failed_providers": ( - hybrid_url_result.get( - "failed_providers", - [], - ) - ), + # URL이 있는 경우에만 URL 추적 및 보안 분석 실행 + if has_url: + original_url = urls[0] - "pending_providers": ( - hybrid_url_result.get( - "pending_providers", - [], - ) - ), + async def url_track(): + """단축 URL을 추적한 뒤 보안 엔진으로 검사""" - "provider_error_codes": ( - hybrid_url_result.get( - "provider_error_codes", - {}, - ) - ), + traced = await trace_url( + original_url + ) - "error_message": ( - hybrid_url_result.get( - "error_message" - ) + analysis = ( + await self.url_analyzer.scan_url( + traced + ) + ) + + return traced, analysis + + url_task = asyncio.create_task( + url_track() + ) + + # 텍스트와 URL 분석을 가능한 한 병렬로 실행 + if url_task is not None: + try: + ( + text_analysis, + ( + traced_url, + hybrid_url_result, ), - } - else: - real_url_analysis = None - - # 원문 메시지를 응답 또는 로그에 추가 X - return SmishingAnalysisResponse( - status="SUCCESS", - message=( - "3중 가중치 결합 스미싱 " - "통합 분석이 완료되었습니다." - ), - final_score=final_score, - risk_grade=risk_grade, - contribution_breakdown=breakdown, - evidence=build_evidence( - rule_analysis=rule_result, - url_analysis=real_url_analysis, - text_analysis=text_analysis, + ) = await asyncio.gather( + text_task, + url_task, + ) + except BaseException: + # 한 트랙이 실패하면 아직 실행 중인 형제 task를 + # 취소하고 두 결과를 모두 회수해 orphan task와 + # "Task exception was never retrieved"를 방지한다. + for task in (text_task, url_task): + if not task.done(): + task.cancel() + + await asyncio.gather( + text_task, + url_task, + return_exceptions=True, + ) + raise + else: + text_analysis = await text_task + + traced_url = None + + # URL이 없는 것은 URL 분석 실패가 아님 + # 단순히 URL 트랙이 적용되지 않은 상태 + hybrid_url_result = { + "is_malicious": False, + "url_risk_score": 0.0, + "source": ( + "Pre-Processing-Filter" ), - text_analysis=text_analysis, - url_analysis=real_url_analysis, - rule_analysis=rule_result, - ) + "available": False, + "failed_providers": [], + "pending_providers": [], + "provider_error_codes": {}, + "error_message": None, + "is_gsb_confirmed": False, + "is_vt_confirmed": False, + } + + return _TrackOutcome( + text_analysis=text_analysis, + has_url=has_url, + original_url=original_url, + traced_url=traced_url, + hybrid_url_result=hybrid_url_result, + ) + + def _run_rules( + self, + text: str, + traced_url: str | None, + ) -> dict: + """로컬 규칙 분석을 실행하고 실패 시 안전한 기본 결과를 반환한다.""" + try: + return self.rule_analyzer( + text, + traced_url, + ) except Exception as exception: - # 예외 타입만 로그에 기록 + # 원문 메시지나 예외 메시지는 로그에 기록 X logger.error( - "[Analysis Service] 파이프라인 실패. " + "[Analysis Service] 규칙 분석 실패. " "error_type=%s", type(exception).__name__, ) - # 전체 파이프라인 실패를 0점/LOW로 반환하면 - # 장애가 안전 판정으로 해석되는 fail-open이 발생 - return SmishingAnalysisResponse( - status="ERROR", - message=( - "분석 파이프라인 처리 중 " - "오류가 발생했습니다." + return { + "rule_score": 0, + "has_malicious_domain_pattern": ( + False ), - final_score=( - RiskScoringEngine - .PIPELINE_FAILURE_FALLBACK_SCORE + "matched_rules": [], + "error_message": ( + "RULE_ANALYSIS_FAILED" ), - risk_grade=RiskGrade.MEDIUM, - contribution_breakdown=( - ContributionBreakdown( - llm=0, - hybrid_url=0, - rules=0, - ) - ), - text_analysis=None, - url_analysis=None, - rule_analysis=None, - ) + } - async def scan_message_text( + def _assess_track_availability( self, - message: str, - ) -> UrlAnalysisResponse: - """메시지에 포함된 첫 번째 URL을 추적하고 검사""" - - urls = extract_urls(message) - - if not urls: - return UrlAnalysisResponse( - has_url=False, - original_url=None, - traced_url=None, - is_url_malicious=False, - url_risk_score=0.0, - engine_source=( - "Pre-Processing-Filter" - ), + *, + text_analysis: dict, + has_url: bool, + hybrid_url_result: dict, + rule_result: dict, + ) -> _TrackAvailability: + """각 트랙의 사용 가능 여부와 스코어링 입력 점수를 판정한다.""" + + # HybridTextAnalyzer가 선택한 최종 텍스트 결과 + text_result = ( + text_analysis.get("result") or {} + ) + + # 하이브리드 분석기가 최종 선택한 텍스트 점수 + selected_text_score = ( + text_result.get("risk_score") + ) + + # 두 텍스트 엔진이 모두 실패한 경우 scoring engine이 fail-safe 점수를 적용할 수 있도록 0을 전달 + scoring_text_score = ( + int(selected_text_score) + if selected_text_score is not None + else 0 + ) + + # HybridTextAnalyzer가 선택한 단일 결과가 있을 때만 + # 텍스트 트랙을 사용 가능한 상태로 본다. + text_available = selected_text_score is not None + + rules_available = not bool( + rule_result.get("error_message") + ) + + # URL이 존재하면서 URL 보안 공급자 중 하나 이상이 + # 정상 결과를 제공했을 때만 URL 트랙을 available로 봄 + url_available = ( + has_url + and hybrid_url_result.get( + "available", + False, ) - traced_url = await trace_url(urls[0]) + ) - result = await self.url_analyzer.scan_url( - traced_url + return _TrackAvailability( + text_available=text_available, + scoring_text_score=scoring_text_score, + url_available=url_available, + rules_available=rules_available, ) - return UrlAnalysisResponse( - has_url=True, - original_url=urls[0], - traced_url=traced_url, - is_url_malicious=result[ + def _apply_confirmed_malicious_boost( + self, + *, + hybrid_url_result: dict, + rule_result: dict, + ) -> bool: + """확정 악성 신호를 반영해 URL 위험도를 보정하고 확정 여부를 반환한다.""" + + # 로컬 도메인 규칙에서 명확한 악성 패턴이 발견되면 + # URL 결과의 최소 위험도를 0.75로 올림 + if rule_result.get( + "has_malicious_domain_pattern", + False, + ): + hybrid_url_result[ "is_malicious" - ], - url_risk_score=result[ + ] = True + + hybrid_url_result[ "url_risk_score" - ], - engine_source=result["source"], - error_message=result[ - "error_message" - ], + ] = max( + hybrid_url_result.get( + "url_risk_score", + 0.0, + ), + 0.75, + ) + + # 신뢰도가 높은 확정 악성 신호를 확인 + # 다음 신호 중 하나라도 있으면 scoring engine에서 최종 등급을 최소 HIGH로 보정 + is_confirmed_malicious = ( + hybrid_url_result.get( + "is_gsb_confirmed", + False, + ) + or hybrid_url_result.get( + "is_vt_confirmed", + False, + ) + or rule_result.get( + "has_malicious_domain_pattern", + False, + ) + ) + + return is_confirmed_malicious + + def _assemble_success_response( + self, + *, + has_url: bool, + original_url: str | None, + traced_url: str | None, + hybrid_url_result: dict, + url_available: bool, + final_score: int, + risk_grade: RiskGrade, + breakdown: ContributionBreakdown, + rule_result: dict, + text_analysis: dict, + ) -> SmishingAnalysisResponse: + """분석 결과를 최종 응답 스키마로 조립한다.""" + + # URL이 실제로 포함된 경우에만 + # URL 분석 상세 결과를 응답에 포함 + if has_url: + real_url_analysis = { + "has_url": True, + + "is_shortened": ( + original_url != traced_url + ), + + "origin_url": traced_url, + "original_url": original_url, + + "is_url_malicious": ( + hybrid_url_result.get( + "is_malicious", + False, + ) + ), + + "url_risk_score": ( + hybrid_url_result.get( + "url_risk_score", + 0.0, + ) + ), + + "engine_source": ( + hybrid_url_result.get( + "source", + "Hybrid-Engine", + ) + ), + + "available": url_available, + + "failed_providers": ( + hybrid_url_result.get( + "failed_providers", + [], + ) + ), + + "pending_providers": ( + hybrid_url_result.get( + "pending_providers", + [], + ) + ), + + "provider_error_codes": ( + hybrid_url_result.get( + "provider_error_codes", + {}, + ) + ), + + "error_message": ( + hybrid_url_result.get( + "error_message" + ) + ), + } + else: + real_url_analysis = None + + # 원문 메시지를 응답 또는 로그에 추가 X + return SmishingAnalysisResponse( + status="SUCCESS", + message=( + "3중 가중치 결합 스미싱 " + "통합 분석이 완료되었습니다." + ), + final_score=final_score, + risk_grade=risk_grade, + contribution_breakdown=breakdown, + evidence=build_evidence( + rule_analysis=rule_result, + url_analysis=real_url_analysis, + text_analysis=text_analysis, + ), + text_analysis=text_analysis, + url_analysis=real_url_analysis, + rule_analysis=rule_result, + ) + + def _build_failure_response( + self, + ) -> SmishingAnalysisResponse: + """파이프라인 실패 시 fail-open을 막는 안전 응답을 반환한다.""" + + # 전체 파이프라인 실패를 0점/LOW로 반환하면 + # 장애가 안전 판정으로 해석되는 fail-open이 발생 + return SmishingAnalysisResponse( + status="ERROR", + message=( + "분석 파이프라인 처리 중 " + "오류가 발생했습니다." + ), + final_score=( + RiskScoringEngine + .PIPELINE_FAILURE_FALLBACK_SCORE + ), + risk_grade=RiskGrade.MEDIUM, + contribution_breakdown=( + ContributionBreakdown( + llm=0, + hybrid_url=0, + rules=0, + ) + ), + text_analysis=None, + url_analysis=None, + rule_analysis=None, ) diff --git a/app/chat/schemas.py b/app/chat/schemas.py index e0ffb61..77d572c 100644 --- a/app/chat/schemas.py +++ b/app/chat/schemas.py @@ -23,7 +23,11 @@ class Indicator(BaseModel): @field_validator("type", "description") @classmethod def within_context_text_limit(cls, value: str) -> str: +<<<<<<< HEAD # LLM 프롬프트로 주입되므로 무제한 유입을 막는다(issue #120). +======= + """LLM 프롬프트로 주입되는 근거 문자열의 무제한 유입을 막는다(issue #120).""" +>>>>>>> f57ae1c1b7e725f9df63b20d31f6be5505de2a1c if len(value) > settings.MAX_CHAT_CONTEXT_TEXT_LENGTH: raise ValueError("indicator field exceeds max length") return value @@ -45,7 +49,11 @@ class AnalysisContext(BaseModel): @field_validator("category") @classmethod def category_within_limit(cls, value: str) -> str: +<<<<<<< HEAD # LLM 프롬프트로 주입되므로 무제한 유입을 막는다(issue #120). +======= + """LLM 프롬프트로 주입되는 category의 무제한 유입을 막는다(issue #120).""" +>>>>>>> f57ae1c1b7e725f9df63b20d31f6be5505de2a1c if len(value) > settings.MAX_CHAT_CONTEXT_TEXT_LENGTH: raise ValueError("category exceeds max length") return value @@ -53,6 +61,7 @@ def category_within_limit(cls, value: str) -> str: @field_validator("explanation") @classmethod def explanation_must_not_be_blank(cls, value: str) -> str: + """공백이거나 상한을 초과하는 분석 요약을 거부한다(issue #120).""" if not value.strip(): raise ValueError("explanation must not be blank") if len(value) > settings.MAX_CHAT_CONTEXT_TEXT_LENGTH: @@ -62,6 +71,10 @@ def explanation_must_not_be_blank(cls, value: str) -> str: @field_validator("indicators") @classmethod def indicators_within_limit(cls, value: list[Indicator]) -> list[Indicator]: +<<<<<<< HEAD +======= + """탐지 근거 개수 상한을 강제해 프롬프트 팽창을 막는다(issue #120).""" +>>>>>>> f57ae1c1b7e725f9df63b20d31f6be5505de2a1c if len(value) > settings.MAX_CHAT_INDICATORS: raise ValueError("indicators exceeds max count") return value @@ -77,6 +90,7 @@ class ChatMessage(BaseModel): @field_validator("content") @classmethod def content_must_not_be_blank(cls, value: str) -> str: + """공백이거나 상한을 초과하는 챗 메시지를 거부한다(issue #120).""" if not value.strip(): raise ValueError("content must not be blank") # BE ChatMessage @Size(max=2000)와 정합. @@ -94,6 +108,10 @@ class ChatRequest(BaseModel): @field_validator("messages") @classmethod def messages_within_limit(cls, value: list[ChatMessage]) -> list[ChatMessage]: +<<<<<<< HEAD +======= + """대화 메시지 개수 상한을 강제해 LLM 비용 폭증을 막는다(issue #120).""" +>>>>>>> f57ae1c1b7e725f9df63b20d31f6be5505de2a1c if len(value) > settings.MAX_CHAT_MESSAGES: raise ValueError("messages exceeds max count") return value diff --git a/app/core/exception_handlers.py b/app/core/exception_handlers.py new file mode 100644 index 0000000..c8fd4c6 --- /dev/null +++ b/app/core/exception_handlers.py @@ -0,0 +1,39 @@ +from fastapi import Request +from fastapi.exceptions import RequestValidationError +from fastapi.responses import JSONResponse + +# 422 Unprocessable Content. 상수명이 starlette 버전에 따라 +# (ENTITY/CONTENT) 달라 deprecation을 피하려 리터럴을 사용한다. +_HTTP_422 = 422 + + +def _sanitize_errors(errors: list[dict]) -> list[dict]: + """검증 에러에서 원문(PII 가능) `input`과 직렬화 불가한 `ctx`를 제거한다. + + `type`/`loc`/`msg`만 남겨 클라이언트가 어떤 필드가 왜 거부됐는지는 알되, + 보낸 원문이 응답에 되돌아오지 않도록 한다. + """ + return [ + { + "type": error.get("type"), + "loc": error.get("loc"), + "msg": error.get("msg"), + } + for error in errors + ] + + +async def validation_exception_handler( + request: Request, + exc: RequestValidationError, +) -> JSONResponse: + """요청 본문을 되비추지 않는 422 응답을 반환한다. + + FastAPI 기본 핸들러는 `exc.errors()`를 그대로 직렬화하는데, 이 목록은 + `hide_input_in_errors`를 켜도 원문 `input`을 포함한다. 안티스미싱 서비스에서 + 입력은 문자 본문(PII)이므로 응답에 절대 반사되면 안 된다. + """ + return JSONResponse( + status_code=_HTTP_422, + content={"detail": _sanitize_errors(exc.errors())}, + ) diff --git a/app/core/middleware.py b/app/core/middleware.py index 26eb2a3..f987d74 100644 --- a/app/core/middleware.py +++ b/app/core/middleware.py @@ -1,5 +1,6 @@ from starlette.datastructures import Headers from starlette.responses import JSONResponse +<<<<<<< HEAD from starlette.types import ASGIApp, Receive, Scope, Send @@ -9,6 +10,19 @@ class BodySizeLimitMiddleware: 스키마 레벨 문자 수 제한이 1차 방어이며, 이 미들웨어는 초대형 바디가 메모리에 적재/역직렬화되기 전에 차단하는 방어 심화(defense-in-depth) 계층이다. 정상 JSON 요청(BE의 httpx, 테스트 클라이언트 등)은 Content-Length를 항상 보낸다. +======= +from starlette.types import ASGIApp, Message, Receive, Scope, Send + + +class BodySizeLimitMiddleware: + """요청 본문이 상한을 초과하면 파싱 전에 413으로 거부한다(issue #120). + + 스키마 레벨 문자 수 제한이 1차 방어이며, 이 미들웨어는 초대형 바디가 + 역직렬화되기 전에 차단하는 방어 심화(defense-in-depth) 계층이다. + Content-Length 헤더에만 의존하지 않는다 — 헤더가 없거나 위조된(청크 전송 등) + 경우에도 ASGI `http.request` 프레임을 상한까지만 버퍼링해 초과 시 거부하므로 + 우회할 수 없다. 메모리 사용량은 `max_body_bytes`로 제한된다. +>>>>>>> f57ae1c1b7e725f9df63b20d31f6be5505de2a1c """ def __init__(self, app: ASGIApp, max_body_bytes: int) -> None: @@ -16,16 +30,25 @@ def __init__(self, app: ASGIApp, max_body_bytes: int) -> None: self.max_body_bytes = max_body_bytes async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: +<<<<<<< HEAD +======= + """HTTP 요청 본문 크기를 강제하고, 이내면 버퍼를 앱에 재생(replay)한다.""" +>>>>>>> f57ae1c1b7e725f9df63b20d31f6be5505de2a1c if scope["type"] != "http": await self.app(scope, receive, send) return +<<<<<<< HEAD +======= + # 빠른 경로: 선언된 Content-Length가 이미 상한을 넘으면 즉시 거부. +>>>>>>> f57ae1c1b7e725f9df63b20d31f6be5505de2a1c content_length = Headers(scope=scope).get("content-length") if ( content_length is not None and content_length.isdigit() and int(content_length) > self.max_body_bytes ): +<<<<<<< HEAD response = JSONResponse( status_code=413, content={"detail": "Request body too large"}, @@ -34,3 +57,56 @@ async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: return await self.app(scope, receive, send) +======= + await self._reject(scope, receive, send) + return + + # 헤더와 무관하게 실제 바이트를 상한까지만 누적한다. + body = bytearray() + more_body = True + while more_body: + message = await receive() + if message["type"] == "http.disconnect": + await self._replay(scope, send, bytes(body), disconnected=True) + return + body.extend(message.get("body", b"")) + if len(body) > self.max_body_bytes: + await self._reject(scope, receive, send) + return + more_body = message.get("more_body", False) + + await self._replay(scope, send, bytes(body)) + + async def _reject(self, scope: Scope, receive: Receive, send: Send) -> None: + """413 Payload Too Large 응답을 전송한다(원문 미포함).""" + response = JSONResponse( + status_code=413, + content={"detail": "Request body too large"}, + ) + await response(scope, receive, send) + + async def _replay( + self, + scope: Scope, + send: Send, + body: bytes, + *, + disconnected: bool = False, + ) -> None: + """버퍼링한 본문을 단일 프레임으로 앱에 재생한다.""" + if disconnected: + messages: list[Message] = [{"type": "http.disconnect"}] + else: + messages = [ + {"type": "http.request", "body": body, "more_body": False} + ] + iterator = iter(messages) + + async def replay_receive() -> Message: + try: + return next(iterator) + except StopIteration: + return {"type": "http.disconnect"} + + await self.app(scope, replay_receive, send) +>>>>>>> f57ae1c1b7e725f9df63b20d31f6be5505de2a1c diff --git a/app/infrastructure/rabbitmq/schemas.py b/app/infrastructure/rabbitmq/schemas.py index 09a45a5..37d4d7d 100644 --- a/app/infrastructure/rabbitmq/schemas.py +++ b/app/infrastructure/rabbitmq/schemas.py @@ -37,6 +37,7 @@ class AnalysisRequestedPayload(BaseModel): @field_validator("content") @classmethod def validate_content(cls, value: str) -> str: + """공백이거나 상한을 초과하는 문자 본문을 거부한다(issue #120).""" if not value.strip(): raise ValueError("content must not be blank") if len(value) > settings.MAX_ANALYSIS_CONTENT_LENGTH: diff --git a/app/main.py b/app/main.py index d80d8d4..5ed4dd8 100644 --- a/app/main.py +++ b/app/main.py @@ -2,6 +2,7 @@ from contextlib import asynccontextmanager from fastapi import FastAPI +from fastapi.exceptions import RequestValidationError from fastapi.middleware.cors import CORSMiddleware from app.analysis import router as analyze @@ -9,6 +10,10 @@ from app.analysis.text.stacking_analyzer import is_stacking_model_loaded from app.chat import router as chat from app.core.config import settings +<<<<<<< HEAD +======= +from app.core.exception_handlers import validation_exception_handler +>>>>>>> f57ae1c1b7e725f9df63b20d31f6be5505de2a1c from app.core.middleware import BodySizeLimitMiddleware from app.infrastructure.rabbitmq.connection import ( RabbitMQConnection, @@ -125,6 +130,15 @@ def create_app( max_body_bytes=settings.MAX_REQUEST_BODY_BYTES, ) +<<<<<<< HEAD +======= + # 검증 실패(422) 응답에서 원문(PII)이 반사되지 않도록 처리. + application.add_exception_handler( + RequestValidationError, + validation_exception_handler, + ) + +>>>>>>> f57ae1c1b7e725f9df63b20d31f6be5505de2a1c application.include_router(analyze.router, prefix="/api") application.include_router(chat.router, prefix="/api") diff --git a/tests/core/test_middleware.py b/tests/core/test_middleware.py index 0991b0d..ad7c08e 100644 --- a/tests/core/test_middleware.py +++ b/tests/core/test_middleware.py @@ -42,3 +42,78 @@ def test_multibyte_body_within_limit_passes(): response = client.post("/echo", json={"text": "가" * 100}) assert response.status_code == 200 +<<<<<<< HEAD +======= + + +# --- ASGI 레벨: Content-Length 헤더 없이 다중 프레임으로 오는 요청 방어 --- + + +async def _echo_asgi_app(scope, receive, send) -> None: + """수신 본문 길이를 그대로 돌려주는 최소 ASGI 앱.""" + body = b"" + more_body = True + while more_body: + message = await receive() + if message["type"] != "http.request": + more_body = False + continue + body += message.get("body", b"") + more_body = message.get("more_body", False) + await send({"type": "http.response.start", "status": 200, "headers": []}) + await send({"type": "http.response.body", "body": str(len(body)).encode()}) + + +async def _drive(middleware, frames: list[dict], headers=None) -> list[dict]: + scope = { + "type": "http", + "method": "POST", + "path": "/", + "headers": headers or [], + } + iterator = iter(frames) + + async def receive() -> dict: + try: + return next(iterator) + except StopIteration: + return {"type": "http.disconnect"} + + sent: list[dict] = [] + + async def send(message: dict) -> None: + sent.append(message) + + await middleware(scope, receive, send) + return sent + + +async def test_rejects_multiframe_body_over_limit_without_content_length(): + """Content-Length가 없어도 실제 누적 바이트가 상한을 넘으면 413.""" + middleware = BodySizeLimitMiddleware(_echo_asgi_app, max_body_bytes=50) + frames = [ + {"type": "http.request", "body": b"x" * 30, "more_body": True}, + {"type": "http.request", "body": b"x" * 30, "more_body": False}, + ] + + sent = await _drive(middleware, frames) + + start = next(m for m in sent if m["type"] == "http.response.start") + assert start["status"] == 413 + + +async def test_allows_multiframe_body_within_limit(): + """다중 프레임이라도 상한 이내면 전체 본문이 앱에 전달된다.""" + middleware = BodySizeLimitMiddleware(_echo_asgi_app, max_body_bytes=50) + frames = [ + {"type": "http.request", "body": b"x" * 10, "more_body": True}, + {"type": "http.request", "body": b"x" * 10, "more_body": False}, + ] + + sent = await _drive(middleware, frames) + + start = next(m for m in sent if m["type"] == "http.response.start") + body = next(m for m in sent if m["type"] == "http.response.body") + assert start["status"] == 200 + assert body["body"] == b"20" +>>>>>>> f57ae1c1b7e725f9df63b20d31f6be5505de2a1c diff --git a/tests/core/test_validation_error_handler.py b/tests/core/test_validation_error_handler.py new file mode 100644 index 0000000..e010712 --- /dev/null +++ b/tests/core/test_validation_error_handler.py @@ -0,0 +1,88 @@ +from fastapi import FastAPI +from fastapi.exceptions import RequestValidationError +from fastapi.testclient import TestClient + +from app.analysis.schemas import SmishingAnalysisRequest +from app.chat.schemas import ChatRequest +from app.core.config import settings +from app.core.exception_handlers import validation_exception_handler + + +def _build_client() -> TestClient: + """실제 요청 스키마 + 실제 검증 핸들러를 배선한 최소 앱. + + AWS 등 서비스 의존 없이 422 응답 형태만 검증한다. + """ + application = FastAPI() + application.add_exception_handler( + RequestValidationError, + validation_exception_handler, + ) + + @application.post("/analyze") + async def analyze(payload: SmishingAnalysisRequest) -> dict: + return {"ok": True} + + @application.post("/chat") + async def chat(payload: ChatRequest) -> dict: + return {"ok": True} + + return TestClient(application) + + +def test_oversized_analyze_text_returns_422_without_leaking_pii(): + client = _build_client() + secret = "01012345678-비밀번호-보이스피싱" + + response = client.post( + "/analyze", + json={"text": secret + "가" * settings.MAX_ANALYSIS_CONTENT_LENGTH}, + ) + + assert response.status_code == 422 + assert secret not in response.text + # 어떤 필드가 왜 거부됐는지는 여전히 알려준다. + assert response.json()["detail"][0]["loc"][-1] == "text" + + +def test_oversized_chat_content_returns_422_without_leaking_pii(): + client = _build_client() + secret = "01012345678-비밀번호" + + response = client.post( + "/chat", + json={ + "messages": [ + { + "role": "user", + "content": secret + + "가" * settings.MAX_CHAT_CONTENT_LENGTH, + } + ] + }, + ) + + assert response.status_code == 422 + assert secret not in response.text + + +def test_oversized_analysis_context_returns_422_without_leaking_pii(): + client = _build_client() + secret = "01012345678-비밀번호" + + response = client.post( + "/chat", + json={ + "analysisContext": { + "riskScore": 90, + "riskLevel": "HIGH", + "category": "FINANCIAL_INSTITUTION", + "explanation": secret + + "가" * settings.MAX_CHAT_CONTEXT_TEXT_LENGTH, + }, + "messages": [{"role": "user", "content": "이거 진짜인가요?"}], + }, + ) + + assert response.status_code == 422 + assert secret not in response.text