| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293 |
- from __future__ import annotations
- import copy
- import json
- import traceback
- from pathlib import Path
- from types import MappingProxyType
- import pytest
- from sqlalchemy import event, func, select
- from dms.config import (
- DEFAULT_UI_DICTIONARY_CONFIG_PATH,
- load_dms_config,
- )
- from dms.services.ui_dictionary_service import (
- DICTIONARY_ENUMS,
- UiDictionaryConfigError,
- load_ui_dictionary_config,
- )
- def _default_payload() -> dict:
- return json.loads(
- DEFAULT_UI_DICTIONARY_CONFIG_PATH.read_text(encoding="utf-8")
- )
- def _write_payload(path: Path, payload: dict) -> Path:
- path.write_text(
- json.dumps(payload, ensure_ascii=False),
- encoding="utf-8",
- )
- return path
- def test_default_ui_dictionary_loads_all_enum_codes_and_chinese_labels():
- config = load_ui_dictionary_config(DEFAULT_UI_DICTIONARY_CONFIG_PATH)
- assert config.version == "2026.1"
- assert config.locale == "zh-CN"
- assert isinstance(config.dictionaries, MappingProxyType)
- assert set(config.dictionaries) == set(DICTIONARY_ENUMS)
- for dictionary_name, enum_type in DICTIONARY_ENUMS.items():
- options = config.dictionaries[dictionary_name]
- assert {option.code for option in options} == {
- member.value for member in enum_type
- }
- assert list(options) == sorted(
- options, key=lambda option: (option.sort_no, option.code)
- )
- assert all(option.label for option in options)
- labels = {
- option.code: option.label
- for option in config.dictionaries["securityLevels"]
- }
- assert labels["INTERNAL"] == "内部"
- assert {
- option.code: option.label
- for option in config.dictionaries["visibilityTypes"]
- }["ALL_AUTHENTICATED"] == "全体已登录用户"
- assert {
- option.code: option.label
- for option in config.dictionaries["documentStatuses"]
- }["PUBLISHED"] == "已发布"
- def test_external_ui_dictionary_path_is_loaded_once(
- monkeypatch, tmp_path: Path
- ):
- payload = _default_payload()
- payload["dictionaries"]["securityLevels"][1]["label"] = "内部资料"
- path = _write_payload(tmp_path / "external.json", payload)
- monkeypatch.setenv("DMS_UI_DICTIONARY_CONFIG_PATH", str(path))
- environment_config = load_dms_config()
- loaded = load_ui_dictionary_config(
- environment_config["DMS_UI_DICTIONARY_CONFIG_PATH"]
- )
- path.write_text("{broken after startup", encoding="utf-8")
- response_data = loaded.to_dict()
- labels = {
- item["code"]: item["label"]
- for item in response_data["dictionaries"]["securityLevels"]
- }
- assert labels["INTERNAL"] == "内部资料"
- assert "DMS_UI_DICTIONARY_CONFIG_PATH" not in response_data
- def test_loaded_configuration_is_immutable():
- config = load_ui_dictionary_config(DEFAULT_UI_DICTIONARY_CONFIG_PATH)
- with pytest.raises(TypeError):
- config.dictionaries["extra"] = () # type: ignore[index]
- with pytest.raises(Exception):
- config.dictionaries["roleCodes"][0].label = "变化" # type: ignore[misc]
- @pytest.mark.parametrize(
- "mutation",
- [
- "missing_dictionary",
- "duplicate_code",
- "missing_code",
- "unknown_code",
- "empty_label",
- "negative_sort",
- "boolean_sort",
- "invalid_version",
- "invalid_locale",
- ],
- )
- def test_invalid_ui_dictionary_config_is_rejected_without_path_leak(
- mutation: str, tmp_path: Path
- ):
- payload = copy.deepcopy(_default_payload())
- roles = payload["dictionaries"]["roleCodes"]
- if mutation == "missing_dictionary":
- del payload["dictionaries"]["auditTargets"]
- elif mutation == "duplicate_code":
- roles.append(copy.deepcopy(roles[0]))
- elif mutation == "missing_code":
- roles.pop()
- elif mutation == "unknown_code":
- roles[0]["code"] = "UNKNOWN_ROLE"
- elif mutation == "empty_label":
- roles[0]["label"] = " "
- elif mutation == "negative_sort":
- roles[0]["sortNo"] = -1
- elif mutation == "boolean_sort":
- roles[0]["sortNo"] = True
- elif mutation == "invalid_version":
- payload["version"] = "latest"
- elif mutation == "invalid_locale":
- payload["locale"] = "../zh-CN"
- path = _write_payload(tmp_path / "sensitive-name.json", payload)
- with pytest.raises(UiDictionaryConfigError) as captured:
- load_ui_dictionary_config(path)
- assert str(path) not in str(captured.value)
- assert str(tmp_path) not in str(captured.value)
- def test_missing_malformed_and_non_utf8_config_are_rejected(tmp_path: Path):
- paths = [
- tmp_path / "missing.json",
- tmp_path / "malformed.json",
- tmp_path / "non-utf8.json",
- ]
- paths[1].write_text("{", encoding="utf-8")
- paths[2].write_bytes(b"\xff\xfe\x00")
- for path in paths:
- with pytest.raises(UiDictionaryConfigError) as captured:
- load_ui_dictionary_config(path)
- assert str(tmp_path) not in str(captured.value)
- rendered = "".join(
- traceback.format_exception(
- type(captured.value),
- captured.value,
- captured.value.__traceback__,
- )
- )
- assert str(tmp_path) not in rendered
- def test_application_startup_fails_for_invalid_external_config(
- monkeypatch, tmp_path: Path
- ):
- from flask import Flask
- from dms import init_dms
- path = tmp_path / "invalid.json"
- path.write_text("{", encoding="utf-8")
- monkeypatch.setenv("DMS_UI_DICTIONARY_CONFIG_PATH", str(path))
- with pytest.raises(UiDictionaryConfigError):
- init_dms(Flask("invalid-ui-dictionary"))
- def test_ui_dictionary_requires_valid_token(b2_client):
- missing = b2_client.get("/api/v1/config/ui-dictionaries")
- invalid = b2_client.get(
- "/api/v1/config/ui-dictionaries",
- headers={"Authorization": "Bearer invalid-token"},
- )
- assert missing.status_code == 401
- assert invalid.status_code == 401
- for response, expected_code in (
- (missing, "TOKEN_INVALID"),
- (invalid, "TOKEN_INVALID"),
- ):
- payload = response.get_json()
- assert payload["code"] == expected_code
- assert payload["requestId"] == response.headers["X-Request-Id"]
- def test_all_roles_receive_stable_dictionary_without_audit_or_business_query(
- b2_app, b2_client, token_for
- ):
- from dms.extensions import db
- from dms.models import AuditLog
- tokens = {
- username: token_for(username)
- for username in ("admin", "user")
- }
- with b2_app.app_context():
- audit_count_before = db.session.scalar(
- select(func.count()).select_from(AuditLog)
- )
- statements: list[str] = []
- def record_statement(
- _connection,
- _cursor,
- statement,
- _parameters,
- _context,
- _executemany,
- ):
- statements.append(statement.lower())
- event.listen(db.engine, "before_cursor_execute", record_statement)
- try:
- responses = [
- b2_client.get(
- "/api/v1/config/ui-dictionaries",
- headers={
- "Authorization": f"Bearer {token}",
- "Origin": "http://127.0.0.1:9346",
- },
- )
- for token in tokens.values()
- ]
- finally:
- event.remove(
- db.engine, "before_cursor_execute", record_statement
- )
- audit_count_after = db.session.scalar(
- select(func.count()).select_from(AuditLog)
- )
- assert all(response.status_code == 200 for response in responses)
- payloads = [response.get_json() for response in responses]
- assert all(
- payload["data"] == payloads[0]["data"] for payload in payloads[1:]
- )
- for response, payload in zip(responses, payloads):
- assert payload["requestId"] == response.headers["X-Request-Id"]
- assert "X-Request-Id" in response.headers[
- "Access-Control-Expose-Headers"
- ]
- assert response.headers["Cache-Control"] == "private, no-store"
- assert "configPath" not in payload["data"]
- assert "loadedAt" not in payload["data"]
- assert audit_count_after == audit_count_before
- assert len(statements) == len(tokens)
- assert all(
- statement.lstrip().startswith("select") and "sys_user" in statement
- for statement in statements
- )
- def test_openapi_ui_dictionary_contract_and_no_write_operation():
- import yaml
- backend_root = Path(__file__).resolve().parents[2]
- openapi_root = backend_root / "openapi"
- document = yaml.safe_load(
- (openapi_root / "openapi.yaml").read_text(encoding="utf-8")
- )
- path_item = yaml.safe_load(
- (openapi_root / "paths" / "configuration.yaml").read_text(
- encoding="utf-8"
- )
- )["UiDictionaries"]
- schemas = yaml.safe_load(
- (openapi_root / "components" / "schemas.yaml").read_text(
- encoding="utf-8"
- )
- )
- assert document["info"]["version"] == "1.5.0"
- assert set(path_item) == {"get"}
- assert path_item["get"]["operationId"] == "getUiDictionaries"
- assert {"200", "401", "500"} <= set(path_item["get"]["responses"])
- assert set(schemas["UiDictionaryGroups"]["required"]) == set(
- DICTIONARY_ENUMS
- )
|