From c563528a2d35ec274e025d795802bc9da283c980 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=9E=E6=B3=95=E6=93=8D=E4=BD=9C?= Date: Wed, 26 Aug 2026 06:32:13 +0000 Subject: [PATCH] refactor(api): standardize compliance downloads (#41057) --- api/.importlinter | 2 +- api/controllers/console/billing/compliance.py | 102 ++++++++----- api/controllers/console/billing/error.py | 15 ++ api/controllers/console/error.py | 6 - api/extensions/ext_application_services.py | 11 ++ api/openapi/markdown/console-openapi.md | 14 +- api/services/billing_service.py | 43 +++--- api/services/compliance_download_service.py | 50 +++++++ api/services/errors/billing.py | 4 + .../console/billing/test_compliance.py | 92 ++++++++++++ .../controllers/console/test_wraps.py | 18 +++ .../unit_tests/controllers/test_swagger.py | 18 ++- .../test_ext_application_services.py | 45 ++++++ .../services/test_billing_service.py | 130 ++++++----------- .../test_compliance_download_service.py | 136 ++++++++++++++++++ .../api/console/compliance/types.gen.ts | 36 ++++- .../api/console/compliance/zod.gen.ts | 40 +++++- 17 files changed, 611 insertions(+), 151 deletions(-) create mode 100644 api/services/compliance_download_service.py create mode 100644 api/tests/unit_tests/controllers/console/billing/test_compliance.py create mode 100644 api/tests/unit_tests/services/test_compliance_download_service.py diff --git a/api/.importlinter b/api/.importlinter index 9d77c891864..227d512372f 100644 --- a/api/.importlinter +++ b/api/.importlinter @@ -124,7 +124,6 @@ ignore_imports = services.account_service -> controllers services.account_service -> controllers.console.error services.app_generate_service -> controllers.console.app.workflow - services.billing_service -> controllers.console.error [importlinter:contract:no-direct-rsa-imports] # Note: `libs` itself is deliberately excluded from source_modules -- import-linter's @@ -417,6 +416,7 @@ name = Billing application services and ports are framework and infrastructure n type = forbidden source_modules = services.billing_portal_service + services.compliance_download_service services.partner_tenant_binding_service forbidden_modules = configs diff --git a/api/controllers/console/billing/compliance.py b/api/controllers/console/billing/compliance.py index f62d6f8e417..b3e30fef993 100644 --- a/api/controllers/console/billing/compliance.py +++ b/api/controllers/console/billing/compliance.py @@ -1,40 +1,47 @@ -from typing import Any - from flask import request from flask_restx import Resource -from pydantic import BaseModel, Field, RootModel +from pydantic import BaseModel, Field -from controllers.common.schema import query_params_from_model, register_response_schema_models -from libs.helper import extract_remote_ip -from libs.login import login_required -from models import Account -from services.billing_service import BillingService - -from ...common.schema import DEFAULT_REF_TEMPLATE_OPENAPI_3_0 -from .. import console_ns -from ..wraps import ( - account_initialization_required, - model_validate, - only_edition_cloud, - setup_required, - with_current_tenant_id, - with_current_user, +from controllers.common.schema import ( + query_params_from_model, + register_response_schema_models, + register_schema_models, ) +from controllers.console import console_ns +from controllers.console.billing.error import ( + BillingOperationFailedErrorResponse, + BillingUnavailableErrorResponse, + BillingUnprocessableEntityErrorResponse, + ComplianceRateLimitErrorResponse, + to_billing_request_error, +) +from controllers.console.flask_admission import console_account_admission +from controllers.console.wraps import model_validate +from enums import DeploymentEdition +from extensions.ext_application_services import application_services +from fields.base import ResponseModel +from libs.helper import dump_response, extract_remote_ip +from machinery.context import RequestContext +from services.errors.billing import BillingError class ComplianceDownloadQuery(BaseModel): doc_name: str = Field(..., description="Compliance document name") -class ComplianceDownloadResponse(RootModel[dict[str, Any]]): - root: dict[str, Any] +class ComplianceDownloadResponse(ResponseModel): + url: str -console_ns.schema_model( - ComplianceDownloadQuery.__name__, - ComplianceDownloadQuery.model_json_schema(ref_template=DEFAULT_REF_TEMPLATE_OPENAPI_3_0), +register_schema_models(console_ns, ComplianceDownloadQuery) +register_response_schema_models( + console_ns, + BillingOperationFailedErrorResponse, + BillingUnavailableErrorResponse, + BillingUnprocessableEntityErrorResponse, + ComplianceDownloadResponse, + ComplianceRateLimitErrorResponse, ) -register_response_schema_models(console_ns, ComplianceDownloadResponse) @console_ns.route("/compliance/download") @@ -43,21 +50,38 @@ class ComplianceApi(Resource): @console_ns.doc("download_compliance_document") @console_ns.doc(description="Get compliance document download link") @console_ns.response(200, "Success", console_ns.models[ComplianceDownloadResponse.__name__]) - @setup_required - @login_required - @account_initialization_required - @only_edition_cloud - @with_current_user - @with_current_tenant_id + @console_ns.response( + 422, + "Invalid compliance download query", + console_ns.models[BillingUnprocessableEntityErrorResponse.__name__], + ) + @console_ns.response( + 429, + "Compliance download rate limit exceeded", + console_ns.models[ComplianceRateLimitErrorResponse.__name__], + ) + @console_ns.response( + 502, + "Compliance download failed", + console_ns.models[BillingOperationFailedErrorResponse.__name__], + ) + @console_ns.response( + 503, + "Billing unavailable", + console_ns.models[BillingUnavailableErrorResponse.__name__], + ) + @console_account_admission(editions=frozenset({DeploymentEdition.CLOUD})) @model_validate(ComplianceDownloadQuery) - def get(self, req_data: ComplianceDownloadQuery, current_tenant_id: str, current_user: Account): - + def get(self, req_data: ComplianceDownloadQuery, request_context: RequestContext): ip_address = extract_remote_ip(request) device_info = request.headers.get("User-Agent", "Unknown device") - return BillingService.get_compliance_download_link( - doc_name=req_data.doc_name, - account_id=current_user.id, - tenant_id=current_tenant_id, - ip=ip_address, - device_info=device_info, - ) + try: + data = application_services().compliance_downloads.get_link( + request_context=request_context, + document_name=req_data.doc_name, + ip_address=ip_address, + device_info=device_info, + ) + except BillingError as error: + raise to_billing_request_error(error) from error + return dump_response(ComplianceDownloadResponse, data) diff --git a/api/controllers/console/billing/error.py b/api/controllers/console/billing/error.py index 6d35a8ec9f4..9c3ca678c72 100644 --- a/api/controllers/console/billing/error.py +++ b/api/controllers/console/billing/error.py @@ -6,6 +6,7 @@ from services.errors.billing import ( BillingError, BillingUpstreamInvalidResponseError, BillingUpstreamUnavailableError, + ComplianceRateLimitExceededError, ) @@ -15,6 +16,12 @@ class BillingUnprocessableEntityErrorResponse(ResponseModel): status: Literal[422] +class ComplianceRateLimitErrorResponse(ResponseModel): + code: Literal["compliance_rate_limit"] + message: str + status: Literal[429] + + class BillingOperationFailedErrorResponse(ResponseModel): code: Literal["billing_operation_failed"] message: str @@ -27,6 +34,12 @@ class BillingUnavailableErrorResponse(ResponseModel): status: Literal[503] +class ComplianceRateLimitError(BaseHTTPException): + error_code = "compliance_rate_limit" + description = "Rate limit exceeded for downloading compliance report." + code = 429 + + class BillingOperationFailedError(BaseHTTPException): error_code = "billing_operation_failed" description = "We couldn't complete this request. Please try again. If the problem persists, contact support." @@ -40,6 +53,8 @@ class BillingUnavailableError(BaseHTTPException): def to_billing_request_error(error: BillingError) -> BaseHTTPException: + if isinstance(error, ComplianceRateLimitExceededError): + return ComplianceRateLimitError() if isinstance(error, BillingUpstreamInvalidResponseError): return BillingOperationFailedError() if isinstance(error, BillingUpstreamUnavailableError): diff --git a/api/controllers/console/error.py b/api/controllers/console/error.py index 0ecb552bdb9..0de7c2df2f5 100644 --- a/api/controllers/console/error.py +++ b/api/controllers/console/error.py @@ -113,9 +113,3 @@ class EducationActivateLimitError(BaseHTTPException): error_code = "education_activate_limit" description = "Rate limit exceeded" code = 429 - - -class ComplianceRateLimitError(BaseHTTPException): - error_code = "compliance_rate_limit" - description = "Rate limit exceeded for downloading compliance report." - code = 429 diff --git a/api/extensions/ext_application_services.py b/api/extensions/ext_application_services.py index 0d243ae3836..f1cbb99808a 100644 --- a/api/extensions/ext_application_services.py +++ b/api/extensions/ext_application_services.py @@ -77,6 +77,7 @@ from services.auth.data_source_api_key_auth_gateways import ( from services.auth.data_source_api_key_auth_service import DataSourceApiKeyAuthService from services.billing_portal_service import BillingPortalService from services.billing_service import BillingService +from services.compliance_download_service import ComplianceDownloadService from services.enterprise.enterprise_service import EnterpriseService from services.errors.enterprise import EnterpriseServiceError from services.explore_banner_query_service import ExploreBannerQueryService @@ -148,6 +149,7 @@ class ApplicationServices: app_definitions: AppDefinitionQueryService app_sites: AppSiteService billing_portal: BillingPortalService + compliance_downloads: ComplianceDownloadService data_source_api_key_auth: DataSourceApiKeyAuthService webapp_access: WebAppAccessQueryService web_app_runtime: WebAppRuntimeQueryService @@ -286,6 +288,15 @@ def build_application_services( get_subscription=BillingService.get_subscription, get_invoices=BillingService.get_invoices, ), + compliance_downloads=ComplianceDownloadService( + fetch_link=BillingService.get_compliance_download_link, + rate_limiter=RateLimiter( + prefix="compliance_download_rate_limiter", + max_attempts=4, + time_window=60, + redis_client=redis, + ), + ), data_source_api_key_auth=DataSourceApiKeyAuthService( bindings=data_source_api_key_auth_bindings, validator=ProviderApiKeyAuthCredentialValidator(), diff --git a/api/openapi/markdown/console-openapi.md b/api/openapi/markdown/console-openapi.md index ec9e8f70f3f..25621593ea2 100644 --- a/api/openapi/markdown/console-openapi.md +++ b/api/openapi/markdown/console-openapi.md @@ -5071,6 +5071,10 @@ Get compliance document download link | Code | Description | Schema | | ---- | ----------- | ------ | | 200 | Success | **application/json**: [ComplianceDownloadResponse](#compliancedownloadresponse)
| +| 422 | Invalid compliance download query | **application/json**: [BillingUnprocessableEntityErrorResponse](#billingunprocessableentityerrorresponse)
| +| 429 | Compliance download rate limit exceeded | **application/json**: [ComplianceRateLimitErrorResponse](#complianceratelimiterrorresponse)
| +| 502 | Compliance download failed | **application/json**: [BillingOperationFailedErrorResponse](#billingoperationfailederrorresponse)
| +| 503 | Billing unavailable | **application/json**: [BillingUnavailableErrorResponse](#billingunavailableerrorresponse)
| ### [GET] /data-source/integrates #### Responses @@ -16219,7 +16223,15 @@ TEAM: Team collaboration paid plan | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| ComplianceDownloadResponse | object | | | +| url | string | | Yes | + +#### ComplianceRateLimitErrorResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| code | string | | Yes | +| message | string | | Yes | +| status | integer | | Yes | #### ComposerBindingPayload diff --git a/api/services/billing_service.py b/api/services/billing_service.py index 8edc33a581d..f89547f4558 100644 --- a/api/services/billing_service.py +++ b/api/services/billing_service.py @@ -13,8 +13,8 @@ from werkzeug.exceptions import InternalServerError from core.helper.http_client_pooling import get_pooled_http_client from enums import CloudPlan from extensions.ext_redis import redis_client -from libs.helper import RateLimiter from services.billing_portal_service import BillingPortalLink +from services.compliance_download_service import ComplianceDownloadLink from services.errors.billing import ( BillingUpstreamInvalidResponseError, BillingUpstreamUnavailableError, @@ -71,6 +71,9 @@ class EducationAutocompleteResponseDict(TypedDict): _billing_portal_link_adapter = TypeAdapter(BillingPortalLink) +_compliance_download_link_adapter = TypeAdapter(ComplianceDownloadLink) + + class QuotaReserveResult(TypedDict): reservation_id: str available: int @@ -229,8 +232,6 @@ class BillingService: quota_base_url = os.environ.get("BILLING_QUOTA_API_URL") or base_url secret_key = os.environ.get("BILLING_API_SECRET_KEY", "BILLING_API_SECRET_KEY") - compliance_download_rate_limiter = RateLimiter("compliance_download_rate_limiter", 4, 60) - # Redis key prefix for tenant plan cache _PLAN_CACHE_KEY_PREFIX = "tenant_plan:" # Cache TTL: 10 minutes @@ -466,7 +467,6 @@ class BillingService: base_url: str | None = None, ) -> Any: headers = {"Content-Type": "application/json", "Billing-Api-Secret-Key": cls.secret_key} - url = f"{base_url or cls.base_url}{endpoint}" response = _http_client.request(method, url, json=json, params=params, headers=headers, follow_redirects=True) if method == "GET" and response.status_code != httpx.codes.OK: @@ -482,7 +482,10 @@ class BillingService: if response.status_code != httpx.codes.OK: raise ValueError("Invalid arguments.") if method == "POST" and response.status_code != httpx.codes.OK: - raise ValueError(f"Unable to send request to {url}. Please try again later or contact support.") + raise _BillingHTTPStatusError( + f"Unable to send request to {url}. Please try again later or contact support.", + response.status_code, + ) if method == "DELETE" and response.status_code != httpx.codes.OK: logger.error("billing_service: DELETE response: %s %s", response.status_code, response.text) raise ValueError(f"Unable to process delete request {url}. Please try again later or contact support.") @@ -587,23 +590,31 @@ class BillingService: tenant_id: str, ip: str, device_info: str, - ): - limiter_key = f"{account_id}:{tenant_id}" - if cls.compliance_download_rate_limiter.is_rate_limited(limiter_key): - from controllers.console.error import ComplianceRateLimitError - - raise ComplianceRateLimitError() - - json = { + ) -> ComplianceDownloadLink: + payload = { "doc_name": doc_name, "account_id": account_id, "tenant_id": tenant_id, "ip_address": ip, "device_info": device_info, } - res = cls._send_request("POST", "/compliance/download", json=json) - cls.compliance_download_rate_limiter.increment_rate_limit(limiter_key) - return res + try: + response = cls._send_request("POST", "/compliance/download", json=payload) + result = _compliance_download_link_adapter.validate_python(response) + except _BillingHTTPStatusError as error: + if error.status_code in {httpx.codes.REQUEST_TIMEOUT, httpx.codes.TOO_MANY_REQUESTS} or ( + error.status_code >= 500 + ): + raise BillingUpstreamUnavailableError from error + raise BillingUpstreamInvalidResponseError from error + except httpx.RequestError as error: + raise BillingUpstreamUnavailableError from error + except (json.JSONDecodeError, UnicodeDecodeError, ValidationError) as error: + raise BillingUpstreamInvalidResponseError from error + except ValueError as error: + raise RuntimeError("Unexpected billing service value error") from error + + return result @classmethod def clean_billing_info_cache(cls, tenant_id: str) -> None: diff --git a/api/services/compliance_download_service.py b/api/services/compliance_download_service.py new file mode 100644 index 00000000000..d379d24c8b8 --- /dev/null +++ b/api/services/compliance_download_service.py @@ -0,0 +1,50 @@ +"""Application service for compliance document downloads.""" + +from collections.abc import Callable +from typing import Protocol, TypedDict + +from machinery.context import RequestContext +from machinery.errors import ActiveWorkspaceRequiredError +from services.errors.billing import ComplianceRateLimitExceededError + + +class ComplianceDownloadLink(TypedDict): + url: str + + +class ComplianceDownloadRateLimiter(Protocol): + def is_rate_limited(self, key: str, /) -> bool: ... + + def increment_rate_limit(self, key: str, /) -> None: ... + + +class ComplianceDownloadService: + def __init__( + self, + *, + fetch_link: Callable[[str, str, str, str, str], ComplianceDownloadLink], + rate_limiter: ComplianceDownloadRateLimiter, + ) -> None: + self._fetch_link = fetch_link + self._rate_limiter = rate_limiter + + def get_link( + self, + *, + request_context: RequestContext, + document_name: str, + ip_address: str, + device_info: str, + ) -> ComplianceDownloadLink: + workspace_id = request_context.active_workspace_id + if workspace_id is None: + raise ActiveWorkspaceRequiredError + + account_id = request_context.account_id + limiter_key = f"{account_id}:{workspace_id}" + if self._rate_limiter.is_rate_limited(limiter_key): + raise ComplianceRateLimitExceededError + + link = self._fetch_link(document_name, account_id, workspace_id, ip_address, device_info) + self._rate_limiter.increment_rate_limit(limiter_key) + return link diff --git a/api/services/errors/billing.py b/api/services/errors/billing.py index aa86d754f46..cbbd96fa803 100644 --- a/api/services/errors/billing.py +++ b/api/services/errors/billing.py @@ -2,6 +2,10 @@ class BillingError(Exception): pass +class ComplianceRateLimitExceededError(BillingError): + pass + + class BillingUpstreamInvalidResponseError(BillingError): pass diff --git a/api/tests/unit_tests/controllers/console/billing/test_compliance.py b/api/tests/unit_tests/controllers/console/billing/test_compliance.py new file mode 100644 index 00000000000..fc39909bddb --- /dev/null +++ b/api/tests/unit_tests/controllers/console/billing/test_compliance.py @@ -0,0 +1,92 @@ +from collections.abc import Iterator +from inspect import unwrap +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest +from flask import Flask + +from controllers.console.billing.compliance import ComplianceApi, ComplianceDownloadQuery +from controllers.console.billing.error import ComplianceRateLimitError +from machinery.context import RequestContext +from services.errors.billing import ComplianceRateLimitExceededError + + +@pytest.fixture +def app() -> Flask: + app = Flask(__name__) + app.config["TESTING"] = True + return app + + +@pytest.fixture +def request_context() -> RequestContext: + return RequestContext( + request_id="request-1", + trace_id="trace-1", + account_id="account-1", + active_workspace_id="tenant-1", + ) + + +@pytest.fixture +def compliance_downloads() -> MagicMock: + return MagicMock() + + +@pytest.fixture(autouse=True) +def mock_application_services(compliance_downloads: MagicMock) -> Iterator[None]: + with patch( + "controllers.console.billing.compliance.application_services", + return_value=SimpleNamespace(compliance_downloads=compliance_downloads), + ): + yield + + +def test_get_compliance_download_link( + app: Flask, + request_context: RequestContext, + compliance_downloads: MagicMock, +) -> None: + resource = ComplianceApi() + method = unwrap(resource.get) + query = ComplianceDownloadQuery(doc_name="SOC2_Type_II") + compliance_downloads.get_link.return_value = {"url": "https://example.com/report", "ignored": True} + + with ( + app.test_request_context("/compliance/download", headers={"User-Agent": "test-agent"}), + patch("controllers.console.billing.compliance.extract_remote_ip", return_value="127.0.0.1"), + ): + result = method(resource, query, request_context) + + assert result == {"url": "https://example.com/report"} + compliance_downloads.get_link.assert_called_once_with( + request_context=request_context, + document_name="SOC2_Type_II", + ip_address="127.0.0.1", + device_info="test-agent", + ) + + +def test_get_compliance_download_translates_rate_limit( + app: Flask, + request_context: RequestContext, + compliance_downloads: MagicMock, +) -> None: + resource = ComplianceApi() + method = unwrap(resource.get) + query = ComplianceDownloadQuery(doc_name="SOC2_Type_II") + compliance_downloads.get_link.side_effect = ComplianceRateLimitExceededError + + with ( + app.test_request_context("/compliance/download"), + patch("controllers.console.billing.compliance.extract_remote_ip", return_value="127.0.0.1"), + ): + with pytest.raises(ComplianceRateLimitError) as exc_info: + method(resource, query, request_context) + + assert exc_info.value.data == { + "code": "compliance_rate_limit", + "message": "Rate limit exceeded for downloading compliance report.", + "status": 429, + } diff --git a/api/tests/unit_tests/controllers/console/test_wraps.py b/api/tests/unit_tests/controllers/console/test_wraps.py index 16773c0ac6c..a9688112153 100644 --- a/api/tests/unit_tests/controllers/console/test_wraps.py +++ b/api/tests/unit_tests/controllers/console/test_wraps.py @@ -221,6 +221,24 @@ class TestCurrentContextInjection: assert admission_context.active_workspace_id == "tenant-123" assert route_value == "route-value" + def test_console_account_admission_enforces_declared_edition_first(self): + class Handler: + @flask_admission.console_account_admission(editions=frozenset({DeploymentEdition.CLOUD})) + def get(self, request_context: RequestContext): + return request_context + + with ( + patch( + "controllers.console.flask_admission.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.COMMUNITY, + ), + Flask(__name__).test_request_context(), + pytest.raises(HTTPException) as exc_info, + ): + Handler().get() + + assert exc_info.value.code == 404 + def test_console_account_admission_enforces_legacy_workspace_roles(self): current_user = make_account() current_user.role = TenantAccountRole.NORMAL diff --git a/api/tests/unit_tests/controllers/test_swagger.py b/api/tests/unit_tests/controllers/test_swagger.py index e36d95825ac..af2e3a1396f 100644 --- a/api/tests/unit_tests/controllers/test_swagger.py +++ b/api/tests/unit_tests/controllers/test_swagger.py @@ -648,6 +648,12 @@ def test_console_billing_routes_document_error_responses(monkeypatch: pytest.Mon "502": "BillingOperationFailedErrorResponse", "503": "BillingUnavailableErrorResponse", }, + ("/compliance/download", "get"): { + "422": "BillingUnprocessableEntityErrorResponse", + "429": "ComplianceRateLimitErrorResponse", + "502": "BillingOperationFailedErrorResponse", + "503": "BillingUnavailableErrorResponse", + }, } for (path, method), responses in expected_responses.items(): @@ -656,12 +662,14 @@ def test_console_billing_routes_document_error_responses(monkeypatch: pytest.Mon schema = operation["responses"][status]["content"]["application/json"]["schema"] assert schema["$ref"] == f"#/components/schemas/{model_name}" - forbidden_response = operation["responses"]["403"] - assert forbidden_response["description"] == "Forbidden" - assert "content" not in forbidden_response + if path.startswith("/billing/"): + forbidden_response = operation["responses"]["403"] + assert forbidden_response["description"] == "Forbidden" + assert "content" not in forbidden_response expected_error_contracts = { "BillingUnprocessableEntityErrorResponse": ("unprocessable_entity", 422), + "ComplianceRateLimitErrorResponse": ("compliance_rate_limit", 429), "BillingOperationFailedErrorResponse": ("billing_operation_failed", 502), "BillingUnavailableErrorResponse": ("billing_unavailable", 503), } @@ -671,6 +679,10 @@ def test_console_billing_routes_document_error_responses(monkeypatch: pytest.Mon assert properties["code"]["const"] == error_code assert properties["status"]["const"] == status + compliance_response = schemas["ComplianceDownloadResponse"] + assert set(compliance_response["properties"]) == {"url"} + assert compliance_response["required"] == ["url"] + def test_console_model_provider_checkout_route_is_deprecated(monkeypatch: pytest.MonkeyPatch): from configs import dify_config diff --git a/api/tests/unit_tests/extensions/test_ext_application_services.py b/api/tests/unit_tests/extensions/test_ext_application_services.py index 3b7d840d0ca..b8de66eb541 100644 --- a/api/tests/unit_tests/extensions/test_ext_application_services.py +++ b/api/tests/unit_tests/extensions/test_ext_application_services.py @@ -34,6 +34,7 @@ from services.app_site_service import AppSiteService from services.auth.data_source_api_key_auth_service import DataSourceApiKeyAuthService from services.billing_portal_service import BillingPortalService from services.billing_service import BillingService +from services.compliance_download_service import ComplianceDownloadService from services.enterprise.enterprise_service import WebAppSettings from services.errors.enterprise import EnterpriseAPIError, EnterpriseAPINotFoundError from services.init_validation_service import InvalidInitializationPasswordError @@ -250,6 +251,50 @@ def test_build_application_services_wires_billing_service( sync_partner_tenants_bindings.assert_called_once_with("account-1", "partner-key", "click-1") +def test_build_application_services_wires_compliance_downloads( + sqlite_session_factory: sessionmaker[Session], +) -> None: + redis = MagicMock(spec=RedisClientWrapper) + with ( + patch.object( + BillingService, + "get_compliance_download_link", + return_value={"url": "https://billing.example.com/compliance"}, + ) as fetch_link, + patch("extensions.ext_application_services.RateLimiter") as rate_limiter_type, + ): + rate_limiter = rate_limiter_type.return_value + rate_limiter.is_rate_limited.return_value = False + services = ext_application_services.build_application_services( + database_client=sqlite_session_factory, + deployment_edition=DeploymentEdition.COMMUNITY, + initialization_password="", + redis=redis, + ) + + assert isinstance(services.compliance_downloads, ComplianceDownloadService) + assert services.compliance_downloads.get_link( + request_context=RequestContext( + request_id="request-1", + trace_id="trace-1", + account_id="account-1", + active_workspace_id="workspace-1", + ), + document_name="SOC2_Type_II", + ip_address="127.0.0.1", + device_info="test-agent", + ) == {"url": "https://billing.example.com/compliance"} + rate_limiter_type.assert_any_call( + prefix="compliance_download_rate_limiter", + max_attempts=4, + time_window=60, + redis_client=redis, + ) + rate_limiter.is_rate_limited.assert_called_once_with("account-1:workspace-1") + rate_limiter.increment_rate_limit.assert_called_once_with("account-1:workspace-1") + fetch_link.assert_called_once_with("SOC2_Type_II", "account-1", "workspace-1", "127.0.0.1", "test-agent") + + def test_build_application_services_wires_education_rate_limiters( sqlite_session_factory: sessionmaker[Session], ) -> None: diff --git a/api/tests/unit_tests/services/test_billing_service.py b/api/tests/unit_tests/services/test_billing_service.py index 47b5c62ebcf..cf587aed3fc 100644 --- a/api/tests/unit_tests/services/test_billing_service.py +++ b/api/tests/unit_tests/services/test_billing_service.py @@ -4,7 +4,7 @@ This test module covers all aspects of the billing service including: - HTTP request handling with retry logic - Subscription tier management and billing information retrieval - Usage calculation and credit management (positive/negative deltas) -- Rate limit enforcement for compliance downloads and education features +- Compliance and education billing-provider requests - Account management and permission checks - Cache management for billing data - Partner integration features @@ -222,9 +222,10 @@ class TestBillingServiceSendRequest: mock_httpx_request.return_value = mock_response # Act & Assert - with pytest.raises(ValueError) as exc_info: + with pytest.raises(_BillingHTTPStatusError) as exc_info: BillingService._send_request("POST", "/test", json={"key": "value"}) assert "Unable to send request to" in str(exc_info.value) + assert exc_info.value.status_code == status_code @pytest.mark.parametrize( "status_code", [httpx.codes.BAD_REQUEST, httpx.codes.INTERNAL_SERVER_ERROR, httpx.codes.NOT_FOUND] @@ -264,10 +265,11 @@ class TestBillingServiceSendRequest: mock_httpx_request.return_value = mock_response # Act & Assert - # POST checks status code before calling response.json(), so ValueError is raised - with pytest.raises(ValueError) as exc_info: + # POST checks status code before calling response.json(). + with pytest.raises(_BillingHTTPStatusError) as exc_info: BillingService._send_request("POST", "/test", json={"key": "value"}) assert "Unable to send request to" in str(exc_info.value) + assert exc_info.value.status_code == status_code @pytest.mark.parametrize( "status_code", [httpx.codes.BAD_REQUEST, httpx.codes.INTERNAL_SERVER_ERROR, httpx.codes.NOT_FOUND] @@ -1028,8 +1030,8 @@ class TestBillingServiceQuotaOperations: assert result["api_rate_limit"]["limit"] == -1 -class TestBillingServiceRateLimitEnforcement: - """Unit tests for compliance download rate-limit enforcement.""" +class TestBillingServiceComplianceDownload: + """Unit tests for compliance download requests.""" @pytest.fixture def mock_send_request(self): @@ -1037,67 +1039,56 @@ class TestBillingServiceRateLimitEnforcement: with patch.object(BillingService, "_send_request") as mock: yield mock - def test_compliance_download_rate_limiter_not_limited(self, mock_send_request): - """Test compliance download when rate limit is not exceeded.""" - # Arrange + def test_compliance_download_returns_validated_link(self, mock_send_request): doc_name = "compliance_report.pdf" account_id = "account-123" tenant_id = "tenant-456" ip = "192.168.1.1" device_info = "Mozilla/5.0" - expected_response = {"download_link": "https://example.com/download"} + expected_response = {"url": "https://example.com/download", "ignored": True} + mock_send_request.return_value = expected_response - # Mock the rate limiter to return False (not limited) - with ( - patch.object( - BillingService.compliance_download_rate_limiter, "is_rate_limited", return_value=False - ) as mock_is_limited, - patch.object(BillingService.compliance_download_rate_limiter, "increment_rate_limit") as mock_increment, - ): - mock_send_request.return_value = expected_response + result = BillingService.get_compliance_download_link(doc_name, account_id, tenant_id, ip, device_info) - # Act - result = BillingService.get_compliance_download_link(doc_name, account_id, tenant_id, ip, device_info) + assert result == {"url": "https://example.com/download"} + mock_send_request.assert_called_once_with( + "POST", + "/compliance/download", + json={ + "doc_name": doc_name, + "account_id": account_id, + "tenant_id": tenant_id, + "ip_address": ip, + "device_info": device_info, + }, + ) - # Assert - assert result == expected_response - mock_is_limited.assert_called_once_with(f"{account_id}:{tenant_id}") - mock_send_request.assert_called_once_with( - "POST", - "/compliance/download", - json={ - "doc_name": doc_name, - "account_id": account_id, - "tenant_id": tenant_id, - "ip_address": ip, - "device_info": device_info, - }, - ) - # Verify rate limit was incremented after successful download - mock_increment.assert_called_once_with(f"{account_id}:{tenant_id}") + @pytest.mark.parametrize( + ("request_error", "error_type"), + [ + (_BillingHTTPStatusError("request failed", httpx.codes.BAD_REQUEST), BillingUpstreamInvalidResponseError), + (_BillingHTTPStatusError("request failed", httpx.codes.REQUEST_TIMEOUT), BillingUpstreamUnavailableError), + (_BillingHTTPStatusError("request failed", httpx.codes.TOO_MANY_REQUESTS), BillingUpstreamUnavailableError), + ( + _BillingHTTPStatusError("request failed", httpx.codes.INTERNAL_SERVER_ERROR), + BillingUpstreamUnavailableError, + ), + (httpx.RequestError("request failed"), BillingUpstreamUnavailableError), + ], + ) + def test_compliance_download_maps_request_failure( + self, mock_send_request, request_error: Exception, error_type: type[Exception] + ) -> None: + mock_send_request.side_effect = request_error - def test_compliance_download_rate_limiter_exceeded(self, mock_send_request): - """Test compliance download when rate limit is exceeded.""" - # Arrange - doc_name = "compliance_report.pdf" - account_id = "account-123" - tenant_id = "tenant-456" - ip = "192.168.1.1" - device_info = "Mozilla/5.0" + with pytest.raises(error_type): + BillingService.get_compliance_download_link("SOC2_Type_II", "account-1", "tenant-1", "127.0.0.1", "test") - # Import the error class to properly catch it - from controllers.console.error import ComplianceRateLimitError + def test_compliance_download_rejects_invalid_response(self, mock_send_request) -> None: + mock_send_request.return_value = {} - # Mock the rate limiter to return True (rate limited) - with patch.object( - BillingService.compliance_download_rate_limiter, "is_rate_limited", return_value=True - ) as mock_is_limited: - # Act & Assert - with pytest.raises(ComplianceRateLimitError): - BillingService.get_compliance_download_link(doc_name, account_id, tenant_id, ip, device_info) - - mock_is_limited.assert_called_once_with(f"{account_id}:{tenant_id}") - mock_send_request.assert_not_called() + with pytest.raises(BillingUpstreamInvalidResponseError): + BillingService.get_compliance_download_link("SOC2_Type_II", "account-1", "tenant-1", "127.0.0.1", "test") class TestBillingServiceEducationIdentity: @@ -1815,33 +1806,6 @@ class TestBillingServiceIntegrationScenarios: assert updated_usage["used"] == 0 assert updated_usage["remaining"] == 100 - def test_compliance_download_multiple_requests_within_limit(self, mock_send_request): - """Test multiple compliance downloads within rate limit.""" - # Arrange - account_id = "account-compliance" - tenant_id = "tenant-compliance" - doc_name = "compliance_report.pdf" - ip = "192.168.1.1" - device_info = "Mozilla/5.0" - - # Mock rate limiter to allow 3 requests (under limit of 4) - with ( - patch.object( - BillingService.compliance_download_rate_limiter, "is_rate_limited", side_effect=[False, False, False] - ) as mock_is_limited, - patch.object(BillingService.compliance_download_rate_limiter, "increment_rate_limit") as mock_increment, - ): - mock_send_request.return_value = {"download_link": "https://example.com/download"} - - # Act - Make 3 requests - for i in range(3): - result = BillingService.get_compliance_download_link(doc_name, account_id, tenant_id, ip, device_info) - assert "download_link" in result - - # Assert - All 3 requests succeeded - assert mock_is_limited.call_count == 3 - assert mock_increment.call_count == 3 - class TestBillingServiceSubscriptionInfoDataType: """Unit tests for data type coercion in BillingService.get_info diff --git a/api/tests/unit_tests/services/test_compliance_download_service.py b/api/tests/unit_tests/services/test_compliance_download_service.py new file mode 100644 index 00000000000..c803bb84e8b --- /dev/null +++ b/api/tests/unit_tests/services/test_compliance_download_service.py @@ -0,0 +1,136 @@ +from unittest.mock import MagicMock + +import pytest + +from machinery.context import RequestContext +from machinery.errors import ActiveWorkspaceRequiredError +from services.compliance_download_service import ComplianceDownloadRateLimiter, ComplianceDownloadService +from services.errors.billing import BillingUpstreamUnavailableError, ComplianceRateLimitExceededError + + +@pytest.fixture +def fetch_link() -> MagicMock: + return MagicMock() + + +@pytest.fixture +def rate_limiter() -> MagicMock: + limiter = MagicMock(spec=ComplianceDownloadRateLimiter) + limiter.is_rate_limited.return_value = False + return limiter + + +@pytest.fixture +def request_context() -> RequestContext: + return RequestContext( + request_id="request-1", + trace_id="trace-1", + account_id="account-1", + active_workspace_id="workspace-1", + ) + + +@pytest.fixture +def service( + fetch_link: MagicMock, + rate_limiter: MagicMock, +) -> ComplianceDownloadService: + return ComplianceDownloadService( + fetch_link=fetch_link, + rate_limiter=rate_limiter, + ) + + +def test_get_link_checks_limit_fetches_and_increments( + service: ComplianceDownloadService, + request_context: RequestContext, + fetch_link: MagicMock, + rate_limiter: MagicMock, +) -> None: + events: list[str] = [] + rate_limiter.is_rate_limited.side_effect = lambda _key: events.append("check") or False + fetch_link.side_effect = lambda *_args: events.append("fetch") or {"url": "https://example.com/report"} + rate_limiter.increment_rate_limit.side_effect = lambda _key: events.append("increment") + + result = service.get_link( + request_context=request_context, + document_name="SOC2_Type_II", + ip_address="127.0.0.1", + device_info="test-agent", + ) + + assert result == {"url": "https://example.com/report"} + assert events == ["check", "fetch", "increment"] + rate_limiter.is_rate_limited.assert_called_once_with("account-1:workspace-1") + fetch_link.assert_called_once_with( + "SOC2_Type_II", + "account-1", + "workspace-1", + "127.0.0.1", + "test-agent", + ) + rate_limiter.increment_rate_limit.assert_called_once_with("account-1:workspace-1") + + +def test_get_link_rejects_rate_limited_request( + service: ComplianceDownloadService, + request_context: RequestContext, + fetch_link: MagicMock, + rate_limiter: MagicMock, +) -> None: + rate_limiter.is_rate_limited.return_value = True + + with pytest.raises(ComplianceRateLimitExceededError): + service.get_link( + request_context=request_context, + document_name="SOC2_Type_II", + ip_address="127.0.0.1", + device_info="test-agent", + ) + + fetch_link.assert_not_called() + rate_limiter.increment_rate_limit.assert_not_called() + + +def test_get_link_does_not_increment_after_fetch_failure( + service: ComplianceDownloadService, + request_context: RequestContext, + fetch_link: MagicMock, + rate_limiter: MagicMock, +) -> None: + fetch_link.side_effect = BillingUpstreamUnavailableError + + with pytest.raises(BillingUpstreamUnavailableError): + service.get_link( + request_context=request_context, + document_name="SOC2_Type_II", + ip_address="127.0.0.1", + device_info="test-agent", + ) + + rate_limiter.increment_rate_limit.assert_not_called() + + +def test_get_link_requires_active_workspace( + service: ComplianceDownloadService, + fetch_link: MagicMock, + rate_limiter: MagicMock, +) -> None: + request_context = RequestContext( + request_id="request-1", + trace_id="trace-1", + account_id="account-1", + active_workspace_id=None, + ) + + with pytest.raises(ActiveWorkspaceRequiredError): + service.get_link( + request_context=request_context, + document_name="SOC2_Type_II", + ip_address="127.0.0.1", + device_info="test-agent", + ) + + rate_limiter.is_rate_limited.assert_not_called() + fetch_link.assert_not_called() + rate_limiter.increment_rate_limit.assert_not_called() diff --git a/packages/contracts/generated/api/console/compliance/types.gen.ts b/packages/contracts/generated/api/console/compliance/types.gen.ts index 6a01752eadd..a72362fc15f 100644 --- a/packages/contracts/generated/api/console/compliance/types.gen.ts +++ b/packages/contracts/generated/api/console/compliance/types.gen.ts @@ -5,7 +5,31 @@ export type ClientOptions = { } export type ComplianceDownloadResponse = { - [key: string]: unknown + url: string +} + +export type BillingUnprocessableEntityErrorResponse = { + code: 'unprocessable_entity' + message: string + status: 422 +} + +export type ComplianceRateLimitErrorResponse = { + code: 'compliance_rate_limit' + message: string + status: 429 +} + +export type BillingOperationFailedErrorResponse = { + code: 'billing_operation_failed' + message: string + status: 502 +} + +export type BillingUnavailableErrorResponse = { + code: 'billing_unavailable' + message: string + status: 503 } export type GetComplianceDownloadData = { @@ -17,6 +41,16 @@ export type GetComplianceDownloadData = { url: '/compliance/download' } +export type GetComplianceDownloadErrors = { + 422: BillingUnprocessableEntityErrorResponse + 429: ComplianceRateLimitErrorResponse + 502: BillingOperationFailedErrorResponse + 503: BillingUnavailableErrorResponse +} + +export type GetComplianceDownloadError = + GetComplianceDownloadErrors[keyof GetComplianceDownloadErrors] + export type GetComplianceDownloadResponses = { 200: ComplianceDownloadResponse } diff --git a/packages/contracts/generated/api/console/compliance/zod.gen.ts b/packages/contracts/generated/api/console/compliance/zod.gen.ts index 167051fee89..7e68a451d6c 100644 --- a/packages/contracts/generated/api/console/compliance/zod.gen.ts +++ b/packages/contracts/generated/api/console/compliance/zod.gen.ts @@ -5,7 +5,45 @@ import * as z from 'zod' /** * ComplianceDownloadResponse */ -export const zComplianceDownloadResponse = z.record(z.string(), z.unknown()) +export const zComplianceDownloadResponse = z.object({ + url: z.string(), +}) + +/** + * BillingUnprocessableEntityErrorResponse + */ +export const zBillingUnprocessableEntityErrorResponse = z.object({ + code: z.literal('unprocessable_entity'), + message: z.string(), + status: z.literal(422), +}) + +/** + * ComplianceRateLimitErrorResponse + */ +export const zComplianceRateLimitErrorResponse = z.object({ + code: z.literal('compliance_rate_limit'), + message: z.string(), + status: z.literal(429), +}) + +/** + * BillingOperationFailedErrorResponse + */ +export const zBillingOperationFailedErrorResponse = z.object({ + code: z.literal('billing_operation_failed'), + message: z.string(), + status: z.literal(502), +}) + +/** + * BillingUnavailableErrorResponse + */ +export const zBillingUnavailableErrorResponse = z.object({ + code: z.literal('billing_unavailable'), + message: z.string(), + status: z.literal(503), +}) export const zGetComplianceDownloadQuery = z.object({ doc_name: z.string(),