diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000..70e2880 --- /dev/null +++ b/.dockerignore @@ -0,0 +1,17 @@ +.git +.github +.vscode + +.venv +__pycache__ +.pytest_cache +*.py[cod] + +.env +.env.* + +tests +requirements-dev.txt + +*.log +.DS_Store \ No newline at end of file diff --git a/.github/workflows/ai-ci.yml b/.github/workflows/ai-ci.yml new file mode 100644 index 0000000..250c77f --- /dev/null +++ b/.github/workflows/ai-ci.yml @@ -0,0 +1,54 @@ +name: AI CI + +on: + pull_request: + push: + branches: + - develop + +permissions: + contents: read + +jobs: + test: + name: Python tests + runs-on: ubuntu-latest + timeout-minutes: 10 + + steps: + - name: Checkout repository + uses: actions/checkout@v4 + with: + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@v5 + with: + python-version: "3.11" + cache: pip + cache-dependency-path: | + requirements.txt + requirements-dev.txt + + - name: Install dependencies + run: pip install -r requirements-dev.txt + + - name: Check dependency compatibility + run: pip check + + - name: Run tests + run: python -m pytest -q + + docker-build: + name: Docker build + runs-on: ubuntu-latest + timeout-minutes: 10 + + steps: + - name: Checkout repository + uses: actions/checkout@v4 + with: + persist-credentials: false + + - name: Build Docker image + run: docker build -t safefam-ai:test . diff --git a/Dockerfile b/Dockerfile index f2aeb96..7a74c8c 100644 --- a/Dockerfile +++ b/Dockerfile @@ -1,17 +1,26 @@ FROM python:3.11-slim -# 작업 디렉토리 설정 WORKDIR /workspace -# 필수 패키지 설치를 위한 레이어 캐싱 +ENV PYTHONDONTWRITEBYTECODE=1 +ENV PYTHONUNBUFFERED=1 + COPY requirements.txt . + RUN pip install --no-cache-dir -r requirements.txt -# 소스코드 전체 복사 -COPY . . +COPY app app + +RUN useradd --create-home --shell /usr/sbin/nologin safefam + +USER safefam -# FastAPI 기본 포트 개방 EXPOSE 8000 -# 로컬 개발 및 컨테이너 구동을 위한 기본 명령어 -CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000"] \ No newline at end of file +HEALTHCHECK \ + --interval=30s \ + --timeout=5s \ + --retries=3 \ + CMD python -c "import urllib.request; urllib.request.urlopen('http://localhost:8000/', timeout=3)" + +CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000"] diff --git a/app/analysis/execution.py b/app/analysis/execution.py new file mode 100644 index 0000000..f307451 --- /dev/null +++ b/app/analysis/execution.py @@ -0,0 +1,84 @@ +from dataclasses import dataclass +from enum import Enum + +from app.analysis.schemas import SmishingAnalysisResponse + + +class AnalysisExecutionStatus(str, Enum): + """분석 실행 상태.""" + + COMPLETED = "COMPLETED" + PARTIAL = "PARTIAL" + FAILED = "FAILED" + + +@dataclass(frozen=True) +class AnalysisExecution: + """분석 결과와 실패 트랙 정보를 담는 불변 데이터.""" + + status: AnalysisExecutionStatus + result: SmishingAnalysisResponse + failed_tracks: tuple[str, ...] + + +def classify_execution( + result: SmishingAnalysisResponse, +) -> AnalysisExecution: + """AI 분석 응답을 실행 상태와 실패 트랙 목록으로 분류한다.""" + if result.status == "ERROR": + return AnalysisExecution( + status=AnalysisExecutionStatus.FAILED, + result=result, + failed_tracks=("PIPELINE",), + ) + + failed_tracks: list[str] = [] + + text_analysis = result.text_analysis or {} + text_result = text_analysis.get("result") or {} + stage1_result = text_analysis.get("stage1_naive_bayes") or {} + + text_failed = ( + text_result.get("grade") == "UNKNOWN" + or text_result.get("error_message") + ) + stage1_failed = ( + stage1_result.get("grade") == "UNKNOWN" + or stage1_result.get("error_message") + ) + + if text_failed: + if stage1_result and not stage1_failed: + failed_tracks.append("TEXT:GEMINI") + else: + failed_tracks.append("TEXT") + elif stage1_failed: + failed_tracks.append("TEXT:NAIVE_BAYES") + + url_analysis = result.url_analysis or {} + failed_providers = url_analysis.get("failed_providers") or [] + url_available = url_analysis.get("available", True) + + if result.url_analysis is not None and not url_available: + failed_tracks.append("URL") + else: + failed_tracks.extend( + f"URL:{provider}" + for provider in failed_providers + ) + + rule_analysis = result.rule_analysis or {} + if rule_analysis.get("error_message"): + failed_tracks.append("RULES") + + status = ( + AnalysisExecutionStatus.PARTIAL + if failed_tracks + else AnalysisExecutionStatus.COMPLETED + ) + + return AnalysisExecution( + status=status, + result=result, + failed_tracks=tuple(failed_tracks), + ) diff --git a/app/analysis/schemas.py b/app/analysis/schemas.py index 03daf4e..cbafd1c 100644 --- a/app/analysis/schemas.py +++ b/app/analysis/schemas.py @@ -39,9 +39,9 @@ class RiskGrade(str, Enum): class ContributionBreakdown(BaseModel): # URL이 있으면 LLM 50% / 규칙 20%, URL이 없으면 하이브리드 URL 트랙(30%)이 LLM/규칙으로 # 재배분되어 LLM 65% / 규칙 35%가 되므로 상한이 두 시나리오 중 더 큰 쪽 기준으로 설정됨 - llm: int = Field(..., description="LLM 문맥 분석 기여 점수 (URL 있음: 0~50, URL 없음: 0~65)", ge=0, le=65) - hybrid_url: int = Field(..., description="하이브리드 URL 보안 엔진 기여 점수 (URL 있을 때만 0~30, 없으면 0)", ge=0, le=30) - rules: int = Field(..., description="로컬 가드 규칙 기반 기여 점수 (URL 있음: 0~20, URL 없음: 0~35)", ge=0, le=35) + llm: int = Field(ge=0, le=100) + hybrid_url: int = Field(ge=0, le=100) + rules: int = Field(ge=0, le=100) class SmishingAnalysisResponse(BaseModel): status: str = Field(..., description="응답 상태 (SUCCESS / ERROR)") diff --git a/app/analysis/scoring.py b/app/analysis/scoring.py index b985bcd..161f628 100644 --- a/app/analysis/scoring.py +++ b/app/analysis/scoring.py @@ -70,49 +70,173 @@ def _combine_text_track_score( + RiskScoringEngine.LLM_WEIGHT * llm_score ) + @staticmethod + def _normalize_weights( + text_weight: float, + url_weight: float, + rules_weight: float, + *, + text_available: bool, + url_available: bool, + rules_available: bool, + ) -> tuple[float, float, float]: + """사용 가능한 분석 트랙들의 가중치 합이 1.0(100%)이 되도록 정규화""" + available_text_weight = ( + text_weight if text_available else 0.0 + ) + available_url_weight = ( + url_weight if url_available else 0.0 + ) + available_rules_weight = ( + rules_weight if rules_available else 0.0 + ) + + total_weight = ( + available_text_weight + + available_url_weight + + available_rules_weight + ) + + if total_weight == 0: + raise ValueError( + "No analysis tracks are available" + ) + + return ( + available_text_weight / total_weight, + available_url_weight / total_weight, + available_rules_weight / total_weight, + ) + @staticmethod def calculate_score( - llm_score: int, # LLM(Gemini)이 반환한 위험도 점수 (0~100). 나이브 베이즈 단독 판정 시엔 그 점수와 동일 - is_url_malicious: bool, # 하이브리드 URL 엔진의 최종 악성 판정 여부 - url_risk_score: float, # 하이브리드 URL 엔진이 계산한 위험도 점수 (0.0~1.0) - rule_score: int, # 로컬 규칙 기반 엔진(금융기관 DB, 금융 키워드, 계좌/카드번호 등)의 위험도 점수 (0~100) - has_url: bool = True, # 문자 본문에 URL이 있었는지 여부 -> 트랙 가중치 재배분 기준 - naive_bayes_score: Optional[int] = None, # 1차 나이브 베이즈 위험도 점수 (미수행 시 None) - llm_available: bool = True, # Gemini 2차 검증이 정상적으로 수행되었는지 여부 - is_confirmed_malicious: bool = False # GSB 블랙리스트 등재 / VT 다수 엔진 합의 / 로컬 도메인 룰 중 하나라도 확정된 경우 - ) -> Tuple[int, RiskGrade, ContributionBreakdown]: - - # 텍스트 문맥 점수와 하이브리드 URL 엔진의 결과값을 결합하여 최종 위험도를 산출 - text_track_score = RiskScoringEngine._combine_text_track_score(naive_bayes_score, llm_score, llm_available) + llm_score: int, + is_url_malicious: bool, + url_risk_score: float, + rule_score: int, + has_url: bool = True, + naive_bayes_score: Optional[int] = None, + llm_available: bool = True, + is_confirmed_malicious: bool = False, + text_available: bool = True, + url_available: bool = True, + rules_available: bool = True, + ) -> Tuple[ + int, + RiskGrade, + ContributionBreakdown, + ]: + """모든 분석 트랙의 점수와 가중치를 종합하여 최종 점수, 위험 등급, 트랙별 기여도 계산""" + text_track_score = ( + RiskScoringEngine + ._combine_text_track_score( + naive_bayes_score=naive_bayes_score, + llm_score=llm_score, + llm_available=llm_available, + ) + ) + # URL 유무에 따른 기본 가중치 설정 if has_url: - text_weight = RiskScoringEngine.TEXT_TRACK_WEIGHT_WITH_URL - rules_weight = RiskScoringEngine.RULES_TRACK_WEIGHT_WITH_URL + text_weight = ( + RiskScoringEngine + .TEXT_TRACK_WEIGHT_WITH_URL + ) + url_weight = ( + RiskScoringEngine.URL_TRACK_WEIGHT + ) + rules_weight = ( + RiskScoringEngine + .RULES_TRACK_WEIGHT_WITH_URL + ) + else: + text_weight = ( + RiskScoringEngine + .TEXT_TRACK_WEIGHT_NO_URL + ) + url_weight = 0.0 + rules_weight = ( + RiskScoringEngine + .RULES_TRACK_WEIGHT_NO_URL + ) + url_available = False - # URL 트랙 기여 점수 (만점 30점) - if is_url_malicious or url_risk_score > 0: - url_contrib = round(url_risk_score * RiskScoringEngine.URL_TRACK_WEIGHT * 100) - url_contrib = min(max(url_contrib, 0), 30) - else: - url_contrib = 0 + # 트랙별 가용성 상태를 반영한 가중치 정규화 + ( + text_weight, + url_weight, + rules_weight, + ) = RiskScoringEngine._normalize_weights( + text_weight=text_weight, + url_weight=url_weight, + rules_weight=rules_weight, + text_available=text_available, + url_available=url_available, + rules_available=rules_available, + ) + + # 텍스트 트랙 점수 기여도 계산 + if text_available: + llm_contrib = round( + text_track_score * text_weight + ) + else: + llm_contrib = 0 + + # URL 트랙 점수 기여도 계산 + if url_available: + normalized_url_score = min( + max(url_risk_score, 0.0), + 1.0, + ) + url_contrib = round( + normalized_url_score + * url_weight + * 100 + ) else: - # 채점할 URL 자체가 없으므로 URL 트랙은 성립하지 않음 -> LLM/규칙 트랙으로 재배분 - text_weight = RiskScoringEngine.TEXT_TRACK_WEIGHT_NO_URL - rules_weight = RiskScoringEngine.RULES_TRACK_WEIGHT_NO_URL url_contrib = 0 - # 1. 텍스트 트랙 기여 점수 — 나이브 베이즈 + Gemini 하이브리드 결합, URL 유무에 따라 50% 또는 65% - llm_contrib = round(text_track_score * text_weight) + # 룰 트랙 점수 기여도 계산 + if rules_available: + normalized_rule_score = min( + max(rule_score, 0), + 100, + ) + rules_contrib = round( + normalized_rule_score + * rules_weight + ) + else: + rules_contrib = 0 - # 2. 로컬 규칙 기반 기여 점수 — 금융기관 DB/키워드/계좌·카드번호 룰 엔진 점수(0~100)를 - # URL 유무에 따라 20% 또는 35% 배점으로 환산 - rules_contrib = round(rule_score * rules_weight) - rules_contrib = min(max(rules_contrib, 0), round(100 * rules_weight)) + contribution_total = ( + llm_contrib + + url_contrib + + rules_contrib + ) - # 최종 위험도 점수 합산 - final_score = min(llm_contrib + url_contrib + rules_contrib, 100) + # 반올림 오차로 인해 총합이 100점을 초과하는 경우 보정 + if contribution_total > 100: + overflow = contribution_total - 100 + + if llm_contrib >= max( + url_contrib, + rules_contrib, + ): + llm_contrib -= overflow + elif url_contrib >= rules_contrib: + url_contrib -= overflow + else: + rules_contrib -= overflow - # 최종 점수 기반 임계치 등급 분기 + final_score = ( + llm_contrib + + url_contrib + + rules_contrib + ) + + # 점수 위험 등급 판정 if final_score >= 70: risk_grade = RiskGrade.HIGH elif final_score >= 40: @@ -120,48 +244,70 @@ def calculate_score( else: risk_grade = RiskGrade.LOW - # 확정 악성 URL 하드 오버라이드: 텍스트 문맥이 아무리 평범해도 아래 중 하나라도 확정되면 - # 가중합으로 희석되지 않고 무조건 HIGH 처리 (GSB만큼 신뢰도가 낮은 VT 단독/소수 탐지는 제외). - # - GSB 블랙리스트 실제 등재 확인 - # - VT 다수 엔진(임계치 이상) 동시 합의 탐지 - # - 로컬 도메인 룰(.ru, testsafebrowsing 등) 매치 - # 등급-점수 표기 일관성을 위해 점수도 HIGH 임계치(70점) 이상으로 끌어올림. + # 확정 악성 신호 감지 시 최소 HIGH 등급(70점) 보장 if is_confirmed_malicious: if risk_grade != RiskGrade.HIGH: logger.warning( - f"[Scoring Engine] 확정 악성 URL 감지 -> HIGH 등급 강제 오버라이드 (원래 점수: {final_score})" + "[Scoring Engine] 확정 악성 신호 " + "감지: HIGH 등급으로 조정 " + "(기존 점수: %s)", + final_score, ) - pre_override_score = final_score - final_score = max(final_score, 70) - risk_grade = RiskGrade.HIGH - # 오버라이드로 늘어난 만큼(final_score - 원래 점수)을 breakdown에도 반영해야 - # contribution_breakdown 합계가 final_score와 어긋나지 않는다. 확정 판정의 - # 근거가 URL 트랙이므로 그쪽에 먼저 배정하고, 각 트랙의 스키마 상한(30/35/65)을 - # 넘으면 규칙 -> LLM 순으로 나머지를 채운다. - score_gap = final_score - pre_override_score + target_score = max(final_score, 70) + score_gap = target_score - final_score + if score_gap > 0: - url_add = min(score_gap, 30 - url_contrib) + url_capacity = max( + round(url_weight * 100) - url_contrib, + 0, + ) + url_add = min( + score_gap, + url_capacity, + ) url_contrib += url_add score_gap -= url_add - rules_add = min(score_gap, 35 - rules_contrib) + rules_capacity = max( + round(rules_weight * 100) - rules_contrib, + 0, + ) + rules_add = min( + score_gap, + rules_capacity, + ) rules_contrib += rules_add score_gap -= rules_add - llm_add = min(score_gap, 65 - llm_contrib) - llm_contrib += llm_add - score_gap -= llm_add + text_capacity = max( + round(text_weight * 100) - llm_contrib, + 0, + ) + text_add = min( + score_gap, + text_capacity, + ) + llm_contrib += text_add + + final_score = target_score + risk_grade = RiskGrade.HIGH logger.info( - f"[Scoring Engine] 통합 연산 완료 -> 최종 점수: {final_score} | 등급: {risk_grade} " - f"(LLM: {llm_contrib}, Hybrid-URL: {url_contrib}, Rules: {rules_contrib})" + "[Scoring Engine] 계산 완료 - " + "점수: %s, 등급: %s " + "(TEXT: %s, URL: %s, RULES: %s)", + final_score, + risk_grade, + llm_contrib, + url_contrib, + rules_contrib, ) breakdown = ContributionBreakdown( llm=llm_contrib, hybrid_url=url_contrib, - rules=rules_contrib + rules=rules_contrib, ) return final_score, risk_grade, breakdown diff --git a/app/analysis/service.py b/app/analysis/service.py index 614430c..fc6a552 100644 --- a/app/analysis/service.py +++ b/app/analysis/service.py @@ -67,36 +67,94 @@ async def url_track(): if url_task: (text_analysis, naive_bayes_score, llm_available), (traced_url, hybrid_res) = await asyncio.gather(text_task, url_task) else: - text_analysis, naive_bayes_score, llm_available = await text_task - traced_url, hybrid_res = None, {"is_malicious": False, "url_risk_score": 0.0, "source": "Pre-Processing-Filter", "error_message": None, "is_gsb_confirmed": False, "is_vt_confirmed": False} + ( + text_analysis, + naive_bayes_score, + llm_available, + ) = await text_task + + traced_url = None + hybrid_res = { + "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, + } - llm_score = text_analysis.get("result", {}).get("risk_score", 0) if isinstance(text_analysis, dict) else 0 + text_result = ( + text_analysis.get("result") or {} + if isinstance(text_analysis, dict) + else {} + ) + llm_score = text_result.get("risk_score", 0) + text_available = ( + naive_bayes_score is not None + or llm_available + ) # 로컬 규칙 기반 트랙: 금융기관 DB 대조 + 금융 키워드 + 계좌/카드번호 패턴 + 도메인 룰(.ru 등) - rule_result = self.rule_analyzer(text, traced_url) - if rule_result["has_malicious_domain_pattern"]: + try: + rule_result = self.rule_analyzer(text, traced_url) + except Exception: + logger.exception("[Analysis Service] 규칙 분석 중 오류 발생") + rule_result = { + "rule_score": 0, + "has_malicious_domain_pattern": False, + "matched_rules": [], + "error_message": "RULE_ANALYSIS_FAILED", + } + + rules_available = not bool(rule_result.get("error_message")) + url_available = ( + has_url + and hybrid_res.get("available", False) + ) + + if rule_result.get("has_malicious_domain_pattern", False): hybrid_res["is_malicious"] = True - hybrid_res["url_risk_score"] = max(hybrid_res["url_risk_score"], 0.75) + hybrid_res["url_risk_score"] = max( + hybrid_res.get("url_risk_score", 0.0), + 0.75, + ) # 확정 악성 판정 소스 3종 중 하나라도 해당하면 문맥 점수와 무관하게 HIGH 강제 오버라이드 대상. # (GSB 블랙리스트 등재 / VT 다수 엔진 합의 / 로컬 도메인 룰 매치 — 신뢰도 낮은 VT 소수 탐지는 제외) is_confirmed_malicious = ( hybrid_res.get("is_gsb_confirmed", False) or hybrid_res.get("is_vt_confirmed", False) - or rule_result["has_malicious_domain_pattern"] + or rule_result.get("has_malicious_domain_pattern", False) + ) + + no_reliable_signal = ( + not text_available + and not url_available + and not rules_available ) + if no_reliable_signal: + raise ValueError( + "No reliable analysis signal is available" + ) # 3중 스코어링 최종 계산 (텍스트 트랙은 나이브 베이즈 + Gemini 하이브리드 결합 점수 사용, # URL 없으면 URL 트랙(30%)이 LLM/규칙 트랙으로 재배분됨) final_score, risk_grade, breakdown = RiskScoringEngine.calculate_score( llm_score=int(llm_score), - is_url_malicious=hybrid_res["is_malicious"], - url_risk_score=hybrid_res["url_risk_score"], - rule_score=rule_result["rule_score"], + is_url_malicious=hybrid_res.get("is_malicious", False), + url_risk_score=hybrid_res.get("url_risk_score", 0.0), + rule_score=rule_result.get("rule_score", 0), has_url=has_url, naive_bayes_score=naive_bayes_score, llm_available=llm_available, - is_confirmed_malicious=is_confirmed_malicious + is_confirmed_malicious=is_confirmed_malicious, + text_available=text_available, + url_available=url_available, + rules_available=rules_available, ) # URL 부재 시 예외 방어 및 스켈레톤 분기벽 구축 @@ -106,10 +164,17 @@ async def url_track(): "is_shortened": original_url != traced_url, "origin_url": traced_url, "original_url": original_url, - "is_url_malicious": hybrid_res["is_malicious"], - "url_risk_score": hybrid_res["url_risk_score"], - "engine_source": hybrid_res["source"], - "error_message": hybrid_res["error_message"] + "is_url_malicious": hybrid_res.get("is_malicious", False), + "url_risk_score": hybrid_res.get("url_risk_score", 0.0), + "engine_source": hybrid_res.get("source", "Hybrid-Engine"), + "available": url_available, + "failed_providers": hybrid_res.get("failed_providers", []), + "pending_providers": hybrid_res.get("pending_providers", []), + "provider_error_codes": hybrid_res.get( + "provider_error_codes", + {}, + ), + "error_message": hybrid_res.get("error_message"), } else: real_url_analysis = None diff --git a/app/analysis/text/gemini_analyzer.py b/app/analysis/text/gemini_analyzer.py index cadbacb..842e274 100644 --- a/app/analysis/text/gemini_analyzer.py +++ b/app/analysis/text/gemini_analyzer.py @@ -150,7 +150,6 @@ async def analyze_text_with_gemini(text: str) -> dict: api_url=API_URL, api_key=GEMINI_API_KEY, payload=payload, - timeout_seconds=10.0, ) candidates = result_json.get("candidates", []) diff --git a/app/analysis/url/analyzer.py b/app/analysis/url/analyzer.py index 3aceb87..02ddddc 100644 --- a/app/analysis/url/analyzer.py +++ b/app/analysis/url/analyzer.py @@ -1,12 +1,21 @@ -import os import logging -import asyncio +import os +from typing import ClassVar + from app.analysis.ports import UrlSecurityProvider -from app.infrastructure.virustotal.client import VirusTotalClient -from app.infrastructure.google_safe_browsing.client import GoogleSafeBrowsingClient +from app.infrastructure.google_safe_browsing.client import ( + GoogleSafeBrowsingClient, +) +from app.infrastructure.virustotal.client import ( + VirusTotalClient, +) logger = logging.getLogger(__name__) -MOCK_ENABLED = os.getenv("MOCK_SECURITY_API", "False").lower() in ("true", "1", "t") + +MOCK_ENABLED = ( + os.getenv("MOCK_SECURITY_API", "False").lower() + in ("true", "1", "t") +) # Google Safe Browsing(1차)과 VirusTotal(2차 백업)을 제어하는 하이브리드 URL 분석 코어 엔진 class HybridUrlAnalyzer: @@ -16,6 +25,10 @@ class HybridUrlAnalyzer: # 다수 엔진이 동시에 일치하면 우연한 오탐일 가능성이 낮아짐) VT_CONFIRMED_ENGINE_THRESHOLD = 5 + AVAILABLE_STATUSES: ClassVar[ + frozenset[str] + ] = frozenset({"safe", "completed"}) + def __init__( self, vt_client: UrlSecurityProvider | None = None, @@ -24,6 +37,16 @@ def __init__( self.vt_client = vt_client or VirusTotalClient() self.gsb_client = gsb_client or GoogleSafeBrowsingClient() + @classmethod + def _is_available(cls, result: dict) -> bool: + """분석 결과 사용 가능한 상태""" + return result.get("status") in cls.AVAILABLE_STATUSES + + @staticmethod + def _is_unavailable(result: dict) -> bool: + """분석 결과 불가 상태""" + return result.get("status") == "unavailable" + async def scan_url(self, traced_url: str) -> dict: # 쉘 환경변수에 따른 MOCK 모드 분기 로직 정상화 if MOCK_ENABLED: @@ -33,66 +56,216 @@ async def scan_url(self, traced_url: str) -> dict: "url_risk_score": 0.85, "source": "Hybrid-Engine (MOCK)", "detected_count": 4, + "available": True, + "failed_providers": [], + "pending_providers": [], + "provider_error_codes": {}, "error_message": None, "is_gsb_confirmed": True, - "is_vt_confirmed": False + "is_vt_confirmed": False, } # PROD 운영 모드 가동 - logger.info("[PROD MODE] 1차 방어선: Google Safe Browsing API 가동") - error_logs = [] + logger.info( + "[PROD MODE] Google Safe Browsing 분석 시작" + ) - try: - gsb_result = await self.gsb_client.scan_url(traced_url) - is_gsb_blocked = gsb_result.get("is_malicious", False) - except Exception as e: - logger.error(f"GSB 통신 실패: {str(e)}") - gsb_result = {"is_malicious": False} - is_gsb_blocked = False - error_logs.append(f"GSB Fail ({str(e)[:15]})") - - vt_result = {"is_malicious": False, "detected_count": 0} - is_vt_confirmed = False - - # GSB 악성 확정 시 VT 생략 (Quota 절약) + failed_providers: list[str] = [] + pending_providers: list[str] = [] + error_messages: list[str] = [] + provider_error_codes: dict[str, str] = {} + + # 1차 분석: GSB + gsb_result = await self._scan_gsb(traced_url) + + gsb_available = self._is_available(gsb_result) + is_gsb_blocked = ( + gsb_available + and gsb_result.get("is_malicious", False) + ) + + if self._is_unavailable(gsb_result): + failed_providers.append("GSB") + + error_code = gsb_result.get( + "error_code", + "UNKNOWN", + ) + provider_error_codes["GSB"] = error_code + error_messages.append( + f"GSB unavailable ({error_code})" + ) + + # GSB에서 악성 URL을 확정한 경우 VirusTotal 호출 X if is_gsb_blocked: logger.info(" GSB 악성 판정으로 VirusTotal 호출 생략 (Quota 절약)") - engine_source = "Hybrid-Engine (GSB)" risk_score = gsb_result.get("raw_score", 0.95) if risk_score > 1.0: risk_score /= 100.0 + + return { + "is_malicious": True, + "url_risk_score": round( + risk_score, + 2, + ), + "source": "Hybrid-Engine (GSB)", + "detected_count": gsb_result.get( + "detected_count", + 1, + ), + "available": True, + "failed_providers": failed_providers, + "pending_providers": pending_providers, + "provider_error_codes": provider_error_codes, + "error_message": ( + " | ".join(error_messages) + if error_messages + else None + ), + "is_gsb_confirmed": True, + "is_vt_confirmed": False, + } + + # 2차 분석: VirusTotal 백업 분석 + logger.info( + "[Hybrid URL] VirusTotal 백업 분석 시작" + ) + + vt_result = await self._scan_virustotal( + traced_url + ) + + vt_status = vt_result.get("status") + vt_available = self._is_available(vt_result) + + if self._is_unavailable(vt_result): + failed_providers.append("VIRUSTOTAL") + + error_code = vt_result.get( + "error_code", + "UNKNOWN", + ) + provider_error_codes[ + "VIRUSTOTAL" + ] = error_code + error_messages.append( + f"VirusTotal unavailable ({error_code})" + ) + + elif vt_status == "scanning": + pending_providers.append("VIRUSTOTAL") + + vt_malicious_count = ( + vt_result.get("detected_count", 0) + if vt_available + else 0 + ) + + # VT 분석 결과를 바탕으로 최종 위험도 점수 재계산 + if vt_available and vt_malicious_count > 0: + base_score = vt_result.get( + "raw_score", + 0.0, + ) + + if base_score > 1.0: + base_score /= 100.0 + + risk_score = max( + base_score, + min( + 0.1 + vt_malicious_count * 0.15, + 0.95, + ), + ) else: - logger.info(" GSB 청정/불확실로 인한 2차 방어선 VirusTotal 백업 가동") - engine_source = "Hybrid-Engine (GSB+VT)" - try: - vt_result = await asyncio.wait_for(self.vt_client.scan_url(traced_url), timeout=4.0) - except Exception as e: - logger.error(f"VirusTotal 통신 실패: {str(e)}") - vt_result = {"is_malicious": False, "detected_count": 0} - error_logs.append(f"VT Fail ({str(e)[:15]})") - - vt_malicious_count = vt_result.get("detected_count", 0) - if vt_malicious_count > 0: - base_score = vt_result.get("raw_score", 0.0) - if base_score > 1.0: - base_score /= 100.0 - risk_score = max(base_score, min(0.1 + (vt_malicious_count * 0.15), 0.95)) - else: - risk_score = 0.0 - - # 다수 백신 엔진이 동시에 악성으로 합의한 경우만 GSB급 확정 신호로 승격 - is_vt_confirmed = vt_malicious_count >= self.VT_CONFIRMED_ENGINE_THRESHOLD - - is_final_malicious = is_gsb_blocked or vt_result.get("is_malicious", False) or vt_result.get("detected_count", 0) >= 1 - combined_error = " | ".join(error_logs) if error_logs else None + risk_score = 0.0 + + is_vt_malicious = ( + vt_available + and vt_result.get( + "is_malicious", + False, + ) + ) + + is_vt_confirmed = ( + vt_available + and vt_malicious_count + >= self.VT_CONFIRMED_ENGINE_THRESHOLD + ) + + is_final_malicious = ( + is_gsb_blocked + or is_vt_malicious + ) + + # 두 제공자 중 하나라도 정상 결과를 반환하면 URL트랙은 사용 가능으로 판정 + available = gsb_available or vt_available + + if not available and not error_messages: + error_messages.append( + "No URL provider returned a completed result" + ) return { "is_malicious": is_final_malicious, "url_risk_score": round(risk_score, 2), - "source": engine_source, - "error_message": combined_error, - # Google Safe Browsing 블랙리스트에 실제로 등재되어 확인된 경우만 True. + "source": "Hybrid-Engine (GSB+VT)", + "detected_count": vt_malicious_count, + "available": available, + "failed_providers": failed_providers, + "pending_providers": pending_providers, + "provider_error_codes": provider_error_codes, + "error_message": ( + " | ".join(error_messages) + if error_messages + else None + ), "is_gsb_confirmed": is_gsb_blocked, - # VT 탐지 엔진 수가 임계치 이상인 "다수 합의" 케이스만 True (스코어링 엔진의 확정 악성 오버라이드 트리거용) - "is_vt_confirmed": is_vt_confirmed + "is_vt_confirmed": is_vt_confirmed, } + + async def _scan_gsb( + self, + traced_url: str, + ) -> dict: + """GSB 클라이언트를 호출하고 예외 발생시 예외 응답 반환""" + try: + return await self.gsb_client.scan_url( + traced_url + ) + + except Exception: + logger.exception( + "[Hybrid URL] GSB 호출 중 예외 발생" + ) + return { + "is_malicious": False, + "raw_score": 0.0, + "detected_count": 0, + "status": "unavailable", + "error_code": "UNEXPECTED_ERROR", + } + async def _scan_virustotal( + self, + traced_url: str, + ) -> dict: + """VirusTotal 클라이언트를 호출하고 실패 결과로 변환""" + try: + return await self.vt_client.scan_url( + traced_url + ) + + except Exception: + logger.exception( + "[Hybrid URL] VirusTotal 호출 중 예외 발생" + ) + return { + "is_malicious": False, + "raw_score": 0.0, + "detected_count": 0, + "status": "unavailable", + "error_code": "UNEXPECTED_ERROR", + } diff --git a/app/analysis/url/tracker.py b/app/analysis/url/tracker.py index 6c0f057..6671e30 100644 --- a/app/analysis/url/tracker.py +++ b/app/analysis/url/tracker.py @@ -3,20 +3,19 @@ import asyncio import ipaddress import logging -from typing import List, Optional -from urllib.parse import urljoin, urlparse import httpx import httpcore +from typing import List, Optional +from urllib.parse import urljoin, urlparse +from app.core.config import settings +from app.infrastructure.http_retry import request_with_retry + logger = logging.getLogger(__name__) # URL 정규표현식 패턴 URL_PATTERN = re.compile(r'https?://[^\s\'"<>]+') -# DNS 조회 자체가 멎어버리는 것(응답 없는 리졸버 등)을 막기 위한 상한 -DNS_RESOLVE_TIMEOUT = 3.0 - - def _is_blocked_ip(ip: "ipaddress.IPv4Address | ipaddress.IPv6Address") -> bool: """사설/루프백/링크로컬 등 외부에 공개되지 않은 주소인지 판별.""" return ( @@ -30,7 +29,8 @@ def _is_blocked_ip(ip: "ipaddress.IPv4Address | ipaddress.IPv6Address") -> bool: async def _is_public_host( - hostname: Optional[str], dns_timeout: float = DNS_RESOLVE_TIMEOUT + hostname: Optional[str], + dns_timeout: float | None = None, ) -> Optional[str]: """ SSRF 방어: 호스트가 실제로 가리키는 IP를 DNS로 확인해서 내부망/사설 대역이면 차단. @@ -41,6 +41,9 @@ async def _is_public_host( 응답이 바뀌는 DNS 리바인딩을 막을 수 있다. 그래서 bool이 아니라 검증에 사용한 IP 문자열(고정할 주소)을 반환하고, 차단 시 None을 반환한다. """ + if dns_timeout is None: + dns_timeout = settings.URL_TRACE_TIMEOUT_SECONDS + if not hostname: return None @@ -145,7 +148,13 @@ def extract_urls(text: str) -> List[str]: # 단축 URL의 리다이렉트를 추적하고 최종 주소를 반환 -async def trace_url(url: str, max_redirects: int = 5, timeout: float = 3.0) -> str: +async def trace_url( + url: str, + max_redirects: int = 5, + timeout: float | None = None, +) -> str: + if timeout is None: + timeout = settings.URL_TRACE_TIMEOUT_SECONDS current_url = url @@ -170,11 +179,27 @@ async def trace_url(url: str, max_redirects: int = 5, timeout: float = 3.0) -> s "User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36" } - response = await client.head(current_url, headers=headers, timeout=timeout) + response = await request_with_retry( + lambda current_url=current_url, headers=headers: client.head( + current_url, + headers=headers, + timeout=timeout, + ), + max_retries=settings.EXTERNAL_API_MAX_RETRIES, + operation_name="URL trace HEAD", + ) # HEAD를 차단하거나 거부하는 서버(400, 404, 405)에 대응하기 위한 GET 폴백 if response.status_code in [400, 404, 405]: - response = await client.get(current_url, headers=headers, timeout=timeout) + response = await request_with_retry( + lambda current_url=current_url, headers=headers: client.get( + current_url, + headers=headers, + timeout=timeout, + ), + max_retries=settings.EXTERNAL_API_MAX_RETRIES, + operation_name="URL trace GET", + ) # HTTP Redirection 상태 코드 판별 (3xx) if response.is_redirect or response.status_code in [301, 302, 303, 307, 308]: @@ -203,4 +228,4 @@ async def trace_url(url: str, max_redirects: int = 5, timeout: float = 3.0) -> s else: logger.warning(f"최대 리다이렉트 횟수({max_redirects}회)를 초과했습니다. 루프 위험 감지.") - return current_url \ No newline at end of file + return current_url diff --git a/app/core/config.py b/app/core/config.py index 97e704d..1ccd349 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -9,6 +9,28 @@ class Settings(BaseSettings): VIRUSTOTAL_API_KEY: str | None = None GOOGLE_SAFE_BROWSING_API_KEY: str | None = None + GEMINI_TIMEOUT_SECONDS: float = Field( + default=10.0, + gt=0, + ) + GSB_TIMEOUT_SECONDS: float = Field( + default=5.0, + gt=0, + ) + VIRUSTOTAL_TIMEOUT_SECONDS: float = Field( + default=5.0, + gt=0, + ) + URL_TRACE_TIMEOUT_SECONDS: float = Field( + default=3.0, + gt=0, + ) + EXTERNAL_API_MAX_RETRIES: int = Field( + default=1, + ge=0, + le=3, + ) + RABBITMQ_URL: str = ( "amqp://safefam:safefam-local@localhost:5672/" @@ -23,6 +45,35 @@ class Settings(BaseSettings): RABBITMQ_PREFETCH_COUNT: int = Field(default=1, ge=1) RABBITMQ_CONSUMER_ENABLED: bool = True + RABBITMQ_ANALYSIS_COMPLETED_ROUTING_KEY: str = ( + "analysis.completed.v1" + ) + RABBITMQ_ANALYSIS_PARTIAL_ROUTING_KEY: str = ( + "analysis.partial.v1" + ) + RABBITMQ_ANALYSIS_FAILED_ROUTING_KEY: str = ( + "analysis.failed.v1" + ) + + RABBITMQ_ANALYSIS_DLQ: str = ( + "safefam.analysis.requested.dlq" + ) + RABBITMQ_ANALYSIS_DLQ_ROUTING_KEY: str = ( + "analysis.requested.dead.v1" + ) + + RABBITMQ_PUBLISH_TIMEOUT_SECONDS: float = Field( + default=5.0, + gt=0, + ) + RABBITMQ_SHUTDOWN_TIMEOUT_SECONDS: float = Field( + default=30.0, + gt=0, + ) + RABBITMQ_REQUEUE_BACKOFF_SECONDS: float = Field( + default=1.0, + ge=0, + ) model_config = SettingsConfigDict( env_file=".env", diff --git a/app/infrastructure/errors.py b/app/infrastructure/errors.py new file mode 100644 index 0000000..488e313 --- /dev/null +++ b/app/infrastructure/errors.py @@ -0,0 +1,18 @@ +class ProcessingError(RuntimeError): + """메시지 처리 오류의 공통 기반 클래스""" + + def __init__( + self, + message: str, + failure_code: str, + ) -> None: + super().__init__(message) + self.failure_code = failure_code + + +class RetryableProcessingError(ProcessingError): + """일시적인 문제로 재시도할 수 있는 오류""" + + +class NonRetryableProcessingError(ProcessingError): + """재시도해도 성공할 가능성이 없는 오류""" \ No newline at end of file diff --git a/app/infrastructure/gemini/client.py b/app/infrastructure/gemini/client.py index 92e815c..31886e2 100644 --- a/app/infrastructure/gemini/client.py +++ b/app/infrastructure/gemini/client.py @@ -4,6 +4,9 @@ import httpx +from app.core.config import settings +from app.infrastructure.http_retry import request_with_retry + class GeminiClient: """Send raw generation requests without applying feature-specific policy.""" @@ -14,18 +17,30 @@ async def generate( api_url: str, api_key: str, payload: dict[str, Any], - timeout_seconds: float = 10.0, + timeout_seconds: float | None = None, ) -> dict[str, Any]: headers = { "x-goog-api-key": api_key, "content-type": "application/json", } + + timeout = ( + timeout_seconds + if timeout_seconds is not None + else settings.GEMINI_TIMEOUT_SECONDS + ) + async with httpx.AsyncClient() as client: - response = await client.post( - api_url, - json=payload, - headers=headers, - timeout=timeout_seconds, + response = await request_with_retry( + lambda: client.post( + api_url, + json=payload, + headers=headers, + timeout=timeout, + ), + max_retries=settings.EXTERNAL_API_MAX_RETRIES, + operation_name="Gemini", ) + response.raise_for_status() - return response.json() + return response.json() \ No newline at end of file diff --git a/app/infrastructure/google_safe_browsing/client.py b/app/infrastructure/google_safe_browsing/client.py index 8370fe0..02bda8c 100644 --- a/app/infrastructure/google_safe_browsing/client.py +++ b/app/infrastructure/google_safe_browsing/client.py @@ -1,70 +1,153 @@ import logging import httpx from app.core.config import settings +from app.infrastructure.http_retry import request_with_retry logger = logging.getLogger(__name__) class GoogleSafeBrowsingClient: + """Google Safe Browsing API를 사용하여 URL의 악성 여부를 검사""" + def __init__(self): self.api_key = settings.GOOGLE_SAFE_BROWSING_API_KEY - self.api_url = f"https://safebrowsing.googleapis.com/v4/threatMatches:find?key={self.api_key}" + self.api_url = ( + "https://safebrowsing.googleapis.com/v4/" + "threatMatches:find" + ) - # Google Safe Browsing API를 사용하여 URL의 실시간 악성 블랙리스트 등재 여부를 검사 - async def scan_url(self, url: str) -> dict: - - # 공통 인터페이스 리턴 규격 스켈레톤 선언 - default_result = {"is_malicious": False, "raw_score": 0.0, "detected_count": 0, "status": "safe"} + @staticmethod + def _safe_result() -> dict: + """GSB가 실제로 URL을 조회한 뒤 안전하다고 판단한 결과""" + return { + # 실제 안전 판정 + "is_malicious": False, + "raw_score": 0.0, + "detected_count": 0, + "status": "safe", + "error_code": None, + } + + @staticmethod + def _unavailable_result(error_code: str) -> dict: + """API 장애로 URL의 안전 여부를 판단하지 못한 결과""" + # API 장애로 판정 불가 + return { + "is_malicious": False, + "raw_score": 0.0, + "detected_count": 0, + "status": "unavailable", + "error_code": error_code, + } + async def scan_url(self, url: str) -> dict: + """입력받은 URL을 GSB API로 검사 후 분석 결과 반환""" if not self.api_key: - logger.warning("[Google Safe Browsing] API Key가 누락되었습니다. 빈 분석 결과를 반환합니다.") - return default_result + logger.warning( + "[Google Safe Browsing] API Key가 누락되어 " + "URL을 분석할 수 없습니다." + ) + return self._unavailable_result( + "MISSING_API_KEY" + ) + # GSB API 요청 페이로드 구성 payload = { "client": { "clientId": "safefam-ai-backend", - "clientVersion": "1.0.0" + "clientVersion": "1.0.0", }, "threatInfo": { "threatTypes": [ - "MALWARE", - "SOCIAL_ENGINEERING", - "UNWANTED_SOFTWARE", - "POTENTIALLY_HARMFUL_APPLICATION" + "MALWARE", + "SOCIAL_ENGINEERING", + "UNWANTED_SOFTWARE", + "POTENTIALLY_HARMFUL_APPLICATION", ], "platformTypes": ["ANY_PLATFORM"], "threatEntryTypes": ["URL"], "threatEntries": [ - {"url": url} - ] - } + {"url": url}, + ], + }, } async with httpx.AsyncClient() as client: try: - response = await client.post(self.api_url, json=payload, timeout=5.0) + response = await request_with_retry( + lambda: client.post( + self.api_url, + params={"key": self.api_key}, + json=payload, + timeout=settings.GSB_TIMEOUT_SECONDS, + ), + max_retries=( + settings.EXTERNAL_API_MAX_RETRIES + ), + operation_name="Google Safe Browsing", + ) response.raise_for_status() result = response.json() + matches = result.get("matches", []) - # 응답 데이터에 'matches' 필드가 있으면 확실한 악성 사이트 상태 - if "matches" in result and len(result["matches"]) > 0: - logger.warning(f"[Google Safe Browsing] 악성 URL 감지됨: {url}") + # 매칭되는 위험 요소가 있는 경우 악성 URL로 처리 + if matches: + logger.warning( + "[Google Safe Browsing] 악성 URL 감지됨: %s", + url, + ) return { "is_malicious": True, - "raw_score": 0.95, - "detected_count": len(result["matches"]), - "status": "completed" + "raw_score": 0.95, + "detected_count": len(matches), + "status": "completed", + "error_code": None, } - - logger.info(f"[Google Safe Browsing] 안전한 URL: {url}") - return default_result - except httpx.HTTPStatusError as e: - logger.error(f"[Google Safe Browsing] API 에러 ({e.response.status_code}): {str(e)}") - return default_result + logger.info( + "[Google Safe Browsing] 안전한 URL: %s", + url, + ) + return self._safe_result() + + except httpx.HTTPStatusError as exc: + status_code = exc.response.status_code + + logger.error( + "[Google Safe Browsing] API 에러 (%s)", + status_code, + ) + + if status_code == 429: + return self._unavailable_result( + "RATE_LIMITED" + ) + + return self._unavailable_result( + f"HTTP_{status_code}" + ) + except httpx.TimeoutException: - logger.error("[Google Safe Browsing] API 요청 타임아웃 발생") - return default_result - except Exception as e: - logger.error(f"[Google Safe Browsing] 연동 중 비정상 에러 발생: {str(e)}") - return default_result + logger.error( + "[Google Safe Browsing] " + "API 요청 타임아웃 발생" + ) + return self._unavailable_result("TIMEOUT") + + except httpx.RequestError as exc: + logger.error( + "[Google Safe Browsing] 네트워크 오류: %s", + type(exc).__name__, + ) + return self._unavailable_result( + "NETWORK_ERROR" + ) + + except Exception: + logger.exception( + "[Google Safe Browsing] " + "연동 중 비정상 오류 발생" + ) + return self._unavailable_result( + "UNEXPECTED_ERROR" + ) diff --git a/app/infrastructure/http_retry.py b/app/infrastructure/http_retry.py new file mode 100644 index 0000000..bd16a76 --- /dev/null +++ b/app/infrastructure/http_retry.py @@ -0,0 +1,123 @@ +import asyncio +import logging +import random +from collections.abc import Awaitable, Callable +from datetime import datetime, timezone +from email.utils import parsedate_to_datetime + +import httpx + +logger = logging.getLogger(__name__) + +# 재시도 대상이 되는 httpx 네트워크 예외 목록 +_RETRYABLE_REQUEST_ERRORS = ( + httpx.TimeoutException, + httpx.ConnectError, + httpx.ReadError, + httpx.WriteError, + httpx.PoolTimeout, + httpx.RemoteProtocolError, +) + + +def is_retryable_status(status_code: int) -> bool: + """HTTP 상태 코드가 재시도 가능한 대상인지 검증""" + return ( + status_code == 429 + or 500 <= status_code <= 599 + ) + + +async def _sleep_before_retry( + attempt: int, + retry_after: str | None = None, +) -> None: + """지수 backoff와 jitter를 적용하고 Retry-After를 우선한다.""" + delay: float | None = None + + if retry_after is not None: + try: + delay = max(float(retry_after), 0.0) + except ValueError: + try: + retry_at = parsedate_to_datetime( + retry_after + ) + if retry_at.tzinfo is None: + retry_at = retry_at.replace( + tzinfo=timezone.utc + ) + delay = max( + ( + retry_at + - datetime.now(timezone.utc) + ).total_seconds(), + 0.0, + ) + except (TypeError, ValueError, OverflowError): + delay = None + + if delay is None: + base_delay = min(0.25 * (2 ** attempt), 5.0) + delay = base_delay + random.uniform( + 0.0, + base_delay * 0.1, + ) + + await asyncio.sleep(min(delay, 30.0)) + + +async def request_with_retry( + operation: Callable[[], Awaitable[httpx.Response]], + *, + max_retries: int, + operation_name: str, +) -> httpx.Response: + """비동기 HTTP 요청 중 네트워크 예외 또는 재시도 대상 상태 코드가 발생하면 지정된 횟수만큼 재시도 수행""" + for attempt in range(max_retries + 1): + try: + response = await operation() + + except _RETRYABLE_REQUEST_ERRORS as exc: + # 최대 재시도 횟수 도달 시 예외 전차 + if attempt >= max_retries: + raise + + logger.warning( + "%s 호출 실패로 재시도합니다. attempt=%s/%s error=%s", + operation_name, + attempt + 1, + max_retries, + type(exc).__name__, + ) + await _sleep_before_retry(attempt) + continue + + # 기존 커넥션 정리 후 재시도 + if ( + is_retryable_status(response.status_code) + and attempt < max_retries + ): + logger.warning( + "%s 응답이 재시도 대상입니다. " + "attempt=%s/%s status=%s", + operation_name, + attempt + 1, + max_retries, + response.status_code, + ) + retry_after = response.headers.get( + "Retry-After" + ) + await response.aclose() + await _sleep_before_retry( + attempt, + retry_after=retry_after, + ) + continue + + return response + + raise RuntimeError( + f"{operation_name} retry loop terminated unexpectedly" + ) diff --git a/app/infrastructure/rabbitmq/connection.py b/app/infrastructure/rabbitmq/connection.py index 7f3bbc7..5316283 100644 --- a/app/infrastructure/rabbitmq/connection.py +++ b/app/infrastructure/rabbitmq/connection.py @@ -29,6 +29,7 @@ def __init__( self.channel: AbstractRobustChannel | None = None self.exchange: AbstractRobustExchange | None = None self.request_queue: AbstractRobustQueue | None = None + self.dead_letter_queue: AbstractRobustQueue | None = None async def connect(self) -> None: """RabbitMQ 연결 및 Exchange, Queue, Binding 초기화""" @@ -45,7 +46,10 @@ async def connect(self) -> None: self.settings.RABBITMQ_URL ) - self.channel = await self.connection.channel() + self.channel = await self.connection.channel( + publisher_confirms=True, + on_return_raises=True, + ) await self.channel.set_qos( prefetch_count=( @@ -72,6 +76,21 @@ async def connect(self) -> None: ), ) + self.dead_letter_queue = ( + await self.channel.declare_queue( + self.settings.RABBITMQ_ANALYSIS_DLQ, + durable=True, + ) + ) + + await self.dead_letter_queue.bind( + self.exchange, + routing_key=( + self.settings + .RABBITMQ_ANALYSIS_DLQ_ROUTING_KEY + ), + ) + logger.info( "RabbitMQ analysis request topology initialized. " "exchange=%s queue=%s routing_key=%s prefetch=%s", @@ -90,6 +109,15 @@ def get_request_queue(self) -> AbstractRobustQueue: return self.request_queue + def get_exchange(self) -> AbstractRobustExchange: + """초기화된 분석 이벤트 Exchange를 반환""" + if self.exchange is None: + raise RabbitMQNotConnectedError( + "RabbitMQ exchange is not initialized." + ) + + return self.exchange + async def close(self) -> None: if self.connection is None: return @@ -101,4 +129,5 @@ async def close(self) -> None: self.connection = None self.channel = None self.exchange = None - self.request_queue = None \ No newline at end of file + self.request_queue = None + self.dead_letter_queue = None \ No newline at end of file diff --git a/app/infrastructure/rabbitmq/consumer.py b/app/infrastructure/rabbitmq/consumer.py index 2589728..d5c04e4 100644 --- a/app/infrastructure/rabbitmq/consumer.py +++ b/app/infrastructure/rabbitmq/consumer.py @@ -1,4 +1,8 @@ +import asyncio +import json import logging +from json import JSONDecodeError +from app.core.config import settings from aio_pika.abc import ( AbstractIncomingMessage, @@ -6,31 +10,65 @@ ) from pydantic import ValidationError +from app.analysis.execution import classify_execution +from app.infrastructure.errors import ( + NonRetryableProcessingError, +) +from app.infrastructure.rabbitmq.dead_letter import ( + DeadLetterPublisher, +) from app.infrastructure.rabbitmq.handler import ( AnalysisRequestHandler, ) +from app.infrastructure.rabbitmq.publisher import ( + AnalysisResultPublisher, +) +from app.infrastructure.rabbitmq.result_factory import ( + AnalysisResultEventFactory, +) from app.infrastructure.rabbitmq.schemas import ( AnalysisRequestedEvent, ) logger = logging.getLogger(__name__) + class AnalysisRequestConsumer: - """Spring의 분석 요청 이벤트를 수신하고 처리를 제어""" + """Spring 분석 요청을 처리하고 결과 이벤트 발행""" + MAX_RETRY_ATTEMPTS = 1 def __init__( self, request_queue: AbstractRobustQueue, handler: AnalysisRequestHandler, + result_publisher: AnalysisResultPublisher, + result_factory: AnalysisResultEventFactory, + dead_letter_publisher: DeadLetterPublisher, + shutdown_timeout_seconds: float = ( + settings.RABBITMQ_SHUTDOWN_TIMEOUT_SECONDS + ), + requeue_backoff_seconds: float = ( + settings.RABBITMQ_REQUEUE_BACKOFF_SECONDS + ), ) -> None: self.request_queue = request_queue self.handler = handler + self.result_publisher = result_publisher + self.result_factory = result_factory + self.dead_letter_publisher = dead_letter_publisher + self.consumer_tag: str | None = None - self.retry_attempts: dict[str, int] = {} + self.shutdown_timeout_seconds = ( + shutdown_timeout_seconds + ) + self.requeue_backoff_seconds = ( + requeue_backoff_seconds + ) + self.in_flight_tasks: set[asyncio.Task] = set() async def start(self) -> None: - """Consumer 수신 시작""" + """분석 요청 Consumer 시작""" if self.consumer_tag is not None: logger.info( "Analysis request consumer is already running. " @@ -39,7 +77,6 @@ async def start(self) -> None: ) return - # Consumer 시작 self.consumer_tag = await self.request_queue.consume( self._on_message, no_ack=False, @@ -52,100 +89,286 @@ async def start(self) -> None: ) async def stop(self) -> None: - """Consumer 수신 중단""" - if self.consumer_tag is None: + """신규 수신을 중단하고 처리 중인 메시지 대기""" + if self.consumer_tag is not None: + consumer_tag = self.consumer_tag + self.consumer_tag = None + + await self.request_queue.cancel( + consumer_tag + ) + + logger.info( + "Analysis request consumer subscription stopped. " + "consumer_tag=%s", + consumer_tag, + ) + + current_task = asyncio.current_task() + + pending_tasks = { + task + for task in self.in_flight_tasks + if task is not current_task + and not task.done() + } + + if not pending_tasks: + logger.info( + "No in-flight analysis requests remain." + ) return - consumer_tag = self.consumer_tag + logger.info( + "Waiting for in-flight analysis requests. " + "count=%s timeout_seconds=%s", + len(pending_tasks), + self.shutdown_timeout_seconds, + ) + + done, pending = await asyncio.wait( + pending_tasks, + timeout=self.shutdown_timeout_seconds, + ) - await self.request_queue.cancel(consumer_tag) - self.consumer_tag = None + if pending: + logger.warning( + "Graceful shutdown timeout reached. " + "completed=%s pending=%s", + len(done), + len(pending), + ) + return logger.info( - "Analysis request consumer stopped. " - "consumer_tag=%s", - consumer_tag, + "All in-flight analysis requests completed. " + "completed=%s", + len(done), ) async def _on_message( self, message: AbstractIncomingMessage, ) -> None: - """메시지를 검증하고 분석 핸들러 호출""" + """현재 메시지 처리 Task 기록""" + task = asyncio.current_task() + + if task is not None: + self.in_flight_tasks.add(task) + try: - event = AnalysisRequestedEvent.model_validate_json( - message.body - ) - except ValidationError: - logger.exception( - "Rejecting invalid analysis request event. " - "message_id=%s", - message.message_id, - ) + await self._process_message(message) + finally: + if task is not None: + self.in_flight_tasks.discard(task) - await message.reject(requeue=False) + async def _process_message( + self, + message: AbstractIncomingMessage, + ) -> None: + """요청 처리 및 결과 발행 성공 후 ACK""" + try: + event = self._parse_event(message.body) + except NonRetryableProcessingError as exception: + await self._route_to_dead_letter( + message=message, + event=None, + failure_code=exception.failure_code, + ) return try: - await self.handler.handle(event) - except Exception: - await self._handle_processing_failure( + result = await self.handler.handle(event) + + execution = classify_execution(result) + + result_event = self.result_factory.create( + request=event, + execution=execution, + ) + + await self.result_publisher.publish( + result_event + ) + except NonRetryableProcessingError as exception: + await self._route_to_dead_letter( + message=message, + event=event, + failure_code=exception.failure_code, + ) + return + except Exception as exception: + await self._handle_retryable_failure( message=message, event=event, + exception=exception, ) return - self.retry_attempts.pop(str(event.eventId), None) await message.ack() logger.info( - "Analysis request acknowledged. " + "Analysis request acknowledged after result publication. " "message_id=%s event_id=%s " - "analysis_id=%s trace_id=%s", + "analysis_id=%s trace_id=%s " + "result_event_id=%s result_event_type=%s", message.message_id, event.eventId, event.analysisId, event.traceId, + result_event.eventId, + result_event.eventType.value, ) - async def _handle_processing_failure( + def _parse_event( + self, + body: bytes, + ) -> AnalysisRequestedEvent: + """JSON과 이벤트 계약을 단계적으로 검증""" + try: + raw_event = json.loads(body) + except (JSONDecodeError, UnicodeDecodeError) as exception: + raise NonRetryableProcessingError( + message="Invalid JSON analysis request", + failure_code="INVALID_JSON", + ) from exception + + if not isinstance(raw_event, dict): + raise NonRetryableProcessingError( + message="Analysis event must be an object", + failure_code="INVALID_EVENT_SCHEMA", + ) + + if raw_event.get("schemaVersion") != "1.0": + raise NonRetryableProcessingError( + message="Unsupported schema version", + failure_code="UNSUPPORTED_SCHEMA_VERSION", + ) + + try: + return AnalysisRequestedEvent.model_validate( + raw_event + ) + except ValidationError as exception: + raise NonRetryableProcessingError( + message="Invalid analysis event schema", + failure_code="INVALID_EVENT_SCHEMA", + ) from exception + + async def _handle_retryable_failure( self, message: AbstractIncomingMessage, event: AnalysisRequestedEvent, + exception: Exception, ) -> None: - """분석 실패 시 1회 재시도 후 최종 Reject 처리""" - event_id = str(event.eventId) - retry_attempt = self.retry_attempts.get(event_id, 0) + """일시적 오류를 1회 재시도한 뒤 DLQ로 격리""" + retry_attempt = self._delivery_attempt(message) if retry_attempt >= self.MAX_RETRY_ATTEMPTS: - self.retry_attempts.pop(event_id, None) - - logger.exception( + logger.error( "Analysis request failed after retry. " - "Rejecting the message. " + "Routing sanitized event to DLQ. " "message_id=%s event_id=%s " - "analysis_id=%s trace_id=%s", + "analysis_id=%s trace_id=%s " + "error_type=%s", message.message_id, event.eventId, event.analysisId, event.traceId, + exception.__class__.__name__, ) - await message.reject(requeue=False) + await self._route_to_dead_letter( + message=message, + event=event, + failure_code=( + "PROCESSING_RETRIES_EXHAUSTED" + ), + ) return - self.retry_attempts[event_id] = retry_attempt + 1 - - logger.exception( + logger.warning( "Analysis request processing failed. " - "Requeueing the message for one retry. " + "Requeueing for retry. " "message_id=%s event_id=%s " - "analysis_id=%s trace_id=%s retry_attempt=%s", + "analysis_id=%s trace_id=%s " + "retry_attempt=%s error_type=%s", message.message_id, event.eventId, event.analysisId, event.traceId, retry_attempt + 1, + exception.__class__.__name__, ) + await asyncio.sleep( + self.requeue_backoff_seconds + ) await message.nack(requeue=True) + + async def _route_to_dead_letter( + self, + *, + message: AbstractIncomingMessage, + event: AnalysisRequestedEvent | None, + failure_code: str, + ) -> None: + """정제된 DLQ 이벤트 발행 성공 후 원본을 ACK""" + try: + await self.dead_letter_publisher.publish( + original_message_id=message.message_id, + failure_code=failure_code, + request_event=event, + ) + except Exception as exception: + logger.error( + "Failed to publish sanitized DLQ event. " + "message_id=%s failure_code=%s " + "error_type=%s", + message.message_id, + failure_code, + exception.__class__.__name__, + ) + + await asyncio.sleep( + self.requeue_backoff_seconds + ) + await message.nack(requeue=True) + return + + await message.ack() + + logger.warning( + "Original analysis request acknowledged " + "after sanitized DLQ publication. " + "message_id=%s analysis_id=%s " + "trace_id=%s failure_code=%s", + message.message_id, + event.analysisId if event else None, + event.traceId if event else None, + failure_code, + ) + + @staticmethod + def _delivery_attempt( + message: AbstractIncomingMessage, + ) -> int: + """Broker가 보존한 전달 상태에서 현재 재시도 횟수를 계산한다.""" + headers = message.headers or {} + + delivery_count = headers.get( + "x-delivery-count" + ) + if delivery_count is not None: + return int(delivery_count) + + x_death = headers.get("x-death") or [] + death_counts = [ + int(entry.get("count", 0)) + for entry in x_death + if isinstance(entry, dict) + ] + if death_counts: + return max(death_counts) + + return 1 if message.redelivered else 0 diff --git a/app/infrastructure/rabbitmq/dead_letter.py b/app/infrastructure/rabbitmq/dead_letter.py new file mode 100644 index 0000000..0c4cd05 --- /dev/null +++ b/app/infrastructure/rabbitmq/dead_letter.py @@ -0,0 +1,126 @@ +import asyncio +import logging +from datetime import datetime, timezone +from uuid import uuid4 + +from aio_pika import DeliveryMode, Message +from aio_pika.abc import AbstractRobustExchange + +from app.core.config import Settings, settings +from app.infrastructure.errors import ( + RetryableProcessingError, +) +from app.infrastructure.rabbitmq.schemas import ( + AnalysisRequestedEvent, + DeadLetterEvent, +) + +logger = logging.getLogger(__name__) + + +class DeadLetterPublisher: + """개인정보가 제거된 실패 이벤트를 DLQ로 발행""" + + def __init__( + self, + exchange: AbstractRobustExchange, + app_settings: Settings = settings, + ) -> None: + self.exchange = exchange + self.settings = app_settings + + async def publish( + self, + *, + original_message_id: str | None, + failure_code: str, + request_event: AnalysisRequestedEvent | None, + ) -> DeadLetterEvent: + """실패 사유와 기본 식별자 정보만 담은 DLQ 이벤트 생성 후 RabbitMQ로 발행""" + + # 요청 이벤트에서 비식별 추적 정보만 추출하여 DLQ 이벤트 생성 + dead_letter_event = DeadLetterEvent( + schemaVersion="1.0", + eventId=uuid4(), + originalMessageId=original_message_id, + analysisId=( + request_event.analysisId + if request_event is not None + else None + ), + traceId=( + request_event.traceId + if request_event is not None + else None + ), + failureCode=failure_code, + failedAt=datetime.now(timezone.utc), + ) + + # DLQ 발행용 aio_pika 메시지 객체 생성 + message = Message( + body=dead_letter_event.model_dump_json( + by_alias=True + ).encode("utf-8"), + content_type="application/json", + delivery_mode=DeliveryMode.PERSISTENT, + message_id=str(dead_letter_event.eventId), + correlation_id=( + str(dead_letter_event.traceId) + if dead_letter_event.traceId is not None + else None + ), + headers={ + "schemaVersion": ( + dead_letter_event.schemaVersion + ), + "failureCode": failure_code, + "sanitized": True, + }, + ) + + try: + # DLQ 라우팅 키로 메시지 발행 + await asyncio.wait_for( + self.exchange.publish( + message, + routing_key=( + self.settings + .RABBITMQ_ANALYSIS_DLQ_ROUTING_KEY + ), + mandatory=True, + ), + timeout=( + self.settings + .RABBITMQ_PUBLISH_TIMEOUT_SECONDS + ), + ) + except TimeoutError as exception: + raise RetryableProcessingError( + message=( + "Sanitized dead-letter event " + "publication timed out" + ), + failure_code="DLQ_PUBLISH_TIMEOUT", + ) from exception + except Exception as exception: + raise RetryableProcessingError( + message=( + "Failed to publish sanitized " + "dead-letter event" + ), + failure_code="DLQ_PUBLISH_FAILED", + ) from exception + + logger.warning( + "Published sanitized dead-letter event. " + "event_id=%s original_message_id=%s " + "analysis_id=%s trace_id=%s failure_code=%s", + dead_letter_event.eventId, + original_message_id, + dead_letter_event.analysisId, + dead_letter_event.traceId, + failure_code, + ) + + return dead_letter_event \ No newline at end of file diff --git a/app/infrastructure/rabbitmq/handler.py b/app/infrastructure/rabbitmq/handler.py index bd486b6..9752547 100644 --- a/app/infrastructure/rabbitmq/handler.py +++ b/app/infrastructure/rabbitmq/handler.py @@ -8,8 +8,6 @@ logger = logging.getLogger(__name__) -class AnalysisPipelineError(RuntimeError): - """AI 분석 파이프라인 처리가 실패한 경우 발생""" class AnalysisRequestHandler: """분석 요청 이벤트를 AI 분석 파이프라인에 연결""" @@ -24,7 +22,7 @@ async def handle( self, event: AnalysisRequestedEvent, ) -> SmishingAnalysisResponse: - """분석 요청을 처리하고 결과의 성공 여부 검증""" + """분석 요청을 실행하고 파이프라인 응답을 그대로 반환""" logger.info( "Starting analysis request processing. " "event_id=%s analysis_id=%s trace_id=%s", @@ -37,22 +35,16 @@ async def handle( event.payload.content ) - if result.status != "SUCCESS": - logger.error( - "Analysis pipeline returned an error. " - "event_id=%s analysis_id=%s trace_id=%s " - "message=%s", + if result.status == "ERROR": + logger.warning( + "Analysis pipeline returned a failed result. " + "event_id=%s analysis_id=%s trace_id=%s", event.eventId, event.analysisId, event.traceId, - result.message, ) - raise AnalysisPipelineError( - "Analysis pipeline failed. " - f"analysis_id={event.analysisId} " - f"reason={result.message}" - ) + return result logger.info( "Analysis request processing completed. " diff --git a/app/infrastructure/rabbitmq/publisher.py b/app/infrastructure/rabbitmq/publisher.py new file mode 100644 index 0000000..e39bd36 --- /dev/null +++ b/app/infrastructure/rabbitmq/publisher.py @@ -0,0 +1,89 @@ +import asyncio +import logging + +from aio_pika import DeliveryMode, Message +from aio_pika.abc import AbstractRobustExchange + +from app.core.config import Settings, settings +from app.infrastructure.errors import ( + RetryableProcessingError, +) +from app.infrastructure.rabbitmq.schemas import ( + AnalysisEventType, + AnalysisResultEvent, +) + +logger = logging.getLogger(__name__) + + +class AnalysisResultPublisher: + """RabbitMQ Exchange로 분석 결과 이벤트 발행""" + + def __init__( + self, + exchange: AbstractRobustExchange, + app_settings: Settings = settings, + ) -> None: + self.exchange = exchange + self.settings = app_settings + + async def publish( + self, + event: AnalysisResultEvent, + ) -> None: + """분석 결과 이벤트를 JSON 메시지로 직렬화하여 라우팅 키로 발행""" + + # 이벤트 타입(COMPLETED, PARTIAL, FAILED)에 따른 라우팅 키 바인딩 + routing_key = { + AnalysisEventType.COMPLETED: + self.settings + .RABBITMQ_ANALYSIS_COMPLETED_ROUTING_KEY, + AnalysisEventType.PARTIAL: + self.settings + .RABBITMQ_ANALYSIS_PARTIAL_ROUTING_KEY, + AnalysisEventType.FAILED: + self.settings + .RABBITMQ_ANALYSIS_FAILED_ROUTING_KEY, + }[event.eventType] + + # aio_pika 메시지 객체 생성 + message = Message( + body=event.model_dump_json( + by_alias=True + ).encode("utf-8"), + content_type="application/json", + delivery_mode=DeliveryMode.PERSISTENT, + message_id=str(event.eventId), + correlation_id=str(event.traceId), + headers={ + "schemaVersion": event.schemaVersion, + "eventType": event.eventType.value, + }, + ) + + try: + await asyncio.wait_for( + self.exchange.publish( + message, + routing_key=routing_key, + mandatory=True, + ), + timeout=( + self.settings + .RABBITMQ_PUBLISH_TIMEOUT_SECONDS + ), + ) + except TimeoutError as exception: + raise RetryableProcessingError( + message=( + "Analysis result publication timed out" + ), + failure_code="RESULT_PUBLISH_TIMEOUT", + ) from exception + except Exception as exception: + raise RetryableProcessingError( + message=( + "Failed to publish analysis result" + ), + failure_code="RESULT_PUBLISH_FAILED", + ) from exception diff --git a/app/infrastructure/rabbitmq/result_factory.py b/app/infrastructure/rabbitmq/result_factory.py new file mode 100644 index 0000000..f7e8043 --- /dev/null +++ b/app/infrastructure/rabbitmq/result_factory.py @@ -0,0 +1,225 @@ +from datetime import datetime, timezone +from uuid import NAMESPACE_URL, uuid5 + +from app.analysis.execution import ( + AnalysisExecution, + AnalysisExecutionStatus, +) +from app.infrastructure.rabbitmq.schemas import ( + AnalysisEventType, + AnalysisRequestedEvent, + AnalysisResultEvent, + AnalysisResultPayload, + RawScores, + RuleAnalysisDetail, + TextAnalysisDetail, + TextAnalysisMethod, + UrlAnalysisDetail, + WeightedContributions, +) + + +class AnalysisResultEventFactory: + """내부 분석 실행 결과를 메시징 규격 이벤트로 변환""" + def create( + self, + request: AnalysisRequestedEvent, + execution: AnalysisExecution, + ) -> AnalysisResultEvent: + """요청 이벤트 데이터와 실행 결과를 결합하여 전달할 결과 이벤트를 생성""" + + # 내부 분석 상태(COMPLETED, PARTIAL, FAILED)를 외부 메시징 이벤트 타입으로 맵핑 + event_type = { + AnalysisExecutionStatus.COMPLETED: + AnalysisEventType.COMPLETED, + AnalysisExecutionStatus.PARTIAL: + AnalysisEventType.PARTIAL, + AnalysisExecutionStatus.FAILED: + AnalysisEventType.FAILED, + }[execution.status] + + result = execution.result + result_event_id = uuid5( + NAMESPACE_URL, + ( + "safefam:analysis-result:" + f"{request.eventId}:" + f"{event_type.value}" + ), + ) + + # 원본 요청의 식별자를 포함하여 결과 이벤트 구성 + return AnalysisResultEvent( + schemaVersion="1.0", + eventId=result_event_id, + causationId=request.eventId, + analysisId=request.analysisId, + clientMessageId=request.clientMessageId, + traceId=request.traceId, + occurredAt=datetime.now(timezone.utc), + eventType=event_type, + payload=build_payload( + result=result, + failed_tracks=execution.failed_tracks, + ), + ) + + +def build_payload( + result, + failed_tracks: tuple[str, ...], +) -> AnalysisResultPayload: + """내부 분석 응답을 외부 결과 이벤트 payload로 변환합니다.""" + text = result.text_analysis or {} + text_result = text.get("result") or {} + url = result.url_analysis or {} + rules = result.rule_analysis or {} + + text_score = _integer_score( + text_result.get("risk_score") + ) + url_score = _url_score( + url.get("url_risk_score") + ) + rule_score = _integer_score( + rules.get("rule_score") + ) + + is_failed = result.status == "ERROR" + + return AnalysisResultPayload( + finalScore=( + None if is_failed else result.final_score + ), + riskGrade=( + None + if is_failed + else getattr( + result.risk_grade, + "value", + result.risk_grade, + ) + ), + phishingType=None, + rawScores=RawScores( + text=text_score, + url=url_score, + rules=rule_score, + ), + weightedContributions=( + None + if is_failed + else WeightedContributions( + text=result.contribution_breakdown.llm, + url=result.contribution_breakdown.hybrid_url, + rules=result.contribution_breakdown.rules, + ) + ), + textAnalysis=( + None + if is_failed or not text + else _text_detail(text, failed_tracks) + ), + urlAnalysis=( + None + if is_failed or not url + else _url_detail(url) + ), + ruleAnalysis=( + None + if is_failed or not rules + else _rule_detail(rules) + ), + failedTracks=list(failed_tracks), + failureCode=( + "PIPELINE_FAILED" if is_failed else None + ), + ) + + +def _text_detail( + text: dict, + failed_tracks: tuple[str, ...], +) -> TextAnalysisDetail: + result = text.get("result") or {} + stage1 = text.get("stage1_naive_bayes") + + if result.get("grade") == "UNKNOWN": + method = TextAnalysisMethod.UNAVAILABLE + elif stage1 is not None: + method = TextAnalysisMethod.NAIVE_BAYES_GEMINI + elif text.get("engine") == "naive_bayes": + method = TextAnalysisMethod.NAIVE_BAYES + else: + method = TextAnalysisMethod.GEMINI + + return TextAnalysisDetail( + method=method, + score=_integer_score(result.get("risk_score")), + grade=result.get("grade"), + reason=result.get("reason"), + evidence=result.get("evidence") or [], + failedEngines=[ + track.removeprefix("TEXT:") + for track in failed_tracks + if track.startswith("TEXT:") + ], + ) + + +def _url_detail(url: dict) -> UrlAnalysisDetail: + return UrlAnalysisDetail( + hasUrl=bool(url.get("has_url")), + originalUrl=url.get("original_url"), + tracedUrl=url.get("origin_url"), + malicious=url.get("is_url_malicious"), + score=_url_score(url.get("url_risk_score")), + engineSource=url.get("engine_source"), + errorCode=_url_error_code(url), + ) + + +def _rule_detail(rules: dict) -> RuleAnalysisDetail: + return RuleAnalysisDetail( + score=_integer_score( + rules.get("rule_score") + ) or 0, + matchedRules=rules.get("matched_rules") or [], + maliciousDomainPattern=bool( + rules.get( + "has_malicious_domain_pattern", + False, + ) + ), + ) + + +def _integer_score(value) -> int | None: + if value is None: + return None + return max(0, min(100, round(float(value)))) + + +def _url_error_code(url: dict) -> str | None: + provider_codes = ( + url.get("provider_error_codes") or {} + ) + + if not provider_codes: + return None + + return ";".join( + f"{provider}:{code}" + for provider, code + in sorted(provider_codes.items()) + ) + + +def _url_score(value) -> int | None: + if value is None: + return None + + numeric = float(value) + if 0 <= numeric <= 1: + numeric *= 100 + return _integer_score(numeric) diff --git a/app/infrastructure/rabbitmq/schemas.py b/app/infrastructure/rabbitmq/schemas.py index 0e4aac7..995fbe2 100644 --- a/app/infrastructure/rabbitmq/schemas.py +++ b/app/infrastructure/rabbitmq/schemas.py @@ -8,13 +8,18 @@ ConfigDict, Field, field_validator, + model_validator, ) +from app.analysis.schemas import RiskGrade + + class AnalysisSource(str, Enum): - """분석 요청이 생성된 경로""" + """분석 요청이 생성된 경로를 정의""" AUTO = "AUTO" MANUAL = "MANUAL" + class AnalysisRequestedPayload(BaseModel): """AI 분석에 필요한 실제 문자 본문 데이터 스키마""" model_config = ConfigDict(extra="forbid") @@ -31,6 +36,7 @@ def validate_content(cls, value: str) -> str: raise ValueError("content must not be blank") return value + class AnalysisRequestedEvent(BaseModel): """Spring 메시징 시스템이 발행하는 ANALYSIS_REQUESTED v1 이벤트 스키마""" model_config = ConfigDict(extra="forbid") @@ -44,4 +50,200 @@ class AnalysisRequestedEvent(BaseModel): ) traceId: UUID occurredAt: AwareDatetime - payload: AnalysisRequestedPayload \ No newline at end of file + payload: AnalysisRequestedPayload + + +class AnalysisEventType(str, Enum): + """분석 결과 이벤트의 종합 처리 상태를 정의""" + COMPLETED = "ANALYSIS_COMPLETED" + PARTIAL = "ANALYSIS_PARTIAL" + FAILED = "ANALYSIS_FAILED" + + +class TextAnalysisMethod(str, Enum): + """텍스트 분석에 사용된 AI 및 알고리즘 방식을 정의""" + NAIVE_BAYES = "NAIVE_BAYES" + GEMINI = "GEMINI" + NAIVE_BAYES_GEMINI = "NAIVE_BAYES_GEMINI" + UNAVAILABLE = "UNAVAILABLE" + + +class RawScores(BaseModel): + """각 분석 트랙별(텍스트, URL, 룰) 원시 점수 데이터 스키마""" + model_config = ConfigDict(extra="forbid") + + text: int | None = Field(default=None, ge=0, le=100) + url: int | None = Field(default=None, ge=0, le=100) + rules: int | None = Field(default=None, ge=0, le=100) + + +class WeightedContributions(BaseModel): + """최종 위험도 점수에 반영된 트랙별 가중치 기여 점수 스키마""" + model_config = ConfigDict(extra="forbid") + + text: int = Field(ge=0, le=100) + url: int = Field(ge=0, le=100) + rules: int = Field(ge=0, le=100) + + +class TextAnalysisDetail(BaseModel): + """텍스트 분석 트랙의 세부 진단 결과 스키마""" + model_config = ConfigDict(extra="forbid") + + method: TextAnalysisMethod + score: int | None = Field(default=None, ge=0, le=100) + grade: str | None = None + reason: str | None = None + evidence: list[str] = Field(default_factory=list) + failedEngines: list[str] = Field(default_factory=list) + + +class UrlAnalysisDetail(BaseModel): + """URL 분석 트랙의 세부 진단 결과 스키마""" + model_config = ConfigDict(extra="forbid") + + hasUrl: bool + originalUrl: str | None = None + tracedUrl: str | None = None + malicious: bool | None = None + score: int | None = Field(default=None, ge=0, le=100) + engineSource: str | None = None + errorCode: str | None = None + + +class RuleAnalysisDetail(BaseModel): + """기반 룰 기반 탐지 트랙의 세부 진단 결과 스키마""" + model_config = ConfigDict(extra="forbid") + + score: int = Field(ge=0, le=100) + matchedRules: list[str] = Field(default_factory=list) + maliciousDomainPattern: bool = False + + +class AnalysisResultPayload(BaseModel): + """분석 결과 이벤트에 포함되는 통합 분석 페이로드 스키마""" + model_config = ConfigDict(extra="forbid") + + finalScore: int | None = Field(default=None, ge=0, le=100) + riskGrade: RiskGrade | None = None + phishingType: str | None = None + + rawScores: RawScores + weightedContributions: WeightedContributions | None = None + + textAnalysis: TextAnalysisDetail | None = None + urlAnalysis: UrlAnalysisDetail | None = None + ruleAnalysis: RuleAnalysisDetail | None = None + + failedTracks: list[str] = Field(default_factory=list) + failureCode: str | None = None + + @model_validator(mode="after") + def validate_failure_payload( + self, + ) -> "AnalysisResultPayload": + if self.failureCode is not None and ( + self.finalScore is not None + or self.riskGrade is not None + or self.weightedContributions is not None + ): + raise ValueError( + "failure payload must not contain " + "successful score fields" + ) + return self + + +class AnalysisResultEvent(BaseModel): + """FastAPI가 처리 후 Spring으로 발행하는 분석 결과 이벤트 스키마""" + model_config = ConfigDict(extra="forbid") + + schemaVersion: Literal["1.0"] + eventId: UUID + causationId: UUID + analysisId: int = Field(gt=0) + clientMessageId: str | None = None + traceId: UUID + occurredAt: AwareDatetime + eventType: AnalysisEventType + payload: AnalysisResultPayload + + @model_validator(mode="after") + def validate_event_result(self) -> "AnalysisResultEvent": + """이벤트 타입(COMPLETED, PARTIAL, FAILED)에 따른 필드 유효성을 검증""" + if self.eventType == AnalysisEventType.COMPLETED: + if self.payload.finalScore is None: + raise ValueError( + "completed event requires finalScore" + ) + if self.payload.riskGrade is None: + raise ValueError( + "completed event requires riskGrade" + ) + if ( + self.payload.failureCode is not None + or self.payload.failedTracks + ): + raise ValueError( + "completed event must not contain " + "failure indicators" + ) + + if self.eventType == AnalysisEventType.PARTIAL: + if not self.payload.failedTracks: + raise ValueError( + "partial event requires failedTracks" + ) + if self.payload.finalScore is None: + raise ValueError( + "partial event requires finalScore" + ) + if self.payload.riskGrade is None: + raise ValueError( + "partial event requires riskGrade" + ) + if self.payload.failureCode is not None: + raise ValueError( + "partial event must not contain " + "failureCode" + ) + + if self.eventType == AnalysisEventType.FAILED: + if not self.payload.failureCode: + raise ValueError( + "failed event requires failureCode" + ) + if ( + self.payload.finalScore is not None + or self.payload.riskGrade is not None + or self.payload.weightedContributions + is not None + ): + raise ValueError( + "failed event must not contain " + "successful score fields" + ) + + return self + +class DeadLetterEvent(BaseModel): + """원문과 개인정보를 제외한 실패 메시지를 격리하는 이벤트 스키마""" + + model_config = ConfigDict(extra="forbid") + + schemaVersion: Literal["1.0"] + eventId: UUID + originalMessageId: str | None = Field( + default=None, + max_length=255, + ) + analysisId: int | None = Field( + default=None, + gt=0, + ) + traceId: UUID | None = None + failureCode: str = Field( + min_length=1, + max_length=100, + ) + failedAt: AwareDatetime diff --git a/app/infrastructure/virustotal/client.py b/app/infrastructure/virustotal/client.py index dfe5e09..5a47ce2 100644 --- a/app/infrastructure/virustotal/client.py +++ b/app/infrastructure/virustotal/client.py @@ -1,102 +1,244 @@ -import logging import base64 +import logging import httpx from app.core.config import settings +from app.infrastructure.http_retry import request_with_retry logger = logging.getLogger(__name__) class VirusTotalClient: + """VirusTotal v3 API를 통해 URL의 악성 여부를 검사하고 스캔 요청""" def __init__(self): self.api_key = settings.VIRUSTOTAL_API_KEY self.base_url = "https://www.virustotal.com/api/v3" self.headers = { - "x-apikey": self.api_key if self.api_key else "", - "accept": "application/json" + "x-apikey": self.api_key or "", + "accept": "application/json", } - # URL을 VirusTotal v3 규격에 맞게 Base64 URL-Safe 인코딩 - def _get_url_id(self, url: str) -> str: + @staticmethod + def _safe_result() -> dict: + """VirusTotal이 실제로 URL을 조회한 뒤 안전하다고 판단한 결과""" + return { + "is_malicious": False, + "raw_score": 0.0, + "detected_count": 0, + "total_engines": 0, + "status": "safe", + "error_code": None, + } - b64_bytes = base64.urlsafe_b64encode(url.encode("utf-8")) + @staticmethod + def _unavailable_result(error_code: str) -> dict: + """API 장애로 URL의 안전 여부를 판단하지 못한 결과""" + return { + "is_malicious": False, + "raw_score": 0.0, + "detected_count": 0, + "total_engines": 0, + "status": "unavailable", + "error_code": error_code, + } + + @staticmethod + def _scanning_result() -> dict: + """신규 분석을 요청했지만 결과가 아직 준비되지 않은 상태""" + return { + "is_malicious": False, + "raw_score": 0.0, + "detected_count": 0, + "total_engines": 0, + "status": "scanning", + "error_code": None, + } + + def _get_url_id(self, url: str) -> str: + """URL을 Base64 URL-safe 식별자로 변환""" + b64_bytes = base64.urlsafe_b64encode( + url.encode("utf-8") + ) return b64_bytes.decode("utf-8").rstrip("=") - # VirusTotal API를 호출해 악성 여부 및 상세 스코어 판정 async def scan_url(self, url: str) -> dict: - - # 기본 안전 상태 스켈레톤 리턴 규격 정의 - default_result = {"is_malicious": False, "raw_score": 0.0, "detected_count": 0, "status": "safe"} - + """VT에 등록된 URL분석 보고서 조회하고 악성 위험도를 계산하여 반환""" if not self.api_key: - logger.warning("[VirusTotal] API Key가 누락되었습니다. 빈 분석 결과를 반환합니다.") - return default_result + logger.warning( + "[VirusTotal] API Key가 누락되어 " + "URL을 분석할 수 없습니다." + ) + return self._unavailable_result( + "MISSING_API_KEY" + ) url_id = self._get_url_id(url) report_url = f"{self.base_url}/urls/{url_id}" async with httpx.AsyncClient() as client: try: - # 기존 분석 보고서 조회 시도 - response = await client.get(report_url, headers=self.headers, timeout=5.0) + response = await request_with_retry( + lambda: client.get( + report_url, + headers=self.headers, + timeout=settings.VIRUSTOTAL_TIMEOUT_SECONDS, + ), + max_retries=settings.EXTERNAL_API_MAX_RETRIES, + operation_name="VirusTotal report", + ) - # 기존 보고서가 없을 때 신규 스캔 요청 분기 + # 기존 분석 보고서가 존재하지 않는 경우 신규 스캔 요청 if response.status_code == 404: - logger.info(f"[VirusTotal] 기존 보고서 없음. 신규 스캔 요청 시작: {url}") - scan_url = f"{self.base_url}/urls" - scan_response = await client.post(scan_url, headers=self.headers, data={"url": url}, timeout=5.0) - - if scan_response.status_code == 429: - logger.error("[VirusTotal] API 호출 한도 초과 (Rate Limit)") - return default_result + return await self._request_new_scan( + client=client, + url=url, + ) - scan_response.raise_for_status() - logger.info(f"[VirusTotal] 신규 스캔 요청 완료: {url}") - return {"is_malicious": False, "raw_score": 0.0, "detected_count": 0, "status": "scanning"} - - # Rate Limit 에러 처리 if response.status_code == 429: - logger.error("[VirusTotal] API 호출 한도 초과 (Rate Limit)") - return default_result + logger.error( + "[VirusTotal] API 호출 한도 초과" + ) + return self._unavailable_result( + "RATE_LIMITED" + ) response.raise_for_status() - report_data = response.json() - # 보고서 데이터 파싱 (오타 수정 및 디테일 파싱) - stats = report_data.get("data", {}).get("attributes", {}).get("last_analysis_stats", {}) + report_data = response.json() + stats = ( + report_data + .get("data", {}) + .get("attributes", {}) + .get("last_analysis_stats", {}) + ) malicious = stats.get("malicious", 0) suspicious = stats.get("suspicious", 0) - # last_analysis_stats에 잡힌 전체 엔진 수 (malicious/suspicious/harmless/undetected/timeout 등 전부 합산) - total_engines = sum(stats.values()) if stats else 0 + total_engines = ( + sum(stats.values()) if stats else 0 + ) logger.info( - f"[VirusTotal] 분석 완료 - 악성: {malicious}, 의심: {suspicious}, 전체 엔진: {total_engines}" + "[VirusTotal] 분석 완료 - " + "악성: %s, 의심: %s, 전체 엔진: %s", + malicious, + suspicious, + total_engines, ) - # 백신 엔진 중 3개 이상이 악성(malicious)이라고 판정하거나, 의심 엔진이 과도하게 많을 때 악성으로 분류 - is_malicious = (malicious >= 3) or (malicious + suspicious >= 5) + # 임계값 기준 + is_malicious = ( + malicious >= 3 + or malicious + suspicious >= 5 + ) - # 악성 판정 엔진 수 비율 기반 위험도 점수 산정 (의심 엔진은 절반 가중치로 반영, 최대 1.0) + # 전체 분석 엔진 대비 위험 비율 기반 가중치 점수 계산 if total_engines > 0: - malicious_ratio = malicious / total_engines - suspicious_ratio = suspicious / total_engines - raw_score = min(malicious_ratio + suspicious_ratio * 0.5, 1.0) + malicious_ratio = ( + malicious / total_engines + ) + suspicious_ratio = ( + suspicious / total_engines + ) + raw_score = min( + malicious_ratio + + suspicious_ratio * 0.5, + 1.0, + ) else: raw_score = 0.0 + status = ( + "completed" + if is_malicious + else "safe" + ) + return { "is_malicious": is_malicious, "raw_score": round(raw_score, 2), "detected_count": malicious, "total_engines": total_engines, - "status": "completed" + "status": status, + "error_code": None, } - except httpx.HTTPStatusError as e: - logger.error(f"[VirusTotal] API 에러 ({e.response.status_code}): {str(e)}") - return default_result + except httpx.HTTPStatusError as exc: + status_code = exc.response.status_code + + logger.error( + "[VirusTotal] API 에러 (%s): %s", + status_code, + exc, + ) + + if status_code == 429: + return self._unavailable_result( + "RATE_LIMITED" + ) + + return self._unavailable_result( + f"HTTP_{status_code}" + ) + except httpx.TimeoutException: - logger.error("[VirusTotal] API 요청 타임아웃 발생") - return default_result - except Exception as e: - logger.error(f"[VirusTotal] 연동 중 비정상 에러 발생: {str(e)}") - return default_result + logger.error( + "[VirusTotal] API 요청 타임아웃 발생" + ) + return self._unavailable_result("TIMEOUT") + + except httpx.RequestError as exc: + logger.error( + "[VirusTotal] 네트워크 오류: %s", + exc, + ) + return self._unavailable_result( + "NETWORK_ERROR" + ) + + except Exception: + logger.exception( + "[VirusTotal] 연동 중 비정상 오류 발생" + ) + return self._unavailable_result( + "UNEXPECTED_ERROR" + ) + + async def _request_new_scan( + self, + *, + client: httpx.AsyncClient, + url: str, + ) -> dict: + """기존 보고서가 없는 URL에 대해 VT에 신규 스캔 분석 요청""" + logger.info( + "[VirusTotal] 기존 보고서 없음. " + "신규 스캔 요청 시작: %s", + url, + ) + + scan_url = f"{self.base_url}/urls" + scan_response = await request_with_retry( + lambda: client.post( + scan_url, + headers=self.headers, + data={"url": url}, + timeout=settings.VIRUSTOTAL_TIMEOUT_SECONDS, + ), + max_retries=settings.EXTERNAL_API_MAX_RETRIES, + operation_name="VirusTotal scan", + ) + + if scan_response.status_code == 429: + logger.error( + "[VirusTotal] 신규 스캔 요청 한도 초과" + ) + return self._unavailable_result( + "RATE_LIMITED" + ) + + scan_response.raise_for_status() + + logger.info( + "[VirusTotal] 신규 스캔 요청 완료: %s", + url, + ) + return self._scanning_result() \ No newline at end of file diff --git a/app/main.py b/app/main.py index 27a7052..3c0cc89 100644 --- a/app/main.py +++ b/app/main.py @@ -16,6 +16,15 @@ from app.infrastructure.rabbitmq.handler import ( AnalysisRequestHandler, ) +from app.infrastructure.rabbitmq.publisher import ( + AnalysisResultPublisher, +) +from app.infrastructure.rabbitmq.result_factory import ( + AnalysisResultEventFactory, +) +from app.infrastructure.rabbitmq.dead_letter import ( + DeadLetterPublisher, +) def create_lifespan( @@ -40,9 +49,22 @@ async def lifespan( analysis_service=SmishingAnalysisService() ) + result_publisher = AnalysisResultPublisher( + exchange=rabbitmq.get_exchange(), + ) + + result_factory = AnalysisResultEventFactory() + + dead_letter_publisher = DeadLetterPublisher( + exchange=rabbitmq.get_exchange(), + ) + consumer = AnalysisRequestConsumer( request_queue=rabbitmq.get_request_queue(), handler=handler, + result_publisher=result_publisher, + result_factory=result_factory, + dead_letter_publisher=dead_letter_publisher, ) await consumer.start() diff --git a/tests/analysis/test_execution.py b/tests/analysis/test_execution.py new file mode 100644 index 0000000..8dfbec3 --- /dev/null +++ b/tests/analysis/test_execution.py @@ -0,0 +1,193 @@ +from app.analysis.execution import ( + AnalysisExecutionStatus, + classify_execution, +) +from app.analysis.schemas import ( + ContributionBreakdown, + RiskGrade, + SmishingAnalysisResponse, +) + + +def _result( + *, + status: str = "SUCCESS", + text_analysis: dict | None = None, + url_analysis: dict | None = None, + rule_analysis: dict | None = None, +) -> SmishingAnalysisResponse: + return SmishingAnalysisResponse( + status=status, + message="analysis result", + final_score=50, + risk_grade=RiskGrade.MEDIUM, + contribution_breakdown=ContributionBreakdown( + llm=30, + hybrid_url=10, + rules=10, + ), + text_analysis=text_analysis, + url_analysis=url_analysis, + rule_analysis=rule_analysis, + ) + + +def test_classifies_successful_execution_as_completed() -> None: + execution = classify_execution( + _result( + text_analysis={ + "result": { + "grade": "SAFE", + "error_message": None, + }, + }, + url_analysis={ + "available": True, + "failed_providers": [], + }, + rule_analysis={"error_message": None}, + ) + ) + + assert execution.status == AnalysisExecutionStatus.COMPLETED + assert execution.failed_tracks == () + + +def test_classifies_provider_failure_as_partial() -> None: + execution = classify_execution( + _result( + text_analysis={ + "result": { + "grade": "SAFE", + "error_message": None, + }, + }, + url_analysis={ + "available": True, + "failed_providers": ["GSB"], + }, + rule_analysis={"error_message": None}, + ) + ) + + assert execution.status == AnalysisExecutionStatus.PARTIAL + assert execution.failed_tracks == ("URL:GSB",) + + +def test_classifies_unavailable_url_track_as_partial() -> None: + execution = classify_execution( + _result( + text_analysis={ + "result": { + "grade": "SAFE", + "error_message": None, + }, + }, + url_analysis={ + "available": False, + "failed_providers": [ + "GSB", + "VIRUSTOTAL", + ], + }, + rule_analysis={"error_message": None}, + ) + ) + + assert execution.status == AnalysisExecutionStatus.PARTIAL + assert execution.failed_tracks == ("URL",) + + +def test_classifies_pipeline_error_as_failed() -> None: + execution = classify_execution( + _result(status="ERROR") + ) + + assert execution.status == AnalysisExecutionStatus.FAILED + assert execution.failed_tracks == ("PIPELINE",) + + +def test_classifies_gemini_failure_with_valid_naive_bayes() -> None: + execution = classify_execution( + _result( + text_analysis={ + "result": { + "grade": "UNKNOWN", + "error_message": "RATE_LIMITED", + }, + "stage1_naive_bayes": { + "grade": "DANGEROUS", + "error_message": None, + }, + }, + rule_analysis={"error_message": None}, + ) + ) + + assert execution.status == AnalysisExecutionStatus.PARTIAL + assert execution.failed_tracks == ("TEXT:GEMINI",) + + +def test_classifies_naive_bayes_failure_with_valid_gemini() -> None: + execution = classify_execution( + _result( + text_analysis={ + "result": { + "grade": "SAFE", + "error_message": None, + }, + "stage1_naive_bayes": { + "grade": "UNKNOWN", + "error_message": "MODEL_UNAVAILABLE", + }, + }, + rule_analysis={"error_message": None}, + ) + ) + + assert execution.status == AnalysisExecutionStatus.PARTIAL + assert execution.failed_tracks == ( + "TEXT:NAIVE_BAYES", + ) + + +def test_classifies_unknown_naive_bayes_without_error_code() -> None: + execution = classify_execution( + _result( + text_analysis={ + "result": { + "grade": "SAFE", + "error_message": None, + }, + "stage1_naive_bayes": { + "grade": "UNKNOWN", + "error_message": None, + }, + }, + rule_analysis={"error_message": None}, + ) + ) + + assert execution.status == AnalysisExecutionStatus.PARTIAL + assert execution.failed_tracks == ( + "TEXT:NAIVE_BAYES", + ) + + +def test_classifies_rule_failure_as_partial() -> None: + execution = classify_execution( + _result( + text_analysis={ + "result": { + "grade": "SAFE", + "error_message": None, + }, + }, + rule_analysis={ + "error_message": "RULE_ANALYSIS_FAILED", + }, + ) + ) + + assert execution.status == AnalysisExecutionStatus.PARTIAL + assert execution.failed_tracks == ("RULES",) diff --git a/tests/analysis/test_scoring.py b/tests/analysis/test_scoring.py index 6e83f02..16181e8 100644 --- a/tests/analysis/test_scoring.py +++ b/tests/analysis/test_scoring.py @@ -1,4 +1,6 @@ from app.analysis.schemas import RiskGrade +import pytest + from app.analysis.scoring import RiskScoringEngine ScoringEngine = RiskScoringEngine @@ -206,3 +208,36 @@ def test_calculate_score_without_url_raises_ceiling_above_old_50_point_cap(): assert breakdown.llm == 65 assert final_score == 65 assert final_score > 50 + +def test_redistributes_url_weight_when_url_unavailable(): + final_score, _, breakdown = ( + RiskScoringEngine.calculate_score( + llm_score=70, + is_url_malicious=False, + url_risk_score=0.0, + rule_score=70, + has_url=True, + url_available=False, + ) + ) + + assert breakdown.hybrid_url == 0 + assert breakdown.llm == 50 + assert breakdown.rules == 20 + assert final_score == 70 + +def test_raises_when_all_tracks_are_unavailable(): + with pytest.raises( + ValueError, + match="No analysis tracks are available", + ): + RiskScoringEngine.calculate_score( + llm_score=0, + is_url_malicious=False, + url_risk_score=0.0, + rule_score=0, + has_url=True, + text_available=False, + url_available=False, + rules_available=False, + ) diff --git a/tests/analysis/test_service.py b/tests/analysis/test_service.py index 243414d..ecd6c8e 100644 --- a/tests/analysis/test_service.py +++ b/tests/analysis/test_service.py @@ -141,7 +141,7 @@ async def test_analyze_pipeline_forces_high_when_local_domain_rule_matches(mock_ @pytest.mark.asyncio @patch("app.analysis.service.analyze_text_with_gemini", new_callable=AsyncMock) @patch("app.analysis.service.analyze_text_with_naive_bayes", new_callable=AsyncMock) -async def test_analyze_pipeline_does_not_fail_open_when_both_text_engines_are_down(mock_nb, mock_gemini): +async def test_analyze_pipeline_uses_available_zero_score_rules_when_text_engines_are_down(mock_nb, mock_gemini): """ 나이브 베이즈 모델 로드 실패 + Gemini 호출도 동시에 실패(rate limit 등)하는 경우, URL/규칙 신호가 전혀 없는 문자라도 최종 등급이 조용히 LOW로 나와선 안 된다. @@ -156,8 +156,10 @@ async def test_analyze_pipeline_does_not_fail_open_when_both_text_engines_are_do service = SmishingAnalysisService() result = await service.analyze_pipeline("URL도 없고 특이사항도 없는 문자") - assert result.risk_grade != "LOW" - assert result.final_score >= 40 + assert result.status == "SUCCESS" + assert result.rule_analysis["rule_score"] == 0 + assert result.risk_grade == "LOW" + assert result.final_score == 0 @pytest.mark.asyncio diff --git a/tests/analysis/url/test_analyzer.py b/tests/analysis/url/test_analyzer.py index 249c1dd..c037abb 100644 --- a/tests/analysis/url/test_analyzer.py +++ b/tests/analysis/url/test_analyzer.py @@ -10,7 +10,11 @@ async def test_gsb_block_marks_is_gsb_confirmed_true_and_skips_vt(): VT는 quota 절약을 위해 호출되지 않아야 한다. """ engine = HybridUrlAnalyzer() - engine.gsb_client.scan_url = AsyncMock(return_value={"is_malicious": True, "raw_score": 0.95}) + engine.gsb_client.scan_url = AsyncMock(return_value={ + "is_malicious": True, + "raw_score": 0.95, + "status": "completed", + }) engine.vt_client.scan_url = AsyncMock() result = await engine.scan_url("https://danger-phishing-test-site.com") @@ -27,8 +31,16 @@ async def test_vt_only_detection_does_not_mark_gsb_confirmed(): is_gsb_confirmed는 False여야 한다 (GSB 확정 오버라이드 트리거 대상이 아님). """ engine = HybridUrlAnalyzer() - engine.gsb_client.scan_url = AsyncMock(return_value={"is_malicious": False}) - engine.vt_client.scan_url = AsyncMock(return_value={"is_malicious": True, "detected_count": 5, "raw_score": 0.8}) + engine.gsb_client.scan_url = AsyncMock(return_value={ + "is_malicious": False, + "status": "safe", + }) + engine.vt_client.scan_url = AsyncMock(return_value={ + "is_malicious": True, + "detected_count": 5, + "raw_score": 0.8, + "status": "completed", + }) result = await engine.scan_url("https://hidden-malware-link.xyz") @@ -43,8 +55,16 @@ async def test_vt_weak_detection_below_threshold_is_not_confirmed(): 기존처럼 가중치 기반 점수로만 반영되어야 한다 (오탐 벤더 섞일 가능성 고려). """ engine = HybridUrlAnalyzer() - engine.gsb_client.scan_url = AsyncMock(return_value={"is_malicious": False}) - engine.vt_client.scan_url = AsyncMock(return_value={"is_malicious": True, "detected_count": 2, "raw_score": 0.3}) + engine.gsb_client.scan_url = AsyncMock(return_value={ + "is_malicious": False, + "status": "safe", + }) + engine.vt_client.scan_url = AsyncMock(return_value={ + "is_malicious": True, + "detected_count": 2, + "raw_score": 0.3, + "status": "completed", + }) result = await engine.scan_url("https://borderline-site.example") @@ -58,8 +78,16 @@ async def test_vt_strong_consensus_at_or_above_threshold_is_confirmed(): VT 탐지 엔진 수가 임계치(5) 이상이면 다수 백신사 합의로 보고 확정 악성으로 승격되어야 한다. """ engine = HybridUrlAnalyzer() - engine.gsb_client.scan_url = AsyncMock(return_value={"is_malicious": False}) - engine.vt_client.scan_url = AsyncMock(return_value={"is_malicious": True, "detected_count": 7, "raw_score": 0.9}) + engine.gsb_client.scan_url = AsyncMock(return_value={ + "is_malicious": False, + "status": "safe", + }) + engine.vt_client.scan_url = AsyncMock(return_value={ + "is_malicious": True, + "detected_count": 7, + "raw_score": 0.9, + "status": "completed", + }) result = await engine.scan_url("https://hidden-malware-link.xyz") @@ -70,7 +98,11 @@ async def test_vt_strong_consensus_at_or_above_threshold_is_confirmed(): async def test_gsb_blocked_branch_skips_vt_so_vt_confirmed_is_false(): """GSB가 이미 차단해서 VT 호출 자체를 생략한 경우 is_vt_confirmed는 False여야 한다.""" engine = HybridUrlAnalyzer() - engine.gsb_client.scan_url = AsyncMock(return_value={"is_malicious": True, "raw_score": 0.95}) + engine.gsb_client.scan_url = AsyncMock(return_value={ + "is_malicious": True, + "raw_score": 0.95, + "status": "completed", + }) engine.vt_client.scan_url = AsyncMock() result = await engine.scan_url("https://danger-phishing-test-site.com") @@ -81,10 +113,41 @@ async def test_gsb_blocked_branch_skips_vt_so_vt_confirmed_is_false(): @pytest.mark.asyncio async def test_clean_url_is_not_gsb_confirmed(): engine = HybridUrlAnalyzer() - engine.gsb_client.scan_url = AsyncMock(return_value={"is_malicious": False}) - engine.vt_client.scan_url = AsyncMock(return_value={"is_malicious": False, "detected_count": 0}) + engine.gsb_client.scan_url = AsyncMock(return_value={ + "is_malicious": False, + "status": "safe", + }) + engine.vt_client.scan_url = AsyncMock(return_value={ + "is_malicious": False, + "detected_count": 0, + "status": "safe", + }) result = await engine.scan_url("https://www.google.com") assert result["is_malicious"] is False assert result["is_gsb_confirmed"] is False + +@pytest.mark.asyncio +async def test_single_vt_detection_is_not_malicious(): + engine = HybridUrlAnalyzer() + engine.gsb_client.scan_url = AsyncMock( + return_value={ + "is_malicious": False, + "status": "safe", + } + ) + engine.vt_client.scan_url = AsyncMock( + return_value={ + "is_malicious": False, + "detected_count": 1, + "raw_score": 0.1, + "status": "safe", + } + ) + + result = await engine.scan_url( + "https://example.com" + ) + + assert result["is_malicious"] is False diff --git a/tests/infrastructure/rabbitmq/test_connection.py b/tests/infrastructure/rabbitmq/test_connection.py index 2721fe7..9107901 100644 --- a/tests/infrastructure/rabbitmq/test_connection.py +++ b/tests/infrastructure/rabbitmq/test_connection.py @@ -22,6 +22,12 @@ def create_fake_settings(): RABBITMQ_ANALYSIS_REQUEST_ROUTING_KEY=( "analysis.requested.v1" ), + RABBITMQ_ANALYSIS_DLQ=( + "safefam.analysis.requested.dlq" + ), + RABBITMQ_ANALYSIS_DLQ_ROUTING_KEY=( + "analysis.requested.dead.v1" + ), RABBITMQ_PREFETCH_COUNT=1, ) @@ -41,10 +47,14 @@ async def test_connect_initializes_request_topology( fake_channel = AsyncMock() fake_exchange = AsyncMock() fake_queue = AsyncMock() + fake_dead_letter_queue = AsyncMock() fake_connection.channel.return_value = fake_channel fake_channel.declare_exchange.return_value = fake_exchange - fake_channel.declare_queue.return_value = fake_queue + fake_channel.declare_queue.side_effect = [ + fake_queue, + fake_dead_letter_queue, + ] mock_connect_robust.return_value = fake_connection @@ -56,7 +66,10 @@ async def test_connect_initializes_request_topology( mock_connect_robust.assert_awaited_once_with( app_settings.RABBITMQ_URL ) - fake_connection.channel.assert_awaited_once() + fake_connection.channel.assert_awaited_once_with( + publisher_confirms=True, + on_return_raises=True, + ) fake_channel.set_qos.assert_awaited_once_with( prefetch_count=1 @@ -68,18 +81,32 @@ async def test_connect_initializes_request_topology( durable=True, ) - fake_channel.declare_queue.assert_awaited_once_with( + fake_channel.declare_queue.assert_any_await( "safefam.analysis.requested.q", durable=True, ) + fake_channel.declare_queue.assert_any_await( + "safefam.analysis.requested.dlq", + durable=True, + ) + assert fake_channel.declare_queue.await_count == 2 fake_queue.bind.assert_awaited_once_with( fake_exchange, routing_key="analysis.requested.v1", ) + fake_dead_letter_queue.bind.assert_awaited_once_with( + fake_exchange, + routing_key="analysis.requested.dead.v1", + ) assert rabbitmq.exchange is fake_exchange assert rabbitmq.request_queue is fake_queue + assert ( + rabbitmq.dead_letter_queue + is fake_dead_letter_queue + ) + assert rabbitmq.get_exchange() is fake_exchange @pytest.mark.asyncio @patch( @@ -90,7 +117,7 @@ async def test_connect_initializes_request_topology( async def test_connect_does_not_open_duplicate_connection( mock_connect_robust: AsyncMock, ): - """중복 연결 방지 테스트.""" + """중복 연결 방지 테스트""" fake_connection = AsyncMock() fake_connection.is_closed = False @@ -124,6 +151,19 @@ def test_get_request_queue_fails_before_connect(): ): rabbitmq.get_request_queue() + +def test_get_exchange_fails_before_connect(): + """연결 전 Exchange 접근 테스트.""" + rabbitmq = RabbitMQConnection( + create_fake_settings() + ) + + with pytest.raises( + RabbitMQNotConnectedError, + match="exchange is not initialized", + ): + rabbitmq.get_exchange() + @pytest.mark.asyncio async def test_close_closes_connection_and_clears_resources(): """연결 종료 테스트""" @@ -138,6 +178,7 @@ async def test_close_closes_connection_and_clears_resources(): rabbitmq.channel = AsyncMock() rabbitmq.exchange = AsyncMock() rabbitmq.request_queue = AsyncMock() + rabbitmq.dead_letter_queue = AsyncMock() await rabbitmq.close() @@ -146,7 +187,8 @@ async def test_close_closes_connection_and_clears_resources(): assert rabbitmq.connection is None assert rabbitmq.channel is None assert rabbitmq.exchange is None - assert rabbitmq.request_queue is None + assert rabbitmq.request_queue is None + assert rabbitmq.dead_letter_queue is None @pytest.mark.asyncio async def test_close_does_not_close_already_closed_connection(): diff --git a/tests/infrastructure/rabbitmq/test_consumer.py b/tests/infrastructure/rabbitmq/test_consumer.py index 21e0b41..7cd382e 100644 --- a/tests/infrastructure/rabbitmq/test_consumer.py +++ b/tests/infrastructure/rabbitmq/test_consumer.py @@ -1,5 +1,6 @@ import json -from unittest.mock import AsyncMock +import asyncio +from unittest.mock import AsyncMock, Mock import pytest @@ -8,6 +9,9 @@ RiskGrade, SmishingAnalysisResponse, ) +from app.infrastructure.errors import ( + NonRetryableProcessingError, +) from app.infrastructure.rabbitmq.consumer import ( AnalysisRequestConsumer, ) @@ -57,6 +61,24 @@ def create_success_result() -> SmishingAnalysisResponse: rule_analysis=None, ) + +def create_error_result() -> SmishingAnalysisResponse: + return SmishingAnalysisResponse( + status="ERROR", + message="All analysis tracks failed.", + final_score=40, + risk_grade=RiskGrade.MEDIUM, + contribution_breakdown=ContributionBreakdown( + llm=0, + hybrid_url=0, + rules=0, + ), + text_analysis=None, + url_analysis=None, + rule_analysis=None, + ) + + def create_message( *, body: bytes | None = None, @@ -67,6 +89,7 @@ def create_message( message.body = body or create_valid_message_body() message.message_id = "rabbit-message-001" message.redelivered = redelivered + message.headers = {} return message @@ -74,9 +97,27 @@ def create_consumer(): request_queue = AsyncMock() handler = AsyncMock() + result_publisher = Mock() + result_publisher.publish = AsyncMock() + + result_factory = Mock() + result_factory.create.return_value = Mock( + eventId="result-event-id", + eventType=Mock( + value="ANALYSIS_COMPLETED" + ), + ) + + dead_letter_publisher = Mock() + dead_letter_publisher.publish = AsyncMock() + consumer = AnalysisRequestConsumer( request_queue=request_queue, handler=handler, + result_publisher=result_publisher, + result_factory=result_factory, + dead_letter_publisher=dead_letter_publisher, + requeue_backoff_seconds=0, ) return consumer, request_queue, handler @@ -100,28 +141,89 @@ async def test_consumer_acknowledges_successful_message(): "[국민은행] 계좌가 정지되었습니다." ) + consumer.result_factory.create.assert_called_once() + factory_arguments = ( + consumer.result_factory.create.call_args.kwargs + ) + assert factory_arguments["request"] is handled_event + assert ( + factory_arguments["execution"].status.value + == "COMPLETED" + ) + + result_event = ( + consumer.result_factory.create.return_value + ) + consumer.result_publisher.publish.assert_awaited_once_with( + result_event + ) + message.ack.assert_awaited_once() message.nack.assert_not_awaited() message.reject.assert_not_awaited() + +@pytest.mark.asyncio +async def test_consumer_does_not_ack_when_publication_fails(): + """결과 이벤트 발행 실패 시 요청 메시지를 ACK X""" + consumer, _, handler = create_consumer() + message = create_message() + + handler.handle.return_value = create_success_result() + consumer.result_publisher.publish.side_effect = ( + RuntimeError("RabbitMQ publish failed") + ) + + await consumer._on_message(message) + + consumer.result_publisher.publish.assert_awaited_once() + message.ack.assert_not_awaited() + message.nack.assert_awaited_once_with(requeue=True) + message.reject.assert_not_awaited() + + @pytest.mark.asyncio -async def test_consumer_rejects_invalid_json_without_requeue(): - """잘못된 JSON reject 테스트""" +async def test_consumer_publishes_failed_result_and_acks(): + """ERROR 응답을 FAILED 결과 이벤트로 발행한 뒤 ACK""" + consumer, _, handler = create_consumer() + message = create_message() + + handler.handle.return_value = create_error_result() + + await consumer._on_message(message) + + factory_arguments = ( + consumer.result_factory.create.call_args.kwargs + ) + assert ( + factory_arguments["execution"].status.value + == "FAILED" + ) + consumer.result_publisher.publish.assert_awaited_once() + message.ack.assert_awaited_once() + message.nack.assert_not_awaited() + +@pytest.mark.asyncio +async def test_consumer_routes_invalid_json_to_sanitized_dlq(): + """잘못된 JSON은 정제된 DLQ 이벤트 발행 후 ACK합니다.""" consumer, _, handler = create_consumer() message = create_message(body=b"{invalid-json") await consumer._on_message(message) handler.handle.assert_not_awaited() - message.reject.assert_awaited_once_with( - requeue=False + consumer.dead_letter_publisher.publish.assert_awaited_once_with( + original_message_id=message.message_id, + failure_code="INVALID_JSON", + request_event=None, ) - message.ack.assert_not_awaited() + message.ack.assert_awaited_once() message.nack.assert_not_awaited() + message.reject.assert_not_awaited() @pytest.mark.asyncio -async def test_consumer_rejects_invalid_event_schema(): - """잘못된 이벤트 스키마 reject 테스트""" +async def test_consumer_routes_unsupported_schema_to_dlq(): + """지원하지 않는 버전은 정제된 DLQ 이벤트로 격리합니다.""" consumer, _, handler = create_consumer() invalid_event = { @@ -137,10 +239,39 @@ async def test_consumer_rejects_invalid_event_schema(): await consumer._on_message(message) handler.handle.assert_not_awaited() - message.reject.assert_awaited_once_with( - requeue=False + consumer.dead_letter_publisher.publish.assert_awaited_once_with( + original_message_id=message.message_id, + failure_code="UNSUPPORTED_SCHEMA_VERSION", + request_event=None, ) - message.ack.assert_not_awaited() + message.ack.assert_awaited_once() + message.nack.assert_not_awaited() + message.reject.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_consumer_routes_invalid_event_schema_to_dlq(): + """v1 형식 오류는 INVALID_EVENT_SCHEMA로 격리합니다.""" + consumer, _, handler = create_consumer() + message = create_message( + body=json.dumps( + { + "schemaVersion": "1.0", + "eventId": "not-a-uuid", + "analysisId": 0, + } + ).encode("utf-8") + ) + + await consumer._on_message(message) + + handler.handle.assert_not_awaited() + consumer.dead_letter_publisher.publish.assert_awaited_once_with( + original_message_id=message.message_id, + failure_code="INVALID_EVENT_SCHEMA", + request_event=None, + ) + message.ack.assert_awaited_once() message.nack.assert_not_awaited() @pytest.mark.asyncio @@ -163,8 +294,8 @@ async def test_consumer_requeues_first_processing_failure(): message.reject.assert_not_awaited() @pytest.mark.asyncio -async def test_consumer_does_not_treat_redelivery_as_retry_attempt(): - """Broker redelivery is not an application retry attempt.""" +async def test_consumer_treats_redelivery_as_retry_attempt(): + """Broker redelivery state survives restarts and bounds retries.""" consumer, _, handler = create_consumer() message = create_message(redelivered=True) @@ -175,15 +306,13 @@ async def test_consumer_does_not_treat_redelivery_as_retry_attempt(): await consumer._on_message(message) handler.handle.assert_awaited_once() - message.nack.assert_awaited_once_with( - requeue=True - ) - message.ack.assert_not_awaited() - message.reject.assert_not_awaited() + consumer.dead_letter_publisher.publish.assert_awaited_once() + message.ack.assert_awaited_once() + message.nack.assert_not_awaited() @pytest.mark.asyncio -async def test_consumer_rejects_after_recorded_retry_fails(): - """Reject after the recorded application retry also fails.""" +async def test_consumer_routes_to_dlq_after_retry_fails(): + """기록된 재시도까지 실패하면 정제 DLQ로 격리합니다.""" consumer, _, handler = create_consumer() first_message = create_message(redelivered=False) second_message = create_message(redelivered=True) @@ -198,24 +327,80 @@ async def test_consumer_rejects_after_recorded_retry_fails(): first_message.nack.assert_awaited_once_with( requeue=True ) - second_message.reject.assert_awaited_once_with( - requeue=False + dlq_arguments = ( + consumer.dead_letter_publisher + .publish.call_args.kwargs ) - assert consumer.retry_attempts == {} + assert ( + dlq_arguments["original_message_id"] + == second_message.message_id + ) + assert ( + dlq_arguments["failure_code"] + == "PROCESSING_RETRIES_EXHAUSTED" + ) + assert dlq_arguments["request_event"].analysisId == 123 + second_message.ack.assert_awaited_once() + second_message.reject.assert_not_awaited() + @pytest.mark.asyncio -async def test_consumer_clears_retry_state_after_success(): - """Clear application retry state after successful processing.""" +async def test_consumer_requeues_when_dlq_publication_fails(): + """DLQ 발행 실패 시 원본 유실을 막기 위해 requeue합니다.""" + consumer, _, handler = create_consumer() + message = create_message(body=b"{invalid-json") + + consumer.dead_letter_publisher.publish.side_effect = ( + RuntimeError("DLQ unavailable") + ) + + await consumer._on_message(message) + + handler.handle.assert_not_awaited() + message.ack.assert_not_awaited() + message.nack.assert_awaited_once_with(requeue=True) + message.reject.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_consumer_routes_non_retryable_processing_error_to_dlq(): + """처리 중 재시도 불가 오류는 즉시 DLQ로 격리합니다.""" consumer, _, handler = create_consumer() message = create_message() - event_id = "1fb898fa-d89d-4d0b-a43f-a8b00daeb765" - consumer.retry_attempts[event_id] = 1 - handler.handle.return_value = create_success_result() + handler.handle.side_effect = NonRetryableProcessingError( + message="Invalid provider credentials", + failure_code="INVALID_PROVIDER_CREDENTIALS", + ) + + await consumer._on_message(message) + + dlq_arguments = ( + consumer.dead_letter_publisher + .publish.call_args.kwargs + ) + assert ( + dlq_arguments["failure_code"] + == "INVALID_PROVIDER_CREDENTIALS" + ) + assert dlq_arguments["request_event"].analysisId == 123 + message.ack.assert_awaited_once() + message.nack.assert_not_awaited() + +@pytest.mark.asyncio +async def test_consumer_uses_broker_delivery_count(): + """Quorum queue delivery count is the authoritative retry state.""" + consumer, _, handler = create_consumer() + message = create_message() + message.headers = {"x-delivery-count": 1} + + handler.handle.side_effect = RuntimeError( + "Temporary analysis failure" + ) await consumer._on_message(message) - assert event_id not in consumer.retry_attempts + consumer.dead_letter_publisher.publish.assert_awaited_once() message.ack.assert_awaited_once() @pytest.mark.asyncio @@ -271,3 +456,43 @@ async def test_consumer_stop_before_start_does_nothing(): await consumer.stop() request_queue.cancel.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_consumer_stop_waits_for_in_flight_task(): + consumer, request_queue, _ = create_consumer() + consumer.consumer_tag = "analysis-consumer-tag" + consumer.shutdown_timeout_seconds = 1 + + task = asyncio.create_task(asyncio.sleep(0)) + consumer.in_flight_tasks.add(task) + + await consumer.stop() + + request_queue.cancel.assert_awaited_once_with( + "analysis-consumer-tag" + ) + assert task.done() + + +@pytest.mark.asyncio +async def test_consumer_stop_returns_after_shutdown_timeout(): + consumer, request_queue, _ = create_consumer() + consumer.consumer_tag = "analysis-consumer-tag" + consumer.shutdown_timeout_seconds = 0 + + blocker = asyncio.Event() + task = asyncio.create_task(blocker.wait()) + consumer.in_flight_tasks.add(task) + + try: + await consumer.stop() + + request_queue.cancel.assert_awaited_once_with( + "analysis-consumer-tag" + ) + assert not task.done() + finally: + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task diff --git a/tests/infrastructure/rabbitmq/test_dead_letter.py b/tests/infrastructure/rabbitmq/test_dead_letter.py new file mode 100644 index 0000000..0b92598 --- /dev/null +++ b/tests/infrastructure/rabbitmq/test_dead_letter.py @@ -0,0 +1,183 @@ +import asyncio +import json +from datetime import datetime, timezone +from types import SimpleNamespace +from unittest.mock import AsyncMock +from uuid import uuid4 + +import pytest +from aio_pika import DeliveryMode + +from app.infrastructure.errors import ( + RetryableProcessingError, +) +from app.infrastructure.rabbitmq.dead_letter import ( + DeadLetterPublisher, +) +from app.infrastructure.rabbitmq.schemas import ( + AnalysisRequestedEvent, +) + + +SENSITIVE_CONTENT = ( + "[국민은행] 계좌가 정지되었습니다. " + "https://malicious.example/login" +) + + +def create_settings( + *, + timeout: float = 1.0, +) -> SimpleNamespace: + return SimpleNamespace( + RABBITMQ_ANALYSIS_DLQ_ROUTING_KEY=( + "analysis.requested.dead.v1" + ), + RABBITMQ_PUBLISH_TIMEOUT_SECONDS=timeout, + ) + + +def create_request_event() -> AnalysisRequestedEvent: + return AnalysisRequestedEvent.model_validate( + { + "schemaVersion": "1.0", + "eventId": str(uuid4()), + "analysisId": 123, + "clientMessageId": "sms-dlq-001", + "traceId": str(uuid4()), + "occurredAt": datetime.now(timezone.utc), + "payload": { + "sender": "1588-0000", + "content": SENSITIVE_CONTENT, + "receivedAt": datetime.now(timezone.utc), + "source": "AUTO", + }, + } + ) + + +@pytest.mark.asyncio +async def test_publish_sends_sanitized_persistent_event(): + """DLQ 메시지에는 추적 정보만 포함하고 원문은 제외합니다.""" + exchange = AsyncMock() + publisher = DeadLetterPublisher( + exchange=exchange, + app_settings=create_settings(), + ) + request_event = create_request_event() + + dead_letter_event = await publisher.publish( + original_message_id="rabbit-message-001", + failure_code="PROCESSING_RETRIES_EXHAUSTED", + request_event=request_event, + ) + + exchange.publish.assert_awaited_once() + message = exchange.publish.await_args.args[0] + arguments = exchange.publish.await_args.kwargs + + assert ( + arguments["routing_key"] + == "analysis.requested.dead.v1" + ) + assert arguments["mandatory"] is True + assert message.delivery_mode == DeliveryMode.PERSISTENT + assert message.content_type == "application/json" + assert message.message_id == str( + dead_letter_event.eventId + ) + assert message.correlation_id == str( + request_event.traceId + ) + assert message.headers["sanitized"] is True + + body_text = message.body.decode("utf-8") + body = json.loads(body_text) + + assert body["analysisId"] == 123 + assert body["traceId"] == str( + request_event.traceId + ) + assert ( + body["failureCode"] + == "PROCESSING_RETRIES_EXHAUSTED" + ) + assert "payload" not in body + assert "content" not in body + assert "sender" not in body + assert "originalUrl" not in body + assert SENSITIVE_CONTENT not in body_text + assert request_event.payload.sender not in body_text + + +@pytest.mark.asyncio +async def test_publish_without_valid_request_uses_no_analysis_data(): + """파싱 불가능한 원본은 메시지 ID와 실패 코드만 남깁니다.""" + exchange = AsyncMock() + publisher = DeadLetterPublisher( + exchange=exchange, + app_settings=create_settings(), + ) + + event = await publisher.publish( + original_message_id="rabbit-invalid-001", + failure_code="INVALID_JSON", + request_event=None, + ) + + assert event.originalMessageId == "rabbit-invalid-001" + assert event.analysisId is None + assert event.traceId is None + + +@pytest.mark.asyncio +async def test_publish_timeout_becomes_retryable_error(): + """DLQ timeout은 원본 재전달을 위한 재시도 오류로 변환합니다.""" + async def delayed_publish(*args, **kwargs): + await asyncio.sleep(0.1) + + exchange = AsyncMock() + exchange.publish.side_effect = delayed_publish + + publisher = DeadLetterPublisher( + exchange=exchange, + app_settings=create_settings(timeout=0.01), + ) + + with pytest.raises( + RetryableProcessingError, + match="publication timed out", + ) as error: + await publisher.publish( + original_message_id="rabbit-message-001", + failure_code="INVALID_JSON", + request_event=None, + ) + + assert error.value.failure_code == "DLQ_PUBLISH_TIMEOUT" + + +@pytest.mark.asyncio +async def test_publish_failure_becomes_retryable_error(): + """브로커 발행 실패는 재시도 가능한 오류로 변환합니다.""" + exchange = AsyncMock() + exchange.publish.side_effect = RuntimeError( + "RabbitMQ unavailable" + ) + + publisher = DeadLetterPublisher( + exchange=exchange, + app_settings=create_settings(), + ) + + with pytest.raises( + RetryableProcessingError, + match="Failed to publish", + ) as error: + await publisher.publish( + original_message_id="rabbit-message-001", + failure_code="INVALID_JSON", + request_event=None, + ) + + assert error.value.failure_code == "DLQ_PUBLISH_FAILED" diff --git a/tests/infrastructure/rabbitmq/test_handler.py b/tests/infrastructure/rabbitmq/test_handler.py index b4b74a7..eed10c2 100644 --- a/tests/infrastructure/rabbitmq/test_handler.py +++ b/tests/infrastructure/rabbitmq/test_handler.py @@ -9,7 +9,6 @@ SmishingAnalysisResponse, ) from app.infrastructure.rabbitmq.handler import ( - AnalysisPipelineError, AnalysisRequestHandler, ) from app.infrastructure.rabbitmq.schemas import ( @@ -111,23 +110,23 @@ async def test_handler_passes_event_content_to_analysis_pipeline(): assert actual_result is expected_result @pytest.mark.asyncio -async def test_handler_raises_when_pipeline_returns_error(): +async def test_handler_returns_pipeline_error_result(): """실패 변환 테스트""" event = create_analysis_requested_event() error_result = create_error_result() analysis_service = AsyncMock() - analysis_service.analyze_pipeline.return_value = error_result + analysis_service.analyze_pipeline.return_value = ( + error_result + ) handler = AnalysisRequestHandler( analysis_service=analysis_service ) - with pytest.raises( - AnalysisPipelineError, - match="Analysis pipeline failed", - ): - await handler.handle(event) + actual_result = await handler.handle(event) + + assert actual_result is error_result analysis_service.analyze_pipeline.assert_awaited_once_with( event.payload.content @@ -157,7 +156,7 @@ async def test_handler_propagates_analysis_service_exception(): async def test_handler_logs_event_tracking_identifiers( caplog: pytest.LogCaptureFixture, ): - """추적 식별자 로그 테스트.""" + """추적 식별자 로그 테스트""" event = create_analysis_requested_event() expected_result = create_success_result() @@ -185,7 +184,7 @@ async def test_handler_logs_event_tracking_identifiers( async def test_handler_does_not_log_message_content( caplog: pytest.LogCaptureFixture, ): - """문자 원문을 로그에 남기지 않는지 테스트.""" + """문자 원문을 로그에 남기지 않는지 테스트""" event = create_analysis_requested_event() expected_result = create_success_result() @@ -207,7 +206,7 @@ async def test_handler_does_not_log_message_content( async def test_handler_logs_tracking_identifiers_on_failure( caplog: pytest.LogCaptureFixture, ): - """실패 로그 식별자 테스트.""" + """실패 로그 식별자 테스트""" event = create_analysis_requested_event() error_result = create_error_result() @@ -218,10 +217,10 @@ async def test_handler_logs_tracking_identifiers_on_failure( analysis_service=analysis_service ) - with caplog.at_level(logging.ERROR): - with pytest.raises(AnalysisPipelineError): - await handler.handle(event) + with caplog.at_level(logging.WARNING): + actual_result = await handler.handle(event) + assert actual_result is error_result assert str(event.eventId) in caplog.text assert str(event.analysisId) in caplog.text assert str(event.traceId) in caplog.text diff --git a/tests/infrastructure/rabbitmq/test_lifecycle.py b/tests/infrastructure/rabbitmq/test_lifecycle.py index 8266445..3e9503a 100644 --- a/tests/infrastructure/rabbitmq/test_lifecycle.py +++ b/tests/infrastructure/rabbitmq/test_lifecycle.py @@ -30,10 +30,17 @@ def test_lifespan_starts_and_stops_rabbitmq_consumer(): fake_connection.get_request_queue.return_value = ( fake_queue ) + fake_exchange = MagicMock() + fake_connection.get_exchange.return_value = ( + fake_exchange + ) fake_consumer = MagicMock() fake_consumer.start = AsyncMock() fake_consumer.stop = AsyncMock() + fake_result_publisher = MagicMock() + fake_result_factory = MagicMock() + fake_dead_letter_publisher = MagicMock() application = create_app( rabbitmq_consumer_enabled=True @@ -47,7 +54,19 @@ def test_lifespan_starts_and_stops_rabbitmq_consumer(): patch( "app.main.AnalysisRequestConsumer", return_value=fake_consumer, + ) as mock_consumer_class, + patch( + "app.main.AnalysisResultPublisher", + return_value=fake_result_publisher, + ) as mock_result_publisher_class, + patch( + "app.main.AnalysisResultEventFactory", + return_value=fake_result_factory, ), + patch( + "app.main.DeadLetterPublisher", + return_value=fake_dead_letter_publisher, + ) as mock_dead_letter_publisher_class, ): with TestClient(application) as test_client: response = test_client.get("/") @@ -56,6 +75,32 @@ def test_lifespan_starts_and_stops_rabbitmq_consumer(): fake_connection.connect.assert_awaited_once() fake_consumer.start.assert_awaited_once() + mock_result_publisher_class.assert_called_once_with( + exchange=fake_exchange, + ) + mock_dead_letter_publisher_class.assert_called_once_with( + exchange=fake_exchange, + ) + consumer_arguments = ( + mock_consumer_class.call_args.kwargs + ) + assert ( + consumer_arguments["request_queue"] + is fake_queue + ) + assert ( + consumer_arguments["result_publisher"] + is fake_result_publisher + ) + assert ( + consumer_arguments["result_factory"] + is fake_result_factory + ) + assert ( + consumer_arguments["dead_letter_publisher"] + is fake_dead_letter_publisher + ) + fake_consumer.stop.assert_awaited_once() fake_connection.close.assert_awaited_once() diff --git a/tests/infrastructure/rabbitmq/test_publisher.py b/tests/infrastructure/rabbitmq/test_publisher.py new file mode 100644 index 0000000..6afc396 --- /dev/null +++ b/tests/infrastructure/rabbitmq/test_publisher.py @@ -0,0 +1,83 @@ +from types import SimpleNamespace +from unittest.mock import AsyncMock, Mock + +import pytest + +from app.infrastructure.errors import ( + RetryableProcessingError, +) +from app.infrastructure.rabbitmq.publisher import ( + AnalysisResultPublisher, +) +from app.infrastructure.rabbitmq.schemas import ( + AnalysisEventType, +) + + +def _publisher(): + exchange = AsyncMock() + app_settings = SimpleNamespace( + RABBITMQ_ANALYSIS_COMPLETED_ROUTING_KEY=( + "analysis.completed.v1" + ), + RABBITMQ_ANALYSIS_PARTIAL_ROUTING_KEY=( + "analysis.partial.v1" + ), + RABBITMQ_ANALYSIS_FAILED_ROUTING_KEY=( + "analysis.failed.v1" + ), + RABBITMQ_PUBLISH_TIMEOUT_SECONDS=1, + ) + return ( + AnalysisResultPublisher( + exchange, + app_settings, + ), + exchange, + ) + + +def _event(): + event = Mock() + event.eventType = AnalysisEventType.COMPLETED + event.eventId = "result-event-id" + event.traceId = "trace-id" + event.schemaVersion = "1.0" + event.model_dump_json.return_value = "{}" + return event + + +@pytest.mark.asyncio +async def test_publisher_classifies_timeout_as_retryable(): + publisher, exchange = _publisher() + exchange.publish.side_effect = TimeoutError + + with pytest.raises( + RetryableProcessingError, + match="timed out", + ) as captured: + await publisher.publish(_event()) + + assert ( + captured.value.failure_code + == "RESULT_PUBLISH_TIMEOUT" + ) + + +@pytest.mark.asyncio +async def test_publisher_classifies_publish_failure_as_retryable(): + publisher, exchange = _publisher() + exchange.publish.side_effect = RuntimeError( + "unroutable" + ) + + with pytest.raises( + RetryableProcessingError, + match="Failed to publish", + ) as captured: + await publisher.publish(_event()) + + assert ( + captured.value.failure_code + == "RESULT_PUBLISH_FAILED" + ) diff --git a/tests/infrastructure/rabbitmq/test_result_factory.py b/tests/infrastructure/rabbitmq/test_result_factory.py new file mode 100644 index 0000000..53b9a31 --- /dev/null +++ b/tests/infrastructure/rabbitmq/test_result_factory.py @@ -0,0 +1,255 @@ +from datetime import datetime, timezone +from uuid import uuid4 + +import pytest + +from app.analysis.execution import ( + AnalysisExecution, + AnalysisExecutionStatus, +) +from app.analysis.schemas import ( + ContributionBreakdown, + RiskGrade, + SmishingAnalysisResponse, +) +from app.infrastructure.rabbitmq.result_factory import ( + AnalysisResultEventFactory, +) +from app.infrastructure.rabbitmq.schemas import ( + AnalysisEventType, + AnalysisRequestedEvent, + TextAnalysisMethod, +) + + +def _request() -> AnalysisRequestedEvent: + return AnalysisRequestedEvent.model_validate( + { + "schemaVersion": "1.0", + "eventId": str(uuid4()), + "analysisId": 123, + "clientMessageId": "sms-result-001", + "traceId": str(uuid4()), + "occurredAt": datetime.now(timezone.utc), + "payload": { + "sender": "1588-0000", + "content": "분석 대상 문자", + "receivedAt": datetime.now(timezone.utc), + "source": "AUTO", + }, + } + ) + + +def _result( + *, + status: str = "SUCCESS", +) -> SmishingAnalysisResponse: + return SmishingAnalysisResponse( + status=status, + message="analysis result", + final_score=82 if status == "SUCCESS" else 40, + risk_grade=( + RiskGrade.HIGH + if status == "SUCCESS" + else RiskGrade.MEDIUM + ), + contribution_breakdown=ContributionBreakdown( + llm=42, + hybrid_url=25, + rules=15, + ), + text_analysis=( + { + "result": { + "risk_score": 84, + "grade": "DANGEROUS", + "reason": "금융기관 사칭", + "evidence": ["인증 요구"], + }, + "stage1_naive_bayes": { + "risk_score": 70, + "grade": "SUSPICIOUS", + }, + } + if status == "SUCCESS" + else None + ), + url_analysis=( + { + "has_url": True, + "original_url": "https://short.example/a", + "origin_url": "https://example.test/login", + "is_url_malicious": True, + "url_risk_score": 0.9, + "engine_source": "Hybrid-Engine", + "error_message": None, + } + if status == "SUCCESS" + else None + ), + rule_analysis=( + { + "rule_score": 75, + "matched_rules": ["financial impersonation"], + "has_malicious_domain_pattern": True, + } + if status == "SUCCESS" + else None + ), + ) + + +@pytest.mark.parametrize( + ("status", "event_type", "failed_tracks"), + [ + ( + AnalysisExecutionStatus.COMPLETED, + AnalysisEventType.COMPLETED, + (), + ), + ( + AnalysisExecutionStatus.PARTIAL, + AnalysisEventType.PARTIAL, + ("URL:VIRUSTOTAL",), + ), + ], +) +def test_factory_maps_successful_execution( + status: AnalysisExecutionStatus, + event_type: AnalysisEventType, + failed_tracks: tuple[str, ...], +) -> None: + request = _request() + event = AnalysisResultEventFactory().create( + request=request, + execution=AnalysisExecution( + status=status, + result=_result(), + failed_tracks=failed_tracks, + ), + ) + + assert event.eventType == event_type + assert event.causationId == request.eventId + assert event.analysisId == request.analysisId + assert event.traceId == request.traceId + assert event.payload.finalScore == 82 + assert event.payload.riskGrade == "HIGH" + assert event.payload.rawScores.text == 84 + assert event.payload.rawScores.url == 90 + assert event.payload.rawScores.rules == 75 + assert event.payload.failedTracks == list(failed_tracks) + + +def test_factory_maps_failed_execution_without_message_content() -> None: + request = _request() + event = AnalysisResultEventFactory().create( + request=request, + execution=AnalysisExecution( + status=AnalysisExecutionStatus.FAILED, + result=_result(status="ERROR"), + failed_tracks=("PIPELINE",), + ), + ) + + assert event.eventType == AnalysisEventType.FAILED + assert event.payload.finalScore is None + assert event.payload.riskGrade is None + assert event.payload.failureCode == "PIPELINE_FAILED" + assert request.payload.content not in event.model_dump_json() + + +def test_factory_uses_deterministic_result_event_id() -> None: + request = _request() + execution = AnalysisExecution( + status=AnalysisExecutionStatus.COMPLETED, + result=_result(), + failed_tracks=(), + ) + factory = AnalysisResultEventFactory() + + first = factory.create( + request=request, + execution=execution, + ) + second = factory.create( + request=request, + execution=execution, + ) + + assert first.eventId == second.eventId + + +def test_factory_maps_machine_readable_url_error_codes() -> None: + request = _request() + result = _result() + result.url_analysis[ + "provider_error_codes" + ] = { + "VIRUSTOTAL": "RATE_LIMITED", + "GSB": "TIMEOUT", + } + + event = AnalysisResultEventFactory().create( + request=request, + execution=AnalysisExecution( + status=AnalysisExecutionStatus.PARTIAL, + result=result, + failed_tracks=( + "URL:GSB", + "URL:VIRUSTOTAL", + ), + ), + ) + + assert event.payload.urlAnalysis is not None + assert event.payload.urlAnalysis.errorCode == ( + "GSB:TIMEOUT;VIRUSTOTAL:RATE_LIMITED" + ) + + +def test_factory_identifies_gemini_only_text_analysis() -> None: + request = _request() + result = _result() + result.text_analysis = { + "engine": "gemini", + "result": { + "risk_score": 55, + "grade": "SUSPICIOUS", + "evidence": [], + }, + } + + event = AnalysisResultEventFactory().create( + request=request, + execution=AnalysisExecution( + status=AnalysisExecutionStatus.COMPLETED, + result=result, + failed_tracks=(), + ), + ) + + assert event.payload.textAnalysis is not None + assert event.payload.textAnalysis.method == ( + TextAnalysisMethod.GEMINI + ) + + +def test_factory_maps_unit_url_score_to_one_hundred() -> None: + request = _request() + result = _result() + result.url_analysis["url_risk_score"] = 1 + + event = AnalysisResultEventFactory().create( + request=request, + execution=AnalysisExecution( + status=AnalysisExecutionStatus.COMPLETED, + result=result, + failed_tracks=(), + ), + ) + + assert event.payload.rawScores.url == 100 + assert event.payload.urlAnalysis is not None + assert event.payload.urlAnalysis.score == 100 diff --git a/tests/infrastructure/rabbitmq/test_result_schemas.py b/tests/infrastructure/rabbitmq/test_result_schemas.py new file mode 100644 index 0000000..20a1594 --- /dev/null +++ b/tests/infrastructure/rabbitmq/test_result_schemas.py @@ -0,0 +1,105 @@ +from datetime import datetime, timezone +from uuid import uuid4 + +import pytest +from pydantic import ValidationError + +from app.infrastructure.rabbitmq.schemas import ( + AnalysisEventType, + AnalysisResultEvent, + AnalysisResultPayload, + RawScores, + WeightedContributions, +) + + +def test_completed_event_is_valid() -> None: + """정상 완료(COMPLETED) 이벤트 직렬화 및 직렬화 복원 테스트""" + event = AnalysisResultEvent( + schemaVersion="1.0", + eventId=uuid4(), + causationId=uuid4(), + analysisId=1, + clientMessageId="message-1", + traceId=uuid4(), + occurredAt=datetime.now(timezone.utc), + eventType=AnalysisEventType.COMPLETED, + payload=AnalysisResultPayload( + finalScore=74, + riskGrade="HIGH", + phishingType=None, + rawScores=RawScores( + text=80, + url=50, + rules=20, + ), + weightedContributions=WeightedContributions( + text=52, + url=15, + rules=7, + ), + ), + ) + + restored = AnalysisResultEvent.model_validate_json( + event.model_dump_json() + ) + + assert restored == event + assert restored.payload.phishingType is None + + +def test_partial_event_requires_failed_tracks() -> None: + """부분 성공(PARTIAL) 이벤트 시 failedTracks 누락을 검증 테스트""" + with pytest.raises(ValidationError): + AnalysisResultEvent( + schemaVersion="1.0", + eventId=uuid4(), + causationId=uuid4(), + analysisId=1, + traceId=uuid4(), + occurredAt=datetime.now(timezone.utc), + eventType=AnalysisEventType.PARTIAL, + payload=AnalysisResultPayload( + finalScore=60, + riskGrade="MEDIUM", + rawScores=RawScores(text=70), + failedTracks=[], + ), + ) + + +def test_failed_event_requires_failure_code() -> None: + """분석 실패(FAILED) 이벤트 시 failureCode 누락 검증 테스트""" + with pytest.raises(ValidationError): + AnalysisResultEvent( + schemaVersion="1.0", + eventId=uuid4(), + causationId=uuid4(), + analysisId=1, + traceId=uuid4(), + occurredAt=datetime.now(timezone.utc), + eventType=AnalysisEventType.FAILED, + payload=AnalysisResultPayload( + rawScores=RawScores(), + ), + ) + + +def test_failed_event_cannot_be_low_risk() -> None: + """분석 실패(FAILED) 이벤트에 LOW 위험 등급 설정 금지 검증 테스트""" + with pytest.raises(ValidationError): + AnalysisResultEvent( + schemaVersion="1.0", + eventId=uuid4(), + causationId=uuid4(), + analysisId=1, + traceId=uuid4(), + occurredAt=datetime.now(timezone.utc), + eventType=AnalysisEventType.FAILED, + payload=AnalysisResultPayload( + riskGrade="LOW", + rawScores=RawScores(), + failureCode="ALL_TRACKS_FAILED", + ), + ) \ No newline at end of file diff --git a/tests/infrastructure/test_http_retry.py b/tests/infrastructure/test_http_retry.py new file mode 100644 index 0000000..d9f0e45 --- /dev/null +++ b/tests/infrastructure/test_http_retry.py @@ -0,0 +1,68 @@ +from datetime import datetime, timedelta, timezone +from email.utils import format_datetime +from unittest.mock import AsyncMock, patch + +import pytest + +from app.infrastructure.http_retry import _sleep_before_retry + + +@pytest.mark.asyncio +async def test_sleep_honors_numeric_retry_after() -> None: + sleep = AsyncMock() + + with patch( + "app.infrastructure.http_retry.asyncio.sleep", + sleep, + ): + await _sleep_before_retry( + attempt=0, + retry_after="7", + ) + + sleep.assert_awaited_once_with(7.0) + + +@pytest.mark.asyncio +async def test_sleep_honors_http_date_retry_after() -> None: + retry_at = datetime.now(timezone.utc) + timedelta( + seconds=20 + ) + sleep = AsyncMock() + + with patch( + "app.infrastructure.http_retry.asyncio.sleep", + sleep, + ): + await _sleep_before_retry( + attempt=0, + retry_after=format_datetime( + retry_at, + usegmt=True, + ), + ) + + delay = sleep.await_args.args[0] + assert 18.0 <= delay <= 20.0 + + +@pytest.mark.asyncio +async def test_sleep_falls_back_for_invalid_retry_after() -> None: + sleep = AsyncMock() + + with ( + patch( + "app.infrastructure.http_retry.asyncio.sleep", + sleep, + ), + patch( + "app.infrastructure.http_retry.random.uniform", + return_value=0.0, + ), + ): + await _sleep_before_retry( + attempt=0, + retry_after="not-a-date", + ) + + sleep.assert_awaited_once_with(0.25)