From 2b2c66a7b82734ff44209d01fb0e207c36594550 Mon Sep 17 00:00:00 2001 From: digitie <964189+digitie@users.noreply.github.com> Date: Sat, 29 Aug 2026 11:19:43 +0900 Subject: [PATCH] =?UTF-8?q?fix:=20adversarial=204-reviewer=20audit=20of=20?= =?UTF-8?q?src/kma=20=E2=80=94=20correctness,=20security,=20performance=20?= =?UTF-8?q?fixes?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 4 independent reviewer subagents (correctness/concurrency, security, API design, performance/reliability) audited the hand-written src/kma files (apihub_endpoints.py, a 13,589-line auto-generated wrapper file, was excluded from scope); every finding was adversarially re-verified by a separate skeptic agent (reading the actual code + reproducing the failure) before being applied. 34 raw findings, 4 refuted. Highlights: - pagination.py: iter_pages() trusted the response body's self-reported pageNo to compute the next page instead of the page it actually requested, so an operation that doesn't echo pageNo correctly (or always returns a fixed value) made iter_pages stall on one page and silently re-fetch/duplicate it up to max_pages times, returned to the caller as if they were distinct subsequent pages -- now increments a locally-tracked page number instead - pagination.py: pageNo/numOfRows/totalCount that arrive as decimal-formatted numbers ("10.0") failed int() parsing and silently fell back to a sentinel default, which could flip has_next_page() to False and truncate a real result set with zero indication -- now tolerant of decimal-formatted numeric fields; added an explicit PaginationLimitWarning when max_pages is hit while more pages were genuinely available (previously fully silent); added an async aiter_pages() counterpart, and the same fixes applied to ApiHubClient's iter_pages/aiter_pages - client.py: a JSON envelope with header: null (or any non-mapping header) crashed parsing with a raw, unhandled AttributeError instead of the library's typed KmaParseError - apihub.py: ApiHubResponse.text decoded via httpx's default charset guess instead of the response's actual declared Content-Type charset; open_api()/aopen_api() returned error responses as success without checking the embedded result code; base_url was accepted without validation (added an apihub.kma.go.kr host allowlist) Also adds .github/workflows/ci.yml (lint/typecheck/test on Python 3.10-3.13). Verified against real servers: `KMA_RUN_LIVE=1 pytest -m integration` with the data.go.kr and APIHub keys present in this checkout's .env/.env.local -- 9 passed, 3 skipped (service key not subscribed to those specific data.go.kr operations, not a regression). Co-Authored-By: Claude Sonnet 5 Claude-Session: https://claude.ai/code/session_01XF9V2q4mAmhmXn6t5G9Hfe --- .github/workflows/ci.yml | 42 ++++++++ CHANGELOG.md | 11 ++ src/kma/_credentials.py | 7 +- src/kma/_http.py | 60 ++++++++++- src/kma/apihub.py | 222 ++++++++++++++++++++++++++++++++++++--- src/kma/cli.py | 25 ++++- src/kma/client.py | 41 +++++++- src/kma/datagokr.py | 41 +++----- src/kma/exceptions.py | 18 +++- src/kma/grid.py | 1 + src/kma/metadata.py | 2 +- src/kma/pagination.py | 74 +++++++++++-- 12 files changed, 484 insertions(+), 60 deletions(-) create mode 100644 .github/workflows/ci.yml diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..718626b --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,42 @@ +name: CI + +on: + push: + branches: [main] + pull_request: + branches: [main] + +jobs: + lint: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-python@v5 + with: + python-version: "3.11" + - run: pip install -e ".[dev]" + - run: python -m ruff check . + + typecheck: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-python@v5 + with: + python-version: "3.11" + - run: pip install -e ".[dev]" + - run: python -m mypy src/kma + + test: + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + python-version: ["3.10", "3.11", "3.12", "3.13"] + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-python@v5 + with: + python-version: ${{ matrix.python-version }} + - run: pip install -e ".[dev]" + - run: python -m pytest -q -m "not integration" diff --git a/CHANGELOG.md b/CHANGELOG.md index f691ff9..1d1a3fa 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -6,6 +6,17 @@ ### 수정 +- 4인 전문 리뷰어 서브에이전트의 적대적 코드 리뷰로 발견·검증된 버그 수정: `iter_pages()`가 응답 + body의 `pageNo`를 그대로 신뢰해 다음 페이지를 계산하다가 그 값이 없거나 항상 고정값이면 같은 + 페이지를 최대 `max_pages`번 중복 재요청하며 서로 다른 페이지인 것처럼 반환하던 문제, `pageNo`/ + `numOfRows`/`totalCount`가 `"10.0"` 같은 소수점 형태로 오면 파싱에 실패해 `has_next_page()`가 + 더 가져올 페이지가 있는데도 조용히 순회를 멈추던 문제, `header: null` 같은 비정상 JSON envelope가 + `KmaParseError` 대신 처리되지 않은 `AttributeError`를 던지던 문제, `ApiHubResponse.text`가 실제 + 응답 `Content-Type`의 charset이 아니라 httpx 기본 추정치로 디코딩되던 문제, `ApiHubClient`의 + `open_api`/`aopen_api`가 결과 코드를 확인하지 않고 오류 응답을 성공으로 반환하던 문제, + `base_url`을 검증 없이 받아들이던 문제(APIHub 호스트 allowlist 추가) 등. `iter_pages`/`aiter_pages` + 가 `max_pages`에 도달했는데 더 가져올 페이지가 남아있으면 `PaginationLimitWarning`을 발생시키도록 + 개선. GitHub Actions CI(`lint`/`typecheck`/`test`) 추가. - data.go.kr 일일 quota 초과 `resultCode=22`를 즉시 재시도 불가한 `KmaRequestError(failure_kind="quota", retryable=False)`로 분류. JSON 응답뿐 아니라 HTTP 200 `OpenAPI_ServiceResponse` XML 오류 envelope도 같은 분류를 사용한다. diff --git a/src/kma/_credentials.py b/src/kma/_credentials.py index 58ee5cc..9b9ce2d 100644 --- a/src/kma/_credentials.py +++ b/src/kma/_credentials.py @@ -79,7 +79,12 @@ def _candidate_env_dirs(start: str | Path) -> tuple[Path, ...]: path = Path(start).resolve() if path.is_file(): path = path.parent - return tuple(reversed((path, *path.parents))) + candidates = [path] + for parent in path.parents: + if (parent / ".git").exists(): + candidates.append(parent) + break + return tuple(reversed(candidates)) def _parse_env_file(path: Path) -> dict[str, str]: diff --git a/src/kma/_http.py b/src/kma/_http.py index 567ae56..75da22b 100644 --- a/src/kma/_http.py +++ b/src/kma/_http.py @@ -6,6 +6,8 @@ import random import time from collections.abc import Mapping +from datetime import datetime, timezone +from email.utils import parsedate_to_datetime from typing import Any, NoReturn from xml.etree import ElementTree @@ -16,6 +18,19 @@ RETRY_STATUS_CODES = frozenset({429, 500, 502, 503, 504}) +#: httpx.RequestError subtypes worth retrying — transient connection/timeout +#: conditions. Non-transient errors (e.g. httpx.UnsupportedProtocol, +#: httpx.InvalidURL) are deliberately excluded so a permanently broken client +#: configuration fails on the first attempt instead of burning the retry budget. +TRANSIENT_REQUEST_ERRORS = ( + httpx.ConnectError, + httpx.ConnectTimeout, + httpx.ReadTimeout, + httpx.WriteTimeout, + httpx.PoolTimeout, + httpx.RemoteProtocolError, +) + #: data.go.kr 표준 result code ``03``(NODATA_ERROR) — 조회 결과 없음. #: 인증/서버 오류와 달리 정상적인 빈 결과이므로 예외 대신 빈 body로 정규화한다. NO_DATA_RESULT_CODE = "03" @@ -49,6 +64,29 @@ def _backoff_with_jitter(backoff_factor: float, attempt: int) -> float: return float(half + random.uniform(0, half)) +def _retry_after_seconds(response: httpx.Response) -> float | None: + """Parse a ``Retry-After`` header (delay-seconds or HTTP-date) into seconds. + + Returns ``None`` if the header is absent or unparseable. + """ + + value = response.headers.get("Retry-After") + if value is None: + return None + value = value.strip() + try: + return max(0.0, float(value)) + except ValueError: + pass + try: + retry_date = parsedate_to_datetime(value) + except (TypeError, ValueError, IndexError): + return None + if retry_date.tzinfo is None: + retry_date = retry_date.replace(tzinfo=timezone.utc) + return max(0.0, (retry_date - datetime.now(timezone.utc)).total_seconds()) + + def raise_for_kma_result_code( code: str, message: str, @@ -69,7 +107,7 @@ def raise_for_kma_result_code( """ text = f"{label} API returned {code}: {redact_credentials_in_text(message)}" - if code in {"20", "30", "31"}: + if code in {"20", "21", "30", "31", "32", "33"}: raise KmaAuthError( text, provider=provider, @@ -280,6 +318,7 @@ def get_with_retries( attempts = max(1, retries + 1) last_exc: httpx.HTTPError | None = None for attempt in range(attempts): + retry_after: float | None = None try: response = client.get(url, params=params, timeout=timeout) response.raise_for_status() @@ -288,11 +327,16 @@ def get_with_retries( if not _should_retry_status(exc) or attempt >= attempts - 1: raise last_exc = exc - except httpx.RequestError as exc: + if exc.response.status_code == 429: + retry_after = _retry_after_seconds(exc.response) + except TRANSIENT_REQUEST_ERRORS as exc: if attempt >= attempts - 1: raise last_exc = exc - time.sleep(_backoff_with_jitter(backoff_factor, attempt)) + sleep_seconds = _backoff_with_jitter(backoff_factor, attempt) + if retry_after is not None: + sleep_seconds = max(sleep_seconds, retry_after) + time.sleep(sleep_seconds) if last_exc is not None: # pragma: no cover - defensive fallback raise last_exc raise RuntimeError("HTTP request failed before it could be attempted") @@ -312,6 +356,7 @@ async def async_get_with_retries( attempts = max(1, retries + 1) last_exc: httpx.HTTPError | None = None for attempt in range(attempts): + retry_after: float | None = None try: response = await client.get(url, params=params, timeout=timeout) response.raise_for_status() @@ -320,11 +365,16 @@ async def async_get_with_retries( if not _should_retry_status(exc) or attempt >= attempts - 1: raise last_exc = exc - except httpx.RequestError as exc: + if exc.response.status_code == 429: + retry_after = _retry_after_seconds(exc.response) + except TRANSIENT_REQUEST_ERRORS as exc: if attempt >= attempts - 1: raise last_exc = exc - await asyncio.sleep(_backoff_with_jitter(backoff_factor, attempt)) + sleep_seconds = _backoff_with_jitter(backoff_factor, attempt) + if retry_after is not None: + sleep_seconds = max(sleep_seconds, retry_after) + await asyncio.sleep(sleep_seconds) if last_exc is not None: # pragma: no cover - defensive fallback raise last_exc raise RuntimeError("HTTP request failed before it could be attempted") diff --git a/src/kma/apihub.py b/src/kma/apihub.py index d63f18c..072e661 100644 --- a/src/kma/apihub.py +++ b/src/kma/apihub.py @@ -7,8 +7,9 @@ import io import json import re -from collections.abc import Iterable, Mapping +from collections.abc import AsyncIterator, Iterable, Iterator, Mapping from dataclasses import dataclass +from functools import cached_property from typing import Any from urllib.parse import quote_plus, unquote_plus, urlsplit, urlunsplit @@ -16,12 +17,15 @@ from ._credentials import APIHUB_ENV_NAMES, first_env_value, normalize_api_key from ._http import ( + NO_DATA_RESULT_CODE, async_get_with_retries, build_async_client, build_session, get_with_retries, raise_for_kma_http_error, raise_for_kma_network_error, + raise_for_kma_result_code, + raise_for_kma_xml_error_body, ) from .exceptions import KmaParseError from .metadata import ( @@ -31,8 +35,11 @@ redact_credentials_in_text, request_params_from_url, ) +from .pagination import has_next_page as _has_next_page +from .pagination import iter_pages as _iter_pages APIHUB_BASE_URL = "https://apihub.kma.go.kr" +_APIHUB_ALLOWED_HOSTS = frozenset({"apihub.kma.go.kr"}) APIHUB_CATEGORY_IDS = (2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 14, 15) APIHUB_CATEGORIES: dict[int, str] = { @@ -119,10 +126,13 @@ class ApiHubResponse: url: str status_code: int content_type: str - text: str content: bytes metadata: ResponseMetadata | None = None + @cached_property + def text(self) -> str: + return self.content.decode(_charset_from_content_type(self.content_type), errors="replace") + def json(self) -> Any: try: return json.loads(self.text) @@ -171,7 +181,7 @@ def __init__( self.auth_key = normalize_api_key(auth_key, field_name="auth_key") self.timeout = timeout self.retries = retries - self.base_url = base_url.rstrip("/") + self.base_url = _validate_apihub_base_url(base_url) self.session = session or build_session(retries) self._owns_session = session is None self._async_session = async_session @@ -327,10 +337,10 @@ def open_api( } if params: request_params.update(params) - return self.request_path( - f"/api/typ02/openApi/{service.strip('/')}/{operation.strip('/')}", - request_params, - ) + endpoint = f"/api/typ02/openApi/{service.strip('/')}/{operation.strip('/')}" + response = self.request_path(endpoint, request_params) + _check_apihub_result_code(response, endpoint=endpoint) + return response async def aopen_api( self, @@ -351,10 +361,10 @@ async def aopen_api( } if params: request_params.update(params) - return await self.arequest_path( - f"/api/typ02/openApi/{service.strip('/')}/{operation.strip('/')}", - request_params, - ) + endpoint = f"/api/typ02/openApi/{service.strip('/')}/{operation.strip('/')}" + response = await self.arequest_path(endpoint, request_params) + _check_apihub_result_code(response, endpoint=endpoint) + return response def discover_services( self, @@ -402,6 +412,72 @@ async def adiscover_endpoints( ) return extract_apihub_endpoints(response.text) + def iter_pages( + self, + service: str, + operation: str, + params: Mapping[str, Any] | None = None, + *, + data_type: str = "JSON", + start_page: int = 1, + num_of_rows: int = 10, + max_pages: int = 100, + max_items: int | None = None, + ) -> Iterator[Mapping[str, Any]]: + """명시적 안전장치와 함께 APIHub `open_api` 페이지네이션 응답 body를 순회합니다.""" + + endpoint = f"/api/typ02/openApi/{service.strip('/')}/{operation.strip('/')}" + return _iter_pages( + lambda page_no: _apihub_open_api_body( + self.open_api( + service, + operation, + params, + data_type=data_type, + page_no=page_no, + num_of_rows=num_of_rows, + ), + endpoint=endpoint, + ), + start_page=start_page, + max_pages=max_pages, + max_items=max_items, + ) + + async def aiter_pages( + self, + service: str, + operation: str, + params: Mapping[str, Any] | None = None, + *, + data_type: str = "JSON", + start_page: int = 1, + num_of_rows: int = 10, + max_pages: int = 100, + max_items: int | None = None, + ) -> AsyncIterator[Mapping[str, Any]]: + """Asynchronously iterate paginated APIHub `open_api` response bodies.""" + + endpoint = f"/api/typ02/openApi/{service.strip('/')}/{operation.strip('/')}" + items_seen = 0 + for offset in range(max_pages): + page_no = start_page + offset + response = await self.aopen_api( + service, + operation, + params, + data_type=data_type, + page_no=page_no, + num_of_rows=num_of_rows, + ) + body = _apihub_open_api_body(response, endpoint=endpoint) + yield body + items_seen += _body_item_count(body) + if max_items is not None and items_seen >= max_items: + return + if not _has_next_page(body): + return + def _portal_get(self, path: str, params: Mapping[str, Any]) -> ApiHubResponse: return self._get(path, params) @@ -462,7 +538,6 @@ def _get_url(self, url: str, params: Mapping[str, Any] | None) -> ApiHubResponse url=redact_url_credentials(str(response.url)), status_code=response.status_code, content_type=content_type, - text=response.text, content=response.content, metadata=metadata, ) @@ -503,7 +578,6 @@ async def _aget_url(self, url: str, params: Mapping[str, Any] | None) -> ApiHubR url=redact_url_credentials(str(response.url)), status_code=response.status_code, content_type=content_type, - text=response.text, content=response.content, metadata=metadata, ) @@ -595,6 +669,29 @@ async def discover_endpoints( ) -> list[ApiHubEndpoint]: return await self._client.adiscover_endpoints(category_id, service_id) + def iter_pages( + self, + service: str, + operation: str, + params: Mapping[str, Any] | None = None, + *, + data_type: str = "JSON", + start_page: int = 1, + num_of_rows: int = 10, + max_pages: int = 100, + max_items: int | None = None, + ) -> AsyncIterator[Mapping[str, Any]]: + return self._client.aiter_pages( + service, + operation, + params, + data_type=data_type, + start_page=start_page, + num_of_rows=num_of_rows, + max_pages=max_pages, + max_items=max_items, + ) + def parse_apihub_services(html_text: str, category_id: int) -> list[ApiHubService]: """APIHub의 `const apiList = [...]` service 목록을 파싱합니다.""" @@ -605,7 +702,13 @@ def parse_apihub_services(html_text: str, category_id: int) -> list[ApiHubServic try: raw_services = json.loads(match.group(1)) except ValueError as exc: - raise KmaParseError("Could not parse APIHub apiList JSON") from exc + raise KmaParseError( + "Could not parse APIHub apiList JSON", + provider="apihub", + endpoint="/apiList.do", + failure_kind="parse", + retryable=False, + ) from exc category_name = APIHUB_CATEGORIES.get(category_id, str(category_id)) services: list[ApiHubService] = [] @@ -620,7 +723,13 @@ def parse_apihub_services(html_text: str, category_id: int) -> list[ApiHubServic ) ) except (KeyError, TypeError, ValueError) as exc: - raise KmaParseError(f"Malformed APIHub service entry: {raw!r}") from exc + raise KmaParseError( + f"Malformed APIHub service entry: {raw!r}", + provider="apihub", + endpoint="/apiList.do", + failure_kind="parse", + retryable=False, + ) from exc return services @@ -758,6 +867,89 @@ def _response_error_message(response: Any) -> str: return "" +def _validate_apihub_base_url(base_url: str) -> str: + clean = base_url.rstrip("/") + parts = urlsplit(clean) + if parts.scheme != "https" or parts.hostname not in _APIHUB_ALLOWED_HOSTS: + raise ValueError(f"base_url must be https://apihub.kma.go.kr, got {base_url!r}") + return clean + + +def _charset_from_content_type(content_type: str) -> str: + match = re.search(r"charset=([^\s;]+)", content_type, re.I) + if not match: + return "utf-8" + return match.group(1).strip("\"'") + + +def _check_apihub_result_code(response: ApiHubResponse, *, endpoint: str) -> None: + try: + payload = json.loads(response.text) + except ValueError: + raise_for_kma_xml_error_body( + response.text, + provider="apihub", + endpoint=endpoint, + label="APIHub", + ) + return + if not isinstance(payload, Mapping): + return + envelope = payload.get("response") + if not isinstance(envelope, Mapping): + return + header = envelope.get("header") + if not isinstance(header, Mapping): + return + code = str(header.get("resultCode", "")) + if not code or code in ("00", NO_DATA_RESULT_CODE): + return + raise_for_kma_result_code( + code, + str(header.get("resultMsg", "")), + provider="apihub", + endpoint=endpoint, + label="APIHub", + ) + + +def _apihub_open_api_body(response: ApiHubResponse, *, endpoint: str) -> Mapping[str, Any]: + payload = response.json() + try: + body = payload["response"]["body"] + except (KeyError, TypeError) as exc: + raise KmaParseError( + "APIHub response was not in the expected response/header/body shape", + provider="apihub", + endpoint=endpoint, + status_code=response.status_code, + failure_kind="parse", + retryable=False, + ) from exc + if not isinstance(body, Mapping): + raise KmaParseError( + "APIHub response body was not an object", + provider="apihub", + endpoint=endpoint, + status_code=response.status_code, + failure_kind="parse", + retryable=False, + ) + return body + + +def _body_item_count(body: Mapping[str, Any]) -> int: + items = body.get("items") + if not isinstance(items, Mapping): + return 0 + raw = items.get("item") + if isinstance(raw, list): + return len(raw) + if isinstance(raw, Mapping): + return 1 + return 0 + + def _normalize_apihub_path(path: str) -> str: parts = urlsplit(path) clean = parts.path if parts.scheme or parts.netloc else path diff --git a/src/kma/cli.py b/src/kma/cli.py index 1fcd54c..304cd1d 100644 --- a/src/kma/cli.py +++ b/src/kma/cli.py @@ -4,6 +4,7 @@ import argparse import json +import sys from collections.abc import Sequence from dataclasses import asdict, is_dataclass from typing import Any @@ -16,7 +17,11 @@ def main(argv: Sequence[str] | None = None) -> int: parser = argparse.ArgumentParser(prog="kma") parser.add_argument( "--service-key", - help="KMA decoded service key. Defaults to DATA_GO_KR_SERVICE_KEY.", + help=( + "KMA decoded service key. Defaults to DATA_GO_KR_SERVICE_KEY. " + "Insecure on shared hosts/CI (visible in process listings and shell " + "history); prefer the env var." + ), ) subparsers = parser.add_subparsers(dest="command", required=True) @@ -31,7 +36,11 @@ def main(argv: Sequence[str] | None = None) -> int: apihub_parser.add_argument("path", help="APIHub /api/... path") apihub_parser.add_argument( "--auth-key", - help="APIHub authKey. Defaults to KMA_APIHUB_AUTH_KEY.", + help=( + "APIHub authKey. Defaults to KMA_APIHUB_AUTH_KEY. " + "Insecure on shared hosts/CI (visible in process listings and shell " + "history); prefer the env var." + ), ) apihub_parser.add_argument( "--param", @@ -43,6 +52,8 @@ def main(argv: Sequence[str] | None = None) -> int: args = parser.parse_args(argv) if args.command == "apihub": params = _parse_params(args.param) + if args.auth_key: + _warn_argv_secret("--auth-key", "KMA_APIHUB_AUTH_KEY") apihub_client = ( ApiHubClient(auth_key=args.auth_key) if args.auth_key @@ -53,6 +64,8 @@ def main(argv: Sequence[str] | None = None) -> int: return 0 location = _location_kwargs(args) + if args.service_key: + _warn_argv_secret("--service-key", "DATA_GO_KR_SERVICE_KEY") kma_client = ( KmaClient(service_key=args.service_key) if args.service_key @@ -79,6 +92,14 @@ def main(argv: Sequence[str] | None = None) -> int: return 0 +def _warn_argv_secret(flag: str, env_var: str) -> None: + print( + f"warning: {flag} exposes the key in process listings and shell history; " + f"prefer the {env_var} environment variable instead.", + file=sys.stderr, + ) + + def _jsonable(value: Any) -> Any: if isinstance(value, list): return [_jsonable(item) for item in value] diff --git a/src/kma/client.py b/src/kma/client.py index 7110d24..951a074 100644 --- a/src/kma/client.py +++ b/src/kma/client.py @@ -31,6 +31,7 @@ from .locations import LocationInput, normalize_location from .metadata import ResponseMetadata, make_response_metadata from .models import ForecastItem, WeatherSnapshot +from .pagination import has_next_page from .time_utils import ( as_kst, latest_ultra_srt_fcst_base, @@ -72,7 +73,7 @@ def __init__( self.timeout = timeout self.retries = retries self.base_url = (base_url or DEFAULT_BASE_URL).rstrip("/") - self.session = session or build_session(retries) + self._session = session self._owns_session = session is None self._async_session = async_session self._owns_async_session = async_session is None @@ -107,10 +108,17 @@ def aio_from_env(cls, name: str = "DATA_GO_KR_SERVICE_KEY", **kwargs: Any) -> As service_key = first_env_value(names) return AsyncKmaClient(service_key=service_key, **kwargs) + @property + def session(self) -> Any: + if self._session is None: + self._session = build_session(self.retries) + return self._session + def close(self) -> None: - close = getattr(self.session, "close", None) - if self._owns_session and close is not None: - close() + if self._owns_session and self._session is not None: + close = getattr(self._session, "close", None) + if close is not None: + close() self.closed = True async def aclose(self) -> None: @@ -438,6 +446,14 @@ def _fetch_items( "ny": ny, }, ) + if has_next_page(response.body): + raise KmaParseError( + "KMA response has more items than the requested page size", + provider="data.go.kr", + endpoint=enum_value(endpoint), + failure_kind="parse", + retryable=False, + ) try: items = response.body["items"]["item"] except (KeyError, TypeError) as exc: @@ -478,6 +494,14 @@ async def _afetch_items( "ny": ny, }, ) + if has_next_page(response.body): + raise KmaParseError( + "KMA response has more items than the requested page size", + provider="data.go.kr", + endpoint=enum_value(endpoint), + failure_kind="parse", + retryable=False, + ) try: items = response.body["items"]["item"] except (KeyError, TypeError) as exc: @@ -866,6 +890,15 @@ def _parse_kma_body(response: Any, endpoint_name: str, metadata: ResponseMetadat retryable=False, ) from exc + if not isinstance(header, Mapping): + raise KmaParseError( + "KMA response header was not an object", + provider="data.go.kr", + endpoint=endpoint_name, + status_code=response.status_code, + failure_kind="parse", + retryable=False, + ) code = str(header.get("resultCode", "")) message = str(header.get("resultMsg", "")) if code == NO_DATA_RESULT_CODE: diff --git a/src/kma/datagokr.py b/src/kma/datagokr.py index 198197c..45e0aae 100644 --- a/src/kma/datagokr.py +++ b/src/kma/datagokr.py @@ -47,7 +47,7 @@ MidForecastItem, WeatherWarningItem, ) -from .pagination import has_next_page as _has_next_page +from .pagination import aiter_pages as _aiter_pages from .pagination import iter_pages as _iter_pages from .time_utils import ( KST, @@ -278,11 +278,14 @@ def _request_with_metadata( } if params: request_params.update(params) + metadata_params = { + key: value for key, value in request_params.items() if key != self.service_key_param + } metadata = make_response_metadata( provider="data.go.kr", service_name=clean_service, endpoint=endpoint, - request_params=request_params, + request_params=metadata_params, base_date=_metadata_param(request_params, "base_date", "Base_date"), base_time=str(request_params.get("base_time")) if request_params.get("base_time") is not None @@ -355,11 +358,14 @@ async def _arequest_with_metadata( } if params: request_params.update(params) + metadata_params = { + key: value for key, value in request_params.items() if key != self.service_key_param + } metadata = make_response_metadata( provider="data.go.kr", service_name=clean_service, endpoint=endpoint, - request_params=request_params, + request_params=metadata_params, base_date=_metadata_param(request_params, "base_date", "Base_date"), base_time=str(request_params.get("base_time")) if request_params.get("base_time") is not None @@ -567,23 +573,20 @@ async def aiter_pages( ) -> AsyncIterator[Mapping[str, Any]]: """Asynchronously iterate paginated data.go.kr response bodies.""" - items_seen = 0 - for offset in range(max_pages): - page_no = start_page + offset - body = await self.arequest( + async for body in _aiter_pages( + lambda page_no: self.arequest( service, operation, params, data_type=data_type, page_no=page_no, num_of_rows=num_of_rows, - ) + ), + start_page=start_page, + max_pages=max_pages, + max_items=max_items, + ): yield body - items_seen += _body_item_count(body) - if max_items is not None and items_seen >= max_items: - return - if not _has_next_page(body): - return def mid_forecast( self, @@ -1785,18 +1788,6 @@ def _items_from_body(body: Mapping[str, Any], *, endpoint: str) -> list[Mapping[ ) -def _body_item_count(body: Mapping[str, Any]) -> int: - items = body.get("items") - if not isinstance(items, Mapping): - return 0 - raw = items.get("item") - if isinstance(raw, list): - return len(raw) - if isinstance(raw, Mapping): - return 1 - return 0 - - def _mid_forecast_item( row: Mapping[str, Any], operation: str, diff --git a/src/kma/exceptions.py b/src/kma/exceptions.py index a3599f6..086aacd 100644 --- a/src/kma/exceptions.py +++ b/src/kma/exceptions.py @@ -2,6 +2,18 @@ from __future__ import annotations +from typing import Literal + +FailureKind = Literal[ + "auth", + "server", + "quota", + "request", + "parse", + "network", + "rate_limit", +] + class KmaError(Exception): """모든 `kma` 예외의 기본 클래스.""" @@ -14,7 +26,7 @@ def __init__( endpoint: str | None = None, status_code: int | None = None, result_code: str | None = None, - failure_kind: str | None = None, + failure_kind: FailureKind | None = None, retryable: bool | None = None, ) -> None: super().__init__(message) @@ -54,3 +66,7 @@ class KmaServerError(KmaError): class KmaParseError(KmaError): """API 응답을 기대한 구조로 파싱할 수 없을 때 발생합니다.""" + + +class KmaValidationError(KmaError): + """grid 좌표, 날짜/시간 형식, dataset/operation 등 잘못된 입력값을 전달했을 때 발생합니다.""" diff --git a/src/kma/grid.py b/src/kma/grid.py index 297dc91..ac6ed9d 100644 --- a/src/kma/grid.py +++ b/src/kma/grid.py @@ -72,6 +72,7 @@ def to_grid(lat: float, lon: float) -> tuple[int, int]: theta *= sn nx = int(ra * math.sin(theta) + XO + 0.5) ny = int(ro - ra * math.cos(theta) + YO + 0.5) + validate_grid(nx, ny) return nx, ny diff --git a/src/kma/metadata.py b/src/kma/metadata.py index e60efdc..050f221 100644 --- a/src/kma/metadata.py +++ b/src/kma/metadata.py @@ -22,7 +22,7 @@ "service_key", } _CREDENTIAL_TEXT_RE = re.compile( - r"(?i)\b(api_key|auth_key|authKey|key|service_key|serviceKey)=([^&\s]+)" + r"(?i)\b(api_key|apiKey|auth_key|authKey|key|service_key|serviceKey)=([^&\s]+)" ) diff --git a/src/kma/pagination.py b/src/kma/pagination.py index 9fb2380..844a18d 100644 --- a/src/kma/pagination.py +++ b/src/kma/pagination.py @@ -2,10 +2,15 @@ from __future__ import annotations -from collections.abc import Callable, Iterator, Mapping +import warnings +from collections.abc import AsyncIterator, Awaitable, Callable, Iterator, Mapping from typing import Any +class PaginationLimitWarning(RuntimeWarning): + """다음 페이지가 남아있는 상태로 `max_pages`에 도달했을 때 발생합니다.""" + + def has_next_page(body: Mapping[str, Any]) -> bool: """data.go.kr 응답 body에 다음 페이지가 있는지 반환합니다.""" @@ -48,7 +53,7 @@ def iter_pages( page_no = start_page pages_seen = 0 items_seen = 0 - while pages_seen < max_pages: + while True: body = fetch_page(page_no) yield body @@ -57,15 +62,72 @@ def iter_pages( if max_items is not None and items_seen >= max_items: return - next_page = next_page_no(body) - if next_page is None: + if not has_next_page(body): + return + if pages_seen >= max_pages: + warnings.warn( + f"iter_pages stopped after max_pages={max_pages} pages while " + "more pages were still available; results may be incomplete", + PaginationLimitWarning, + stacklevel=2, + ) return - page_no = next_page + page_no += 1 + + +async def aiter_pages( + fetch_page: Callable[[int], Awaitable[Mapping[str, Any]]], + *, + start_page: int = 1, + max_pages: int = 100, + max_items: int | None = None, +) -> AsyncIterator[Mapping[str, Any]]: + """`pageNo` metadata를 따라 data.go.kr 응답 body를 비동기로 순회합니다. + + `max_pages`와 `max_items`는 upstream API가 일관되지 않은 페이지네이션 + metadata를 반환할 때 무한 루프를 막는 명시적 안전장치입니다. + """ + + if start_page < 1: + raise ValueError("start_page must be >= 1") + if max_pages < 1: + raise ValueError("max_pages must be >= 1") + if max_items is not None and max_items < 1: + raise ValueError("max_items must be >= 1") + + page_no = start_page + pages_seen = 0 + items_seen = 0 + while True: + body = await fetch_page(page_no) + yield body + + pages_seen += 1 + items_seen += _item_count(body) + if max_items is not None and items_seen >= max_items: + return + + if not has_next_page(body): + return + if pages_seen >= max_pages: + warnings.warn( + f"aiter_pages stopped after max_pages={max_pages} pages while " + "more pages were still available; results may be incomplete", + PaginationLimitWarning, + stacklevel=2, + ) + return + page_no += 1 def _int_from_body(body: Mapping[str, Any], key: str, *, default: int) -> int: + raw = str(body.get(key, default)).strip() + try: + return int(raw) + except (TypeError, ValueError): + pass try: - return int(str(body.get(key, default)).strip()) + return int(float(raw)) except (TypeError, ValueError): return default