Commit f12d2992 by 沈彬

feat: 独立算力服务初始提交

账户/凭证/赠送/模型调用(OpenAI 协议非流式+SSE)/精确日志结算与补偿的独立 Python 服务(FastAPI + SQLAlchemy 2),含 DDL 与对账 SQL、审计缺陷修复与全量测试(SQLite + MySQL 8.0.43)。
parents
__pycache__/
*.py[cod]
.pytest_cache/
*.egg-info/
.venv/
build/
dist/
.env
.DS_Store
.idea/
.claude/
.codex/
.qoder/
# 阶段 1 单阶段构建:python:3.12-slim + 项目依赖(design.md §5)。
FROM python:3.12-slim
ENV PYTHONUNBUFFERED=1 \
PYTHONDONTWRITEBYTECODE=1 \
PIP_NO_CACHE_DIR=1
WORKDIR /app
COPY pyproject.toml ./
COPY app ./app
RUN pip install --no-cache-dir .
# 独立账户服务入口(app.account_main);旧 app.main 仅保留迁移回归,不再作为容器入口。
# 配置文件必须位于项目目录外,由部署方只读挂载到约定路径(JSON 含数据库与网关 PAT,权限 0600)。
# 结算 worker 用同一镜像另起容器:python -m app.account_worker --config <同路径>。
EXPOSE 8000
# 应用只监听容器内私网端口;公网流量经负载均衡/反向代理进入(design.md §12)。
CMD ["python", "-m", "app.account_main", "--host", "0.0.0.0", "--port", "8000", "--config", "/etc/mei1-computing/account.json"]
This diff is collapsed. Click to expand it.
"""美一算力账户独立服务。
阶段进度与对应文档:mei1-computing-service/docs/design.md、
mei1-saas/docs/notes/2026-09-21-python-computing-service-migration.md。
当前实现范围:迁移方案阶段 1(项目基础、AuthContext 预留、员工/门店归属解析、
只读查询),鉴权 AUTH_MODE=disabled,RSA 提供器在阶段 4 实现。
"""
import hashlib
import json
import re
from sqlalchemy import select
from sqlalchemy.exc import IntegrityError
from .account_models import Client, Credential, OperationAudit
from .account_security import AccountError, issue_secret, utc_now, validate_request_id
def bootstrap_client(session_factory, generator, *, client_id, name, category, request_id, deliver_secret):
validate_request_id(request_id)
if not isinstance(client_id, str) or not re.fullmatch(r"[a-z0-9-]{2,50}", client_id):
raise AccountError("INVALID_ARGUMENT")
if not isinstance(name, str) or not name.strip() or len(name) > 100 or category not in {"INTEGRATION", "PLATFORM"}:
raise AccountError("INVALID_ARGUMENT")
digest = hashlib.sha256(json.dumps({"name": name, "category": category}, sort_keys=True).encode()).hexdigest()
credential_id, audit_id = generator.next_id(), generator.next_id()
for attempt in range(2):
inserting_client = False
try:
with session_factory.begin() as session:
client = session.get(Client, client_id, with_for_update=True)
if client is None:
inserting_client = True
session.add(Client(client_id=client_id, name=name))
session.flush()
inserting_client = False
elif client.status != "ACTIVE":
raise AccountError("PERMISSION_DENIED")
elif client.name != name:
raise AccountError("INVALID_ARGUMENT")
previous = session.scalar(select(OperationAudit).where(
OperationAudit.client_id == client_id,
OperationAudit.action == "BOOTSTRAP_MANAGEMENT",
OperationAudit.request_id == request_id,
))
if previous is not None:
if previous.evidence_ref.get("fingerprint") != digest:
raise AccountError("FINGERPRINT_MISMATCH")
row = session.get(Credential, previous.target_id)
return {"credentialId": str(row.id), "status": row.status, "secretAvailable": False}
secret, secret_digest, mask = issue_secret(credential_id)
session.add(Credential(
id=credential_id, category=category, client_id=client_id,
secret_digest=secret_digest, secret_mask=mask, status="ACTIVE", activated_at=utc_now(),
))
session.flush()
session.add(OperationAudit(
id=audit_id, client_id=client_id, action="BOOTSTRAP_MANAGEMENT",
target_type="CREDENTIAL", target_id=credential_id, request_id=request_id,
reason="受控初始化管理凭证", to_status="ACTIVE",
evidence_ref={"fingerprint": digest, "category": category},
))
session.flush()
deliver_secret({"credentialId": str(credential_id), "clientId": client_id,
"category": category, "secret": secret})
return {"credentialId": str(credential_id), "status": "ACTIVE", "secretAvailable": False}
except IntegrityError:
if not inserting_client or attempt:
raise AccountError("SERVICE_UNAVAILABLE") from None
raise AccountError("SERVICE_UNAVAILABLE")
"""External, fail-closed settings for the independent account service."""
import hashlib
import ipaddress
import json
import math
import re
from dataclasses import dataclass, field, fields
from pathlib import Path
from typing import Union
from urllib.parse import unquote, urlsplit, urlunsplit
_PROJECT_ROOT = Path(__file__).resolve().parents[1]
_MAX_CONFIG_BYTES = 1024 * 1024
_MODEL_PATTERN = re.compile(r"[A-Za-z0-9][A-Za-z0-9._:/+\-]{0,99}\Z")
_PAT_PATTERN = re.compile(r"[A-Za-z0-9._\-]+\Z")
class AccountConfigError(RuntimeError):
def __init__(self):
super().__init__("Invalid account configuration")
def _valid_model(value):
return isinstance(value, str) and _MODEL_PATTERN.fullmatch(value) is not None
def _url_parts(value):
if (not isinstance(value, str) or not value
or any(char.isspace() or ord(char) < 32 or ord(char) == 127 for char in value)
or "\\" in value):
raise ValueError
parts = urlsplit(value)
# Accessing port also validates malformed/out-of-range port values.
if parts.port is not None and parts.port <= 0:
raise ValueError
return parts
def _canonical_base_url(value):
parts = _url_parts(value)
if (parts.scheme != "https" or not parts.hostname
or parts.username is not None or parts.password is not None
or "?" in value or "#" in value or parts.netloc.endswith(":")):
raise ValueError
host = parts.hostname.encode("idna").decode("ascii").lower()
if ":" in host:
if "%" in host:
raise ValueError
host = "[" + ipaddress.IPv6Address(host).compressed + "]"
else:
host = host.rstrip(".")
if len(host) > 253 or any(
re.fullmatch(r"[a-z0-9](?:[a-z0-9-]{0,61}[a-z0-9])?", label) is None
for label in host.split(".")):
raise ValueError
# Avoid aliases that HTTP clients or reverse proxies can normalize differently.
if any(unquote(segment) in (".", "..") for segment in parts.path.split("/")):
raise ValueError
if parts.port not in (None, 443):
host += ":" + str(parts.port)
return urlunsplit(("https", host, parts.path.rstrip("/"), "", ""))
def _unique_object(pairs):
result = {}
for key, value in pairs:
if key in result:
raise ValueError
result[key] = value
return result
def _reject_constant(value):
raise ValueError
@dataclass(frozen=True, repr=False)
class AccountSettings:
database_url: str = field(repr=False)
newapi_base_url: str = field(repr=False)
newapi_pat: str = field(repr=False)
newapi_model: str = field(repr=False)
db_pool_size: int = 5
db_max_overflow: int = 0
management_timeout_seconds: float = 30.0
management_concurrency: int = 8
connect_timeout_seconds: float = 10.0
read_timeout_seconds: float = 30.0
@property
def gateway_identity(self) -> str:
try:
canonical = _canonical_base_url(self.newapi_base_url)
return hashlib.sha256(canonical.encode("utf-8")).hexdigest()
except Exception:
raise AccountConfigError() from None
def validate(self) -> None:
try:
database = _url_parts(self.database_url)
schema = unquote(database.path[1:])
if (not self.database_url.startswith("mysql+pymysql://") or not database.path.startswith("/")
or not schema.strip() or schema in (".", "..") or "/" in schema
or "\\" in schema or any(ord(char) < 32 or ord(char) == 127 for char in schema)
or "#" in self.database_url):
raise ValueError
_canonical_base_url(self.newapi_base_url)
if (not isinstance(self.newapi_pat, str)
or _PAT_PATTERN.fullmatch(self.newapi_pat) is None
or not _valid_model(self.newapi_model)):
raise ValueError
if type(self.db_pool_size) is not int or self.db_pool_size <= 0:
raise ValueError
if type(self.db_max_overflow) is not int or self.db_max_overflow < 0:
raise ValueError
if (type(self.management_concurrency) is not int
or not 0 < self.management_concurrency <= 100):
raise ValueError
for value, maximum in (
(self.management_timeout_seconds, 150),
(self.connect_timeout_seconds, 10),
(self.read_timeout_seconds, 120),
):
if (type(value) not in (int, float) or not math.isfinite(value)
or not 0 < value <= maximum):
raise ValueError
except Exception:
raise AccountConfigError() from None
def load_account_settings(path: Union[Path, str]) -> AccountSettings:
"""Read only the explicit external file; never consult environment variables."""
try:
resolved = Path(path).resolve(strict=True)
if resolved == _PROJECT_ROOT or _PROJECT_ROOT in resolved.parents or not resolved.is_file():
raise ValueError
with resolved.open("rb") as source:
content = source.read(_MAX_CONFIG_BYTES + 1)
if len(content) > _MAX_CONFIG_BYTES:
raise ValueError
values = json.loads(content.decode("utf-8"), object_pairs_hook=_unique_object,
parse_constant=_reject_constant)
if not isinstance(values, dict) or set(values) - {item.name for item in fields(AccountSettings)}:
raise ValueError
settings = AccountSettings(**values)
settings.validate()
return settings
except Exception:
raise AccountConfigError() from None
import asyncio
import time
from contextvars import ContextVar
from starlette.concurrency import run_in_threadpool
from .account_security import AccountError
database_deadline = ContextVar("account_database_deadline", default=None)
def check_database_deadline():
deadline = database_deadline.get()
if deadline is not None and time.monotonic() >= deadline:
# asyncio.TimeoutError unifies with builtin TimeoutError on 3.12+ and
# is caught by the same handlers that treat deadline overrun as timeout.
raise asyncio.TimeoutError("Database phase deadline exceeded")
async def run_database(accounts, function, *args, deadline=None, **kwargs):
if not accounts.database_slots.acquire(blocking=False):
raise AccountError("RATE_LIMITED")
def execute():
token = database_deadline.set(deadline)
try:
check_database_deadline()
return function(*args, **kwargs)
finally:
database_deadline.reset(token)
accounts.database_slots.release()
task = asyncio.create_task(run_in_threadpool(execute))
task.add_done_callback(lambda completed: None if completed.cancelled() else completed.exception())
# Keep the slot until the physical transaction ends, even after HTTP cancellation.
return await asyncio.shield(task)
from sqlalchemy import func, select
from .account_models import Account, AccountClient, Call, Consumption, Gift, ManagementRequest
def get_account(session, account_id, client_id):
return session.scalar(
select(Account)
.join(AccountClient, AccountClient.account_id == Account.id)
.where(
Account.id == account_id,
AccountClient.client_id == client_id,
AccountClient.status == "AUTHORIZED",
)
)
def lock_account(session, account_id):
return session.scalar(
select(Account)
.where(Account.id == account_id)
.with_for_update()
.execution_options(populate_existing=True)
)
def find_call(session, account_id, client_id, request_id):
return session.scalar(
select(Call).where(
Call.account_id == account_id,
Call.client_id == client_id,
Call.request_id == request_id,
)
)
def find_gift(session, account_id, client_id, request_id):
return session.scalar(
select(Gift).where(
Gift.account_id == account_id,
Gift.client_id == client_id,
Gift.request_id == request_id,
)
)
def find_management_request(session, client_id, operation_type, request_id):
return session.scalar(
select(ManagementRequest).where(
ManagementRequest.client_id == client_id,
ManagementRequest.operation_type == operation_type,
ManagementRequest.request_id == request_id,
)
)
def count_unsettled_calls(session, account_id):
return session.scalar(
select(func.count()).select_from(Call).where(
Call.account_id == account_id,
Call.billing_status.in_(
("PROCESSING", "SETTLE_PENDING", "UNKNOWN", "SETTLE_FAILED")
),
)
)
def count_uncertain_calls(session, account_id):
return session.scalar(
select(func.count()).select_from(Call).where(
Call.account_id == account_id,
Call.billing_status.in_(("UNKNOWN", "SETTLE_FAILED")),
)
)
def page_consumptions(session, account_id, client_id, *, page=1, size=20, start=None, end=None):
if page < 1 or not 1 <= size <= 200:
raise ValueError("分页范围无效")
conditions = [Consumption.account_id == account_id, Consumption.client_id == client_id]
if start is not None:
conditions.append(Consumption.create_time >= start)
if end is not None:
conditions.append(Consumption.create_time < end)
total = session.scalar(select(func.count()).select_from(Consumption).where(*conditions))
rows = session.scalars(
select(Consumption)
.where(*conditions)
.order_by(Consumption.create_time.desc(), Consumption.id.desc())
.offset((page - 1) * size)
.limit(size)
).all()
return total, list(rows)
from datetime import datetime, timezone
from typing import Annotated, List, Literal, Optional
from pydantic import AfterValidator, BaseModel, ConfigDict, Field, StringConstraints, field_validator, model_validator
def _id_range(value):
if int(value) > 2**64 - 1:
raise ValueError("ID out of range")
return value
def _camel(value):
first, *rest = value.split("_")
return first + "".join(word.title() for word in rest)
Identifier = Annotated[str, StringConstraints(strict=True, pattern=r"^[1-9][0-9]{0,19}$"), AfterValidator(_id_range)]
ClientId = Annotated[str, StringConstraints(strict=True, pattern=r"^[a-z0-9-]{2,50}$")]
RequestId = Annotated[str, StringConstraints(strict=True, pattern=r"^[!-~]{8,100}$")]
class AccountRequest(BaseModel):
model_config = ConfigDict(alias_generator=_camel, populate_by_name=True, extra="forbid", strict=True)
class WriteRequest(AccountRequest):
request_id: Optional[RequestId] = None
@field_validator("reason", check_fields=False)
@classmethod
def require_reason(cls, value):
if not value.strip():
raise ValueError("Blank reason")
return value
class CreateAccountRequest(WriteRequest):
name: str = Field(min_length=1, max_length=100)
remark: Optional[str] = Field(default=None, max_length=500)
client_id: Optional[ClientId] = None
@field_validator("name")
@classmethod
def nonblank_name(cls, value):
if not value.strip():
raise ValueError("Blank name")
return value
class IssueCredentialRequest(WriteRequest):
client_id: Optional[ClientId] = None
class ActivateCredentialRequest(WriteRequest):
client_id: Optional[ClientId] = None
replaces_credential_id: Optional[Identifier] = None
class RevokeCredentialRequest(WriteRequest):
client_id: Optional[ClientId] = None
reason: str = Field(min_length=1, max_length=200)
@field_validator("reason")
@classmethod
def nonblank_reason(cls, value):
if not value.strip():
raise ValueError("Blank reason")
return value
class SetClientRequest(WriteRequest):
client_id: ClientId
status: Literal["AUTHORIZED", "REVOKED"]
reason: str = Field(min_length=1, max_length=200)
class SetAccountStatusRequest(WriteRequest):
status: Literal["ACTIVE", "DISABLED"]
reason: str = Field(min_length=1, max_length=200)
class AccountQueryRequest(AccountRequest):
account_ids: Optional[List[Identifier]] = Field(default=None, min_length=1, max_length=1000)
page: int = Field(default=1, ge=1)
size: int = Field(default=20, ge=1, le=200)
EvidenceRef = Annotated[str, StringConstraints(strict=True, pattern=r"^[!-~]{1,200}$")]
class GiftRequest(WriteRequest):
account_id: Identifier
client_id: ClientId
points: int = Field(ge=1, le=1000000000)
operator_note: Optional[str] = Field(default=None, max_length=200)
exclusive_writer_confirmed: bool
exclusivity_evidence_ref: EvidenceRef
@field_validator("exclusive_writer_confirmed")
@classmethod
def require_exclusive_writer(cls, value):
if value is not True:
raise ValueError("Exclusive writer confirmation required")
return value
class ReconcileGiftRequest(WriteRequest):
account_id: Identifier
client_id: ClientId
action: Literal["CONFIRM_APPLIED", "CLOSE_NOT_APPLIED", "RETRY_UNSENT"]
expected_version: int = Field(ge=0, le=2147483646)
reason: str = Field(min_length=1, max_length=500)
evidence_ref: EvidenceRef
writers_drained: bool
dry_run: bool = True
@field_validator("writers_drained")
@classmethod
def require_drained(cls, value):
if value is not True:
raise ValueError("Drained writer confirmation required")
return value
class LedgerPageRequest(AccountQueryRequest):
owner_client_id: Optional[ClientId] = None
client_id: Optional[ClientId] = None
request_id: Optional[RequestId] = None
start: Optional[datetime] = None
end: Optional[datetime] = None
@field_validator("start", "end", mode="before")
@classmethod
def utc_time(cls, value):
if value is None:
return None
if not isinstance(value, str) or len(value) > 40:
raise ValueError("ISO timestamp required")
parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
if parsed.utcoffset() is None:
raise ValueError("Timezone required")
try:
return parsed.astimezone(timezone.utc)
except OverflowError:
raise ValueError("Timestamp out of range") from None
@model_validator(mode="after")
def time_range(self):
if self.start is not None and self.end is not None and self.start >= self.end:
raise ValueError("Invalid time range")
return self
class ChatMessage(AccountRequest):
role: Literal["system", "user", "assistant"]
content: str = Field(min_length=1, max_length=1048576)
class ChatRequest(WriteRequest):
business_code: Optional[Annotated[str, StringConstraints(pattern=r"^[!-~]{1,100}$")]] = None
business_ref: Optional[Annotated[str, StringConstraints(pattern=r"^[!-~]{1,100}$")]] = None
model: Optional[Annotated[str, StringConstraints(pattern=r"^[A-Za-z0-9][A-Za-z0-9._:/+\-]{0,99}$")]] = None
messages: List[ChatMessage] = Field(min_length=1, max_length=50)
class ReconcileCallRequest(WriteRequest):
account_id: Identifier
client_id: ClientId
action: Literal["BIND_GATEWAY_LOG", "CLOSE_UNBILLED", "SETTLE_MANUAL"]
expected_version: int = Field(ge=0, le=2147483646)
reason: str = Field(min_length=1, max_length=500)
evidence_ref: EvidenceRef
writers_drained: bool
dry_run: bool = True
gateway_request_id: Optional[Annotated[str, StringConstraints(pattern=r"^[!-~]{1,100}$")]] = None
consumed_quota: Optional[Annotated[str, StringConstraints(pattern=r"^(0|[1-9][0-9]{0,18})$")]] = None
@field_validator("writers_drained")
@classmethod
def require_drained(cls, value):
if value is not True:
raise ValueError("Drained writer confirmation required")
return value
@model_validator(mode="after")
def action_fields(self):
if self.action == "BIND_GATEWAY_LOG" and self.gateway_request_id is None:
raise ValueError("Gateway request ID required")
if self.action == "SETTLE_MANUAL":
if self.consumed_quota is None or int(self.consumed_quota) > (2**63 - 1) // 146:
raise ValueError("Exact quota required")
elif self.consumed_quota is not None:
raise ValueError("Quota is only allowed for manual settlement")
if self.action == "CLOSE_UNBILLED" and self.gateway_request_id is not None:
raise ValueError("Unexpected gateway request ID")
return self
import hashlib
import hmac
import re
import secrets
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Optional
from .account_models import Account, AccountClient, Client, Credential
from .errors import AppError
ERRORS = {
"INVALID_ARGUMENT": (400, "请求参数错误"),
"PAYLOAD_TOO_LARGE": (413, "请求体超过大小限制"),
"AUTHENTICATION_FAILED": (401, "凭证校验失败"),
"CREDENTIAL_PENDING": (401, "凭证尚未激活"),
"CREDENTIAL_EXPIRED": (401, "凭证已过期"),
"CREDENTIAL_REVOKED": (401, "凭证已撤销"),
"ACCOUNT_ACCESS_REVOKED": (403, "账户来源授权已撤销"),
"ACCOUNT_DISABLED": (403, "账户已停用"),
"PERMISSION_DENIED": (403, "无权执行此操作"),
"ACCOUNT_NOT_FOUND": (404, "账户不存在或不可见"),
"REQUEST_NOT_FOUND": (404, "请求或凭证不存在或不可见"),
"FINGERPRINT_MISMATCH": (409, "相同请求号的参数与原请求不一致"),
"ACCOUNT_BLOCKED": (409, "当前状态不允许此操作"),
"GIFT_GATE_BUSY": (409, "账户正被其他赠送占用"),
"RATE_LIMITED": (429, "请求数量超过处理容量"),
"DEPENDENCY_UNAVAILABLE": (503, "依赖尚未就绪"),
"SERVICE_UNAVAILABLE": (503, "服务暂不可用,请查询原请求"),
}
KEY_PATTERN = re.compile(r"ck-([1-9][0-9]{17})-([A-Za-z0-9_-]{43})\Z")
REQUEST_PATTERN = re.compile(r"[!-~]{8,100}\Z")
class AccountError(AppError):
def __init__(self, code, *, reason=None):
status, message = ERRORS[code]
super().__init__(code, message)
self.http_status = status
self.data = {"reason": reason} if reason else None
@dataclass(frozen=True)
class Principal:
credential_id: int
client_id: str
category: str
account_id: Optional[int]
def utc_now():
return datetime.now(timezone.utc)
def issue_secret(credential_id):
secret = "ck-%s-%s" % (credential_id, secrets.token_urlsafe(32))
return secret, hashlib.sha256(secret.encode("ascii")).hexdigest(), secret[:4] + "****" + secret[-4:]
def validate_request_id(value):
if not isinstance(value, str) or not REQUEST_PATTERN.fullmatch(value):
raise AccountError("INVALID_ARGUMENT")
return value
def authenticate(session, secret, *, now=None):
match = KEY_PATTERN.fullmatch(secret) if isinstance(secret, str) else None
if match is None:
raise AccountError("AUTHENTICATION_FAILED")
credential = session.get(Credential, int(match.group(1)))
digest = hashlib.sha256(secret.encode("ascii")).hexdigest()
expected = credential.secret_digest if credential is not None else "0" * 64
if not hmac.compare_digest(digest, expected) or credential is None:
raise AccountError("AUTHENTICATION_FAILED")
now = now or utc_now()
if credential.status == "REVOKED":
raise AccountError("CREDENTIAL_REVOKED")
if credential.status == "PENDING":
if credential.pending_expires_at is None or credential.pending_expires_at <= now:
raise AccountError("CREDENTIAL_EXPIRED")
raise AccountError("CREDENTIAL_PENDING")
if credential.valid_until is not None and credential.valid_until <= now:
raise AccountError("CREDENTIAL_EXPIRED")
client = session.get(Client, credential.client_id)
if client is None or client.status != "ACTIVE":
raise AccountError("PERMISSION_DENIED")
if credential.category == "CALL":
account = session.get(Account, credential.account_id)
access = session.get(AccountClient, (credential.account_id, credential.client_id))
if account is None or access is None or access.status != "AUTHORIZED":
raise AccountError("ACCOUNT_ACCESS_REVOKED")
return Principal(credential.id, credential.client_id, credential.category, credential.account_id)
def require_management(principal):
if principal.category not in {"INTEGRATION", "PLATFORM"}:
raise AccountError("PERMISSION_DENIED")
import argparse
import asyncio
import logging
from sqlalchemy.orm import sessionmaker
from .account_config import load_account_settings
from .account_execution import run_database
from .account_main import _database_ready
from .account_model_gateway import ModelGateway
from .account_models import IdSegment
from .account_service import AccountService
from .account_settlement import SettlementService
from .db import create_mysql_engine
from .id_generator import SegmentIDGenerator
logger = logging.getLogger(__name__)
async def run(settings, batch_size):
settings.validate()
engine = create_mysql_engine(settings.database_url, pool_size=settings.db_pool_size,
max_overflow=settings.db_max_overflow, bounded_operations=True)
gateway = None
try:
await asyncio.to_thread(_database_ready, engine)
factory = sessionmaker(engine, expire_on_commit=False)
gateway = ModelGateway(settings)
accounts = AccountService(factory, SegmentIDGenerator(factory, segment_model=IdSegment), gateway, settings)
settlements = SettlementService(accounts)
settled = await settlements.run_once(batch_size)
cleared = await run_database(accounts, settlements.clear_results, batch_size)
logger.info("account worker batch finished, settled=%s, cleared=%s", settled, cleared)
return settled, cleared
finally:
if gateway is not None:
await gateway.aclose()
engine.dispose()
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--config", required=True)
parser.add_argument("--batch-size", type=int, default=50)
args = parser.parse_args()
if not 1 <= args.batch_size <= 200:
parser.error("batch-size must be between 1 and 200")
try:
settings = load_account_settings(args.config)
asyncio.run(run(settings, args.batch_size))
except Exception:
parser.exit(2, "独立算力补偿失败,请检查配置、依赖及待处理状态\n")
if __name__ == "__main__":
main()
"""HTTP 接口层(design.md §5/§7):鉴权依赖、统一响应、异常转换与路由。
只做身份/参数校验、调用 service 与协议转换,不含业务规则与事务编排;
统一响应为 {success, code, message, data, requestId}(design.md §7.2)。
"""
import logging
from fastapi import Depends, Request
from fastapi.exceptions import RequestValidationError
from fastapi.responses import JSONResponse
from sqlalchemy.orm import Session
from . import service
from .auth import AuthContext, REQUEST_ID_PATTERN
from .db import get_session, ping
from .errors import AppError, ERROR_MESSAGES
from .schemas import (
AvailabilityRequest,
ChatCompletionsRequest,
ConsumptionPageRequest,
OptAccountPageRequest,
OptGiftCreateRequest,
OptGiftPageRequest,
)
logger = logging.getLogger(__name__)
def auth_dependency(request: Request) -> AuthContext:
"""鉴权依赖:由提供器产出 AuthContext,业务层不解析鉴权 Header(design.md §5.1)。"""
auth = request.app.state.auth_provider.from_headers(request.headers)
request.state.auth = auth
return auth
def _request_id(request: Request):
auth = getattr(request.state, "auth", None)
if auth is not None:
return auth.request_id
return request.headers.get("x-request-id")
def _success(request: Request, data) -> dict:
return {
"success": True,
"code": "OK",
"message": "",
"data": data,
"requestId": _request_id(request),
}
def _error(request: Request, code: str, message: str, http_status: int) -> JSONResponse:
return JSONResponse(
status_code=http_status,
content={
"success": False,
"code": code,
"message": message,
"data": None,
"requestId": _request_id(request),
},
)
async def app_error_handler(request: Request, exc: AppError) -> JSONResponse:
return _error(request, exc.code, exc.message, exc.http_status)
async def validation_error_handler(request: Request, exc: RequestValidationError) -> JSONResponse:
"""参数校验失败统一映射 INVALID_ARGUMENT,稳定文案,不外泄内部细节。"""
logger.info("invalid argument, path=%s, count=%s", request.url.path, len(exc.errors()))
return _error(request, "INVALID_ARGUMENT", ERROR_MESSAGES["INVALID_ARGUMENT"], 422)
async def unhandled_error_handler(request: Request, exc: Exception) -> JSONResponse:
logger.exception("unhandled error, path=%s", request.url.path)
return _error(request, "INTERNAL_ERROR", ERROR_MESSAGES["INTERNAL_ERROR"], 500)
def register_exception_handlers(app) -> None:
app.add_exception_handler(AppError, app_error_handler)
app.add_exception_handler(RequestValidationError, validation_error_handler)
app.add_exception_handler(Exception, unhandled_error_handler)
def register_routes(app) -> None:
@app.get("/health/live")
def health_live():
"""进程存活检查:不访问外部依赖(design.md §7.8),探针不走 /api/v1 鉴权。"""
return {"success": True, "code": "OK", "message": "", "data": {"status": "UP"}, "requestId": None}
@app.get("/health/ready")
def health_ready(session: Session = Depends(get_session)):
"""就绪检查:验证数据库可连接(design.md §7.8),仅供内网运维。"""
try:
ping(session)
except Exception:
logger.exception("readiness check failed")
return JSONResponse(
status_code=503,
content={
"success": False,
"code": "INTERNAL_ERROR",
"message": "数据库不可用",
"data": None,
"requestId": None,
},
)
return {"success": True, "code": "OK", "message": "", "data": {"status": "UP"}, "requestId": None}
@app.get("/api/v1/merchant/accounts/{merchant_id}")
def get_merchant_account(
merchant_id: int,
request: Request,
auth: AuthContext = Depends(auth_dependency),
session: Session = Depends(get_session),
):
if merchant_id <= 0:
raise AppError("INVALID_ARGUMENT", "merchantId 必须为正整数")
data = service.get_merchant_account(session, auth, merchant_id, request.app.state.settings)
return _success(request, data)
@app.post("/api/v1/merchant/consumptions/page")
def page_merchant_consumptions(
body: ConsumptionPageRequest,
request: Request,
auth: AuthContext = Depends(auth_dependency),
session: Session = Depends(get_session),
):
data = service.page_merchant_consumptions(
session, auth, body.merchant_id, body.start_time, body.end_time, body.page, body.size
)
return _success(request, data)
@app.post("/api/v1/opt/accounts/page")
def page_opt_accounts(
body: OptAccountPageRequest,
request: Request,
auth: AuthContext = Depends(auth_dependency),
session: Session = Depends(get_session),
):
data = service.page_opt_accounts(
session, auth, body.merchant_no, body.merchant_name, body.page, body.size
)
return _success(request, data)
@app.post("/api/v1/opt/gifts")
def create_opt_gift(
body: OptGiftCreateRequest,
request: Request,
auth: AuthContext = Depends(auth_dependency),
session: Session = Depends(get_session),
):
data = service.create_opt_gift(
session, auth, request.app.state.settings, request.app.state.id_generator, body
)
return _success(request, data)
@app.post("/api/v1/opt/gifts/page")
def page_opt_gifts(
body: OptGiftPageRequest,
request: Request,
auth: AuthContext = Depends(auth_dependency),
session: Session = Depends(get_session),
):
data = service.page_opt_gifts(
session, auth, body.merchant_no, body.merchant_name,
body.start_time, body.end_time, body.page, body.size,
)
return _success(request, data)
@app.post("/api/v1/chat/completions")
def chat_completions(
body: ChatCompletionsRequest,
request: Request,
auth: AuthContext = Depends(auth_dependency),
session: Session = Depends(get_session),
):
data = service.create_chat_completion(
session, auth, request.app.state.settings, request.app.state.id_generator, body
)
return _success(request, data)
@app.get("/api/v1/calls/{request_id}")
def get_call(
request_id: str,
request: Request,
auth: AuthContext = Depends(auth_dependency),
session: Session = Depends(get_session),
):
if not REQUEST_ID_PATTERN.match(request_id or ""):
raise AppError("INVALID_ARGUMENT", "调用不存在")
data = service.get_call_status(session, auth, request_id)
return _success(request, data)
@app.post("/api/v1/accounts/availability")
def check_availability(
body: AvailabilityRequest,
request: Request,
auth: AuthContext = Depends(auth_dependency),
session: Session = Depends(get_session),
):
data = service.check_availability(session, auth, body)
return _success(request, data)
"""调用方认证上下文。
阶段 1 仅实现 AUTH_MODE=disabled 提供器:强制 X-App-Id / X-Request-Id,
构造全权限测试 AuthContext,不执行 RSA、nonce、scope 与商户授权校验
(design.md §5.1)。阶段 4 用 RsaAuthProvider 替换提供器实现,业务层零改动。
"""
import re
from dataclasses import dataclass, field
from .errors import AppError
# authentication.md §3 推荐 scopes。
SCOPE_ACCOUNT_READ = "account:read"
SCOPE_CALL_EXECUTE = "call:execute"
SCOPE_CALL_READ = "call:read"
SCOPE_OPT_ACCOUNT_READ = "opt:account:read"
SCOPE_OPT_GIFT_WRITE = "opt:gift:write"
SCOPE_OPT_GIFT_READ = "opt:gift:read"
ALL_SCOPES = frozenset({
SCOPE_ACCOUNT_READ,
SCOPE_CALL_EXECUTE,
SCOPE_CALL_READ,
SCOPE_OPT_ACCOUNT_READ,
SCOPE_OPT_GIFT_WRITE,
SCOPE_OPT_GIFT_READ,
})
APP_ID_PATTERN = re.compile(r"^[A-Za-z0-9._-]{1,64}$")
REQUEST_ID_PATTERN = re.compile(r"^[A-Za-z0-9._:-]{1,100}$")
HEADER_APP_ID = "x-app-id"
HEADER_REQUEST_ID = "x-request-id"
@dataclass(frozen=True)
class AuthContext:
"""已认证调用方上下文。业务层只依赖本对象,不直接解析鉴权 Header(design.md §5.1)。"""
app_id: str
request_id: str
key_id: str = None
scopes: frozenset = field(default=ALL_SCOPES)
merchant_scope: str = "ALL"
bound_merchant_ids: frozenset = field(default=frozenset())
business_codes: frozenset = field(default=None)
def allows_scope(self, scope: str) -> bool:
return scope in self.scopes
def allows_business_code(self, business_code: str) -> bool:
"""businessCode 白名单:None 表示未配置(disabled 测试模式放行,design.md §5.1)。"""
if self.business_codes is None:
return True
return business_code in self.business_codes
def allows_merchant(self, merchant_id) -> bool:
"""merchantScope 校验:ALL 放行,BOUND 必须显式绑定(authentication.md §3)。"""
if self.merchant_scope == "ALL":
return True
return merchant_id is not None and int(merchant_id) in self.bound_merchant_ids
class AuthContextProvider:
"""鉴权提供器接口:输入 HTTP Header,输出可信 AuthContext。"""
def from_headers(self, headers) -> AuthContext:
raise NotImplementedError
class DisabledAuthProvider(AuthContextProvider):
"""AUTH_MODE=disabled:仅校验 X-App-Id / X-Request-Id,构造全权限测试上下文。
disabled 模式下 source_app 仍来自 X-App-Id,不允许 Body 覆盖(design.md §5.1)。
"""
def from_headers(self, headers) -> AuthContext:
app_id = headers.get(HEADER_APP_ID)
if not app_id:
raise AppError("AUTH_HEADER_MISSING")
if not APP_ID_PATTERN.match(app_id):
raise AppError("AUTH_HEADER_INVALID")
request_id = headers.get(HEADER_REQUEST_ID)
if not request_id:
raise AppError("AUTH_HEADER_MISSING")
if not REQUEST_ID_PATTERN.match(request_id):
raise AppError("AUTH_HEADER_INVALID")
return AuthContext(app_id=app_id, request_id=request_id)
def get_auth_provider(auth_mode: str) -> AuthContextProvider:
"""按 AUTH_MODE 返回提供器。RSA 提供器在阶段 4 实现前不可用。"""
if auth_mode == "disabled":
return DisabledAuthProvider()
raise RuntimeError("RSA 鉴权提供器尚未实现(迁移方案阶段 4)")
"""固定计费基线与换算(design.md §6.6/§6.9)。
10000/146/500000 是冻结的计费常量,不是可配置项;展示积分一律 Decimal,
禁止 float 中转。所有整数结果必须落在 signed BIGINT 范围内。
"""
from decimal import Decimal
POINT_SCALE = 10000
POINT_UNITS_PER_QUOTA = 146
EXPECTED_QUOTA_PER_UNIT = 500000
BIGINT_MIN = -(2**63)
BIGINT_MAX = 2**63 - 1
_SCALES = {
POINT_SCALE,
POINT_UNITS_PER_QUOTA,
EXPECTED_QUOTA_PER_UNIT,
}
def require_bigint(value: int, name: str = "value") -> int:
"""校验整数结果在 signed BIGINT 范围内,溢出拒绝落账(design.md §6.6)。"""
if not isinstance(value, int) or isinstance(value, bool):
raise ValueError("%s 必须为整数" % name)
if not BIGINT_MIN <= value <= BIGINT_MAX:
raise OverflowError("%s 超出 signed BIGINT 范围: %d" % (name, value))
return value
def quota_to_point_units(quota: int) -> int:
"""网关日志 quota -> 消耗子单位:max(quota, 0) * 146,无逐笔整分取整。"""
if quota is None:
raise ValueError("quota 不能为空")
units = max(int(quota), 0) * POINT_UNITS_PER_QUOTA
return require_bigint(units, "consumedPointUnits")
def requested_points_to_units(requested_points: int) -> int:
"""申请整数积分 -> 申请子单位:points * 10000。"""
units = int(requested_points) * POINT_SCALE
return require_bigint(units, "requestedUnits")
def point_units_to_quota_half_up(point_units: int) -> int:
"""申请子单位 -> 目标 quota,正值 HALF_UP:(units + 73) // 146。"""
units = int(point_units)
if units <= 0:
return 0
quota = (units + POINT_UNITS_PER_QUOTA // 2) // POINT_UNITS_PER_QUOTA
return require_bigint(quota, "quotaDelta")
def units_to_points(units: int) -> Decimal:
"""子单位 -> 展示积分 Decimal(units)/10000,由调用方固定 4 位格式化。"""
return Decimal(int(units)) / Decimal(POINT_SCALE)
def format_points(units) -> str:
"""子单位 -> 固定 4 位小数字符串;None 返回 None(未结算)。"""
if units is None:
return None
return format(units_to_points(units), ".4f")
def units_json(units) -> str:
"""子单位/quota -> 十进制字符串;None 返回 None。"""
if units is None:
return None
return str(int(units))
"""环境变量配置。全部配置仅来自环境变量(design.md §12),敏感值由部署系统注入。
计费常量 10000/146/500000 为冻结基线(app/billing.py),不提供环境配置入口。
"""
import os
from dataclasses import dataclass
# disabled 鉴权只允许隔离的本地/测试环境;任意其他环境名一律要求 rsa,
# 未枚举的新环境默认拒绝 disabled(design.md §5.1/§12)。
DISABLED_ALLOWED_ENVS = frozenset({"local", "test"})
ALLOWED_AUTH_MODES = frozenset({"disabled", "rsa"})
class StartupError(RuntimeError):
"""启动参数不满足安全约束。"""
@dataclass(frozen=True)
class Settings:
app_env: str
app_name: str
log_level: str
auth_mode: str
database_url: str
db_pool_size: int
db_max_overflow: int
# NewAPI 管理 API(design.md §12):PAT 只能由部署环境注入,不落库不写日志。
newapi_base_url: str = ""
newapi_pat: str = ""
newapi_chat_url: str = ""
newapi_model: str = "deepseek-v4.1-flash"
newapi_token_name_prefix: str = "mei1-saas"
newapi_connect_timeout_seconds: float = 10.0
newapi_read_timeout_seconds: float = 120.0
# 延迟结算(design.md §6.7/§12):超过限制次数转 SETTLE_FAILED。
settlement_retry_limit: int = 24
# 立即结算短重试(design.md §6.7):次数与间隔可配置。
settle_short_retries: int = 3
settle_short_retry_seconds: float = 2.0
@property
def auth_required(self) -> bool:
"""非 local/test 环境均要求 RSA 鉴权。"""
return self.app_env not in DISABLED_ALLOWED_ENVS
def validate(self) -> None:
"""启动期安全校验:disabled 白名单外的环境必须使用 rsa,否则拒绝启动。"""
if self.auth_mode not in ALLOWED_AUTH_MODES:
raise StartupError("AUTH_MODE 必须为 disabled 或 rsa: %s" % self.auth_mode)
if self.auth_required and self.auth_mode != "rsa":
raise StartupError(
"APP_ENV=%s 下必须使用 AUTH_MODE=rsa,禁止以 disabled 启动" % self.app_env
)
def load_settings(environ=None) -> Settings:
env = environ if environ is not None else os.environ
return Settings(
app_env=env.get("APP_ENV", "local"),
app_name=env.get("APP_NAME", "mei1-computing-service"),
log_level=env.get("LOG_LEVEL", "INFO"),
auth_mode=env.get("AUTH_MODE", "disabled"),
database_url=env.get("DATABASE_URL", ""),
db_pool_size=int(env.get("DB_POOL_SIZE", "10")),
db_max_overflow=int(env.get("DB_MAX_OVERFLOW", "20")),
newapi_base_url=env.get("NEWAPI_BASE_URL", ""),
newapi_pat=env.get("NEWAPI_PAT", ""),
newapi_chat_url=env.get("NEWAPI_CHAT_URL", ""),
newapi_model=env.get("NEWAPI_MODEL", "deepseek-v4.1-flash"),
newapi_token_name_prefix=env.get("NEWAPI_TOKEN_NAME_PREFIX", "mei1-saas"),
newapi_connect_timeout_seconds=float(env.get("NEWAPI_CONNECT_TIMEOUT_SECONDS", "10")),
newapi_read_timeout_seconds=float(env.get("NEWAPI_READ_TIMEOUT_SECONDS", "120")),
settlement_retry_limit=int(env.get("SETTLEMENT_RETRY_LIMIT", "24")),
settle_short_retries=int(env.get("SETTLE_SHORT_RETRIES", "3")),
settle_short_retry_seconds=float(env.get("SETTLE_SHORT_RETRY_SECONDS", "2")),
)
"""Engine、Session 与事务辅助(SQLAlchemy 同步 Session,design.md §2)。"""
from contextlib import contextmanager
from sqlalchemy import create_engine, event, text
from sqlalchemy.engine import make_url
from sqlalchemy.orm import Session, sessionmaker
from .config import Settings
_engine = None
_session_factory = None
def create_mysql_engine(database_url: str, *, pool_size=5, max_overflow=0, bounded_operations=False):
url = make_url(database_url)
if url.drivername != "mysql+pymysql" or not url.database:
raise ValueError("需要指定独立数据库的 mysql+pymysql 连接")
engine = create_engine(
url,
pool_size=pool_size,
max_overflow=max_overflow,
pool_timeout=5 if bounded_operations else 30,
pool_pre_ping=True,
pool_recycle=3600,
isolation_level="READ COMMITTED",
hide_parameters=True,
connect_args={"charset": "utf8mb4", "connect_timeout": 5 if bounded_operations else 10,
"read_timeout": 5 if bounded_operations else 30,
"write_timeout": 5 if bounded_operations else 30},
)
@event.listens_for(engine, "connect")
def configure_session(connection, _):
with connection.cursor() as cursor:
cursor.execute("SET SESSION time_zone = '+00:00'")
cursor.execute("SET SESSION innodb_lock_wait_timeout = 5" if bounded_operations
else "SET SESSION innodb_lock_wait_timeout = 10")
cursor.execute("SET SESSION sql_mode = CONCAT_WS(',', @@sql_mode, 'STRICT_TRANS_TABLES')")
if bounded_operations:
cursor.execute("SET SESSION lock_wait_timeout = 5")
cursor.execute("SET SESSION max_execution_time = 5000")
if bounded_operations:
from .account_execution import check_database_deadline
@event.listens_for(engine, "before_cursor_execute")
def before_statement(connection, cursor, statement, parameters, context, executemany):
check_database_deadline()
@event.listens_for(engine, "commit")
def before_commit(connection):
check_database_deadline()
return engine
def init_engine(settings: Settings) -> None:
"""按环境配置初始化全局 Engine。重复调用忽略(uvicorn 多次加载保护)。"""
global _engine, _session_factory
if _engine is not None:
return
if not settings.database_url:
raise RuntimeError("DATABASE_URL 未配置")
_engine = create_engine(
settings.database_url,
pool_size=settings.db_pool_size,
max_overflow=settings.db_max_overflow,
pool_pre_ping=True,
pool_recycle=3600,
future=True,
)
_session_factory = sessionmaker(bind=_engine, expire_on_commit=False, future=True)
def dispose_engine() -> None:
global _engine, _session_factory
if _engine is not None:
_engine.dispose()
_engine = None
_session_factory = None
def get_session():
"""FastAPI 依赖:每请求一个 Session,测试通过 dependency_overrides 替换。"""
if _session_factory is None:
raise RuntimeError("数据库未初始化,请先调用 init_engine")
session = _session_factory()
try:
yield session
finally:
session.close()
def get_session_factory():
"""返回会话工厂,供号段生成器等组件创建独立短事务 Session(design.md §9.6)。"""
if _session_factory is None:
raise RuntimeError("数据库未初始化,请先调用 init_engine")
return _session_factory
@contextmanager
def session_scope(settings: Settings) -> Session:
"""脚本/任务用短事务作用域:成功提交,异常回滚。"""
if _session_factory is None:
init_engine(settings)
session = _session_factory()
try:
yield session
session.commit()
except Exception:
session.rollback()
raise
finally:
session.close()
def ping(session: Session) -> bool:
"""就绪检查:验证数据库可连接(design.md §7.8 /health/ready)。"""
session.execute(text("SELECT 1"))
return True
"""统一错误码与业务异常。错误码全集见 design.md §8 与 authentication.md §12。"""
# 错误码 -> HTTP 状态码。202 用于异步/待结算语义的成功受理。
ERROR_HTTP_STATUS = {
"AUTH_HEADER_MISSING": 401,
"AUTH_HEADER_INVALID": 401,
"AUTH_CLIENT_INVALID": 401,
"AUTH_KEY_INVALID": 401,
"AUTH_TIMESTAMP_EXPIRED": 401,
"AUTH_BODY_DIGEST_MISMATCH": 401,
"AUTH_SIGNATURE_INVALID": 401,
"AUTH_REPLAY_DETECTED": 401,
"AUTH_QUERY_NOT_SUPPORTED": 400,
"FORBIDDEN_SCOPE": 403,
"FORBIDDEN_BUSINESS_CODE": 403,
"FORBIDDEN_MERCHANT": 403,
"INVALID_ARGUMENT": 422,
"ACCOUNT_NOT_OPEN": 409,
"ACCOUNT_DISABLED": 409,
"ACCOUNT_SYNC_FAILED": 409,
"ACCOUNT_INSUFFICIENT": 409,
"DUPLICATE_REQUEST_MISMATCH": 409,
"CALL_PROCESSING": 202,
"CALL_STATUS_UNKNOWN": 409,
"GATEWAY_ERROR": 502,
"GATEWAY_TIMEOUT": 504,
"SETTLEMENT_PENDING": 202,
"SERVICE_UNAVAILABLE": 503,
"INTERNAL_ERROR": 500,
}
# 稳定业务文案(design.md §6.4:失败时返回稳定业务错误码和友好文案)。
ERROR_MESSAGES = {
"AUTH_HEADER_MISSING": "缺少鉴权请求头",
"AUTH_HEADER_INVALID": "鉴权请求头格式错误",
"FORBIDDEN_SCOPE": "调用方没有接口权限",
"FORBIDDEN_BUSINESS_CODE": "调用方没有业务码权限",
"FORBIDDEN_MERCHANT": "调用方没有商户权限",
"INVALID_ARGUMENT": "请求参数错误",
"ACCOUNT_NOT_OPEN": "算力账户未开通",
"ACCOUNT_DISABLED": "算力账户已停用,请联系平台处理",
"ACCOUNT_SYNC_FAILED": "算力账户同步异常,请联系平台处理",
"ACCOUNT_INSUFFICIENT": "算力积分不足,请联系平台赠送算力积分",
"DUPLICATE_REQUEST_MISMATCH": "相同幂等号请求参数与原请求不一致",
"CALL_STATUS_UNKNOWN": "网关是否受理未知,禁止自动重试",
"GATEWAY_ERROR": "模型网关调用失败",
"GATEWAY_TIMEOUT": "模型网关调用超时",
"SERVICE_UNAVAILABLE": "服务暂不可用,请稍后重试",
"INTERNAL_ERROR": "服务内部异常",
}
class AppError(Exception):
"""业务异常。api 层统一转换为 {success, code, message, data, requestId} 响应。"""
def __init__(self, code: str, message: str = None):
self.code = code
self.http_status = ERROR_HTTP_STATUS.get(code, 500)
self.message = message if message is not None else ERROR_MESSAGES.get(code, code)
super().__init__(self.message)
"""固定 18 位数值 ID 号段生成器(design.md §9.6)。
数据库单一分配行推进高水位,进程内锁保护候选段;号段事务与业务事务分离,
业务回滚不回收已分配 ID。不依赖时钟、节点号,重启/fork 丢弃本地余段。
"""
import logging
import os
import threading
import weakref
from sqlalchemy import select
from sqlalchemy.orm import Session
from .errors import AppError
from .models import IdSegment
logger = logging.getLogger(__name__)
ID_LOW = 100_000_000_000_000_000
ID_HIGH_EXCLUSIVE = 1_000_000_000_000_000_000
SEGMENT_SIZE = 1000
GENERATOR_KEY = "computing-global"
_ALLOCATION_ATTEMPTS = 3
_live_generators = weakref.WeakSet()
def _reset_generators_after_fork():
for generator in _live_generators:
generator._lock = threading.Lock()
generator._discard_segment()
generator._pid = os.getpid()
if hasattr(os, "register_at_fork"):
os.register_at_fork(after_in_child=_reset_generators_after_fork)
class SegmentIDGenerator:
"""多 worker/任务/受控工具共用同一分配行的 18 位 ID 发号器。"""
def __init__(self, session_factory, *, segment_model=IdSegment):
self._session_factory = session_factory
self._segment_model = segment_model
self._lock = threading.Lock()
self._cursor = None
self._segment_end = None
self._pid = os.getpid()
_live_generators.add(self)
def next_id(self) -> int:
with self._lock:
if os.getpid() != self._pid:
# fork 后继承的候选段仅供父进程使用,子进程必须重新申请(§9.6)。
self._discard_segment()
self._pid = os.getpid()
if self._cursor is None or self._cursor >= self._segment_end:
self._allocate()
value = self._cursor
self._cursor += 1
return value
def _discard_segment(self) -> None:
self._cursor = None
self._segment_end = None
def _allocate(self) -> None:
"""独立短事务推进高水位;提交确认前不发号,失败废弃候选段重试。"""
for attempt in range(1, _ALLOCATION_ATTEMPTS + 1):
session = None
try:
session = self._session_factory()
row = session.scalars(
select(self._segment_model)
.where(self._segment_model.generator_key == GENERATOR_KEY)
.with_for_update()
).one()
low = int(row.next_id)
if not ID_LOW <= low < ID_HIGH_EXCLUSIVE:
logger.error("id segment exhausted, generator_key=%s", GENERATOR_KEY)
raise AppError("SERVICE_UNAVAILABLE")
high = min(low + SEGMENT_SIZE, ID_HIGH_EXCLUSIVE)
row.next_id = high
session.commit()
except AppError:
if session is not None:
session.rollback()
raise
except Exception:
if session is not None:
session.rollback()
logger.warning(
"id segment allocation failed, attempt=%s, generator_key=%s",
attempt, GENERATOR_KEY,
)
else:
# commit 成功返回即视为持久确认,才允许使用 [low, high)。
self._cursor = low
self._segment_end = high
return
finally:
if session is not None:
session.close()
logger.error("id segment allocation exhausted retries, generator_key=%s", GENERATOR_KEY)
raise AppError("SERVICE_UNAVAILABLE")
def seed_segment(session: Session, next_id: int = ID_LOW, *, segment_model=IdSegment) -> None:
"""受控初始化只允许提高高水位,禁止回收已分配号段。"""
if not ID_LOW <= next_id <= ID_HIGH_EXCLUSIVE:
raise ValueError("号段高水位超出允许范围")
row = session.get(segment_model, GENERATOR_KEY, with_for_update=True, populate_existing=True)
if row is None:
session.add(segment_model(generator_key=GENERATOR_KEY, next_id=next_id))
elif next_id > row.next_id:
row.next_id = next_id
session.flush()
"""SETTLE_PENDING 补偿任务(design.md §6.7/§10.3)。
用法:python -m app.jobs.settle_pending
多实例安全(MySQL 8+):SELECT ... FOR UPDATE SKIP LOCKED 分批领取后立即提交
释放行锁;网关日志查询在事务外;逐条独立短事务复核落账。超过重试上限转
SETTLE_FAILED 保留人工核对入口,不自动记零消耗(design.md §6.7)。
"""
import logging
from datetime import datetime, timedelta
from app import newapi, repository, service
from app.config import load_settings
from app.db import get_session_factory, init_engine
from app.id_generator import SegmentIDGenerator
from app.models import ComputingCall
logger = logging.getLogger(__name__)
# 退避节奏:第 n 次重查失败后延迟 BACKOFF_MINUTES[min(n-1, 3)](design.md §6.7)。
BACKOFF_MINUTES = (1, 5, 15, 60)
def _backoff(retry_count: int) -> timedelta:
index = min(max(int(retry_count), 1) - 1, len(BACKOFF_MINUTES) - 1)
return timedelta(minutes=BACKOFF_MINUTES[index])
def _mark_settle_failed(session, call_id, error_message: str) -> None:
call = session.get(ComputingCall, call_id, with_for_update=True)
if call is not None and call.status == service.CALL_STATUS_SETTLE_PENDING:
call.status = service.CALL_STATUS_SETTLE_FAILED
call.error_code = "SETTLE_FAILED"
call.error_message = service._truncate_error(error_message)
session.commit()
def _process_call(session_factory, gateway, id_generator, call_id: int, settings) -> None:
now = datetime.now()
# 事务 1:复核仍待结算且到期,提交释放行锁(design.md §10.3)。
with session_factory() as session:
call = session.get(ComputingCall, call_id, with_for_update=True)
if call is None or call.status != service.CALL_STATUS_SETTLE_PENDING:
return
if call.next_retry_time is None or call.next_retry_time > now:
return
gateway_request_id = call.gateway_request_id
session.commit()
if not gateway_request_id:
with session_factory() as session:
_mark_settle_failed(session, call_id, "缺少可核对的网关请求ID,需人工核对")
return
# 网关日志查询在事务外(design.md §10.1/§10.3)。
try:
quota = gateway.find_log_quota(gateway_request_id)
except newapi.NewApiError:
logger.warning("settle pending log query failed, callId=%s", call_id)
quota = None
# 事务 2:锁调用复核仍待结算后落账或推迟。
with session_factory() as session:
call = session.get(ComputingCall, call_id, with_for_update=True)
if call is None or call.status != service.CALL_STATUS_SETTLE_PENDING:
return
if quota is not None:
service.settle_call(session, id_generator, call, quota)
return
call.retry_count = int(call.retry_count or 0) + 1
if int(call.retry_count) > int(settings.settlement_retry_limit):
call.status = service.CALL_STATUS_SETTLE_FAILED
call.error_code = "SETTLE_FAILED"
call.error_message = service._truncate_error(
"网关日志超过重试上限仍未生成,需人工核对"
)
else:
call.next_retry_time = datetime.now() + _backoff(int(call.retry_count))
session.commit()
def run_once(session_factory, settings, id_generator, batch_size: int = 50) -> int:
"""领取并处理一批到期 SETTLE_PENDING 调用,返回本批领取数量。"""
gateway = service._newapi_gateway(settings)
now = datetime.now()
with session_factory() as session:
calls = repository.claim_settle_pending_calls(session, batch_size, now)
call_ids = [call.id for call in calls]
session.commit()
for call_id in call_ids:
try:
_process_call(session_factory, gateway, id_generator, call_id, settings)
except Exception:
logger.exception("settle pending processing failed, callId=%s", call_id)
return len(call_ids)
def main() -> None:
logging.basicConfig(level=logging.INFO)
settings = load_settings()
init_engine(settings)
session_factory = get_session_factory()
id_generator = SegmentIDGenerator(session_factory)
processed = run_once(session_factory, settings, id_generator)
logger.info("settle pending batch done, count=%s", processed)
if __name__ == "__main__":
main()
"""应用入口:lifespan 装配配置、Engine 与鉴权提供器(design.md §5)。
启动命令(容器内):uvicorn app.main:app --host 0.0.0.0 --port 8000
AUTH_MODE/APP_ENV 安全约束由 Settings.validate() 在启动期强制执行,
test2/test3/test4/prod 未启用 RSA 时拒绝启动(design.md §5.1)。
"""
import logging
from contextlib import asynccontextmanager
from fastapi import FastAPI
from .api import register_exception_handlers, register_routes
from .auth import get_auth_provider
from .config import load_settings
from .db import dispose_engine, get_session_factory, init_engine
from .id_generator import SegmentIDGenerator
@asynccontextmanager
async def lifespan(app: FastAPI):
settings = load_settings()
settings.validate()
logging.basicConfig(
level=getattr(logging, settings.log_level.upper(), logging.INFO),
format="%(asctime)s %(levelname)s %(name)s %(message)s",
)
init_engine(settings)
app.state.settings = settings
app.state.auth_provider = get_auth_provider(settings.auth_mode)
app.state.id_generator = SegmentIDGenerator(get_session_factory)
yield
dispose_engine()
def create_app() -> FastAPI:
app = FastAPI(title="mei1-computing-service", lifespan=lifespan)
register_exception_handlers(app)
register_routes(app)
return app
app = create_app()
This diff is collapsed. Click to expand it.
This diff is collapsed. Click to expand it.
"""员工/门店归属解析(design.md §6.3)。
解析分两步:确定商户(由调用链上游完成,本模块假定 merchantId 已可信),
再补充归属快照。降级只能发生在已得到可信 merchantId 之后,不得猜测商户;
storeId=null 不阻断调用,本次消耗记商户 MASTER 总账户(2026-09-21 决策:
总部员工无门店是正常业务场景,逐级降级、永不抛错)。
本模块为纯逻辑:员工/门店快照由调用方通过 loader 提供,便于单元测试。
"""
from dataclasses import dataclass
# 与 Java 实体常量对齐:Employee.STATUS_NORMAL=1、POSITION_STATUS_LEAVE=2、Store.STATUS_OPEN=1。
EMPLOYEE_STATUS_NORMAL = 1
POSITION_STATUS_LEAVE = 2
STORE_STATUS_OPEN = 1
# 降级原因(写入 warn 日志的 degradation 标记)。
EMPLOYEE_NOT_FOUND = "EMPLOYEE_NOT_FOUND"
EMPLOYEE_CROSS_MERCHANT = "EMPLOYEE_CROSS_MERCHANT"
EMPLOYEE_INACTIVE = "EMPLOYEE_INACTIVE"
EMPLOYEE_NO_STORE = "EMPLOYEE_NO_STORE"
STORE_NOT_FOUND = "STORE_NOT_FOUND"
STORE_CROSS_MERCHANT = "STORE_CROSS_MERCHANT"
STORE_NOT_OPEN = "STORE_NOT_OPEN"
@dataclass(frozen=True)
class EmployeeSnapshot:
"""t_saas_mb_employee 归属解析所需字段的快照。"""
id: int
merchant_id: int
status: int
position_status: int = None
dimission_date: object = None
store_id: int = None
delete_flag: bool = False
@dataclass(frozen=True)
class StoreSnapshot:
"""t_saas_mb_store 归属解析所需字段的快照。"""
id: int
merchant_id: int
status: int
@dataclass(frozen=True)
class ResolvedIdentity:
"""可信归属输出:merchantId 必有,storeId/employeeId 可空,employeeAccount 为快照。"""
merchant_id: int
store_id: int = None
employee_id: int = None
employee_account: str = None
degradations: tuple = tuple()
def _employee_usable(employee: EmployeeSnapshot) -> bool:
if employee.delete_flag:
return False
if employee.status is not None and employee.status != EMPLOYEE_STATUS_NORMAL:
return False
if employee.position_status is not None and employee.position_status == POSITION_STATUS_LEAVE:
return False
if employee.dimission_date is not None:
return False
return True
def resolve_identity(merchant_id, request_store_id, employee_id, employee_account,
employee_loader, store_loader) -> ResolvedIdentity:
"""按 §6.3 顺序解析:员工覆盖候选门店 -> 候选门店归属校验 -> 降级记总账户。
employee_loader(employee_id) -> EmployeeSnapshot | None
store_loader(store_id) -> StoreSnapshot | None
"""
degradations = []
resolved_store_id = request_store_id
resolved_employee_id = None
if employee_id is not None:
employee = employee_loader(employee_id)
if employee is None or employee.delete_flag:
degradations.append(EMPLOYEE_NOT_FOUND)
elif employee.merchant_id != merchant_id:
degradations.append(EMPLOYEE_CROSS_MERCHANT)
elif not _employee_usable(employee):
degradations.append(EMPLOYEE_INACTIVE)
else:
# 员工可信即保留 employee_id 输出;门店归属缺失不影响身份快照。
resolved_employee_id = employee.id
if employee.store_id is not None:
resolved_store_id = employee.store_id
else:
# 总部员工无门店:忽略员工门店归属,保留请求候选门店并记 warn。
degradations.append(EMPLOYEE_NO_STORE)
if resolved_store_id is not None:
store = store_loader(resolved_store_id)
if store is None:
degradations.append(STORE_NOT_FOUND)
resolved_store_id = None
elif store.merchant_id != merchant_id:
degradations.append(STORE_CROSS_MERCHANT)
resolved_store_id = None
elif store.status is not None and store.status != STORE_STATUS_OPEN:
degradations.append(STORE_NOT_OPEN)
resolved_store_id = None
return ResolvedIdentity(
merchant_id=merchant_id,
store_id=resolved_store_id,
employee_id=resolved_employee_id,
employee_account=employee_account,
degradations=tuple(degradations),
)
"""请求模型(pydantic v2)。JSON 键为 camelCase,与 Java 侧请求契约一致(design.md §7.2)。
extra="forbid":请求 Body 中的未知字段(sourceApp、consumedPoints 等服务端权威
字段)一律拒绝,从契约上杜绝调用方越权覆盖(design.md §7.6)。
"""
from datetime import datetime
from typing import Optional
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from pydantic.alias_generators import to_camel
# 分页 size 上限;超出按参数错误处理(design.md §6.10)。
MAX_PAGE_SIZE = 200
class CamelModel(BaseModel):
"""snake_case 字段名与 camelCase JSON 键双向映射。"""
model_config = ConfigDict(alias_generator=to_camel, populate_by_name=True, extra="forbid")
class PageRequest(CamelModel):
page: int = Field(default=1, ge=1, description="页码,从 1 开始")
size: int = Field(default=20, ge=1, le=MAX_PAGE_SIZE, description="每页条数,上限 200")
class ConsumptionPageRequest(PageRequest):
"""商户端 Token 消耗记录分页请求(design.md §7.9.3)。时间窗口左闭右开。"""
merchant_id: int = Field(ge=1)
start_time: Optional[datetime] = None
end_time: Optional[datetime] = None
@model_validator(mode="after")
def check_time_window(self):
_validate_window(self.start_time, self.end_time)
return self
class OptAccountPageRequest(PageRequest):
"""OPT 账户分页请求(design.md §7.9.2)。M 号/名称为模糊匹配。"""
merchant_no: Optional[str] = Field(default=None, max_length=50)
merchant_name: Optional[str] = Field(default=None, max_length=100)
class OptGiftPageRequest(PageRequest):
"""OPT 赠送记录分页请求(design.md §7.9.4)。时间窗口左闭右开。"""
merchant_no: Optional[str] = Field(default=None, max_length=50)
merchant_name: Optional[str] = Field(default=None, max_length=100)
start_time: Optional[datetime] = None
end_time: Optional[datetime] = None
@model_validator(mode="after")
def check_time_window(self):
_validate_window(self.start_time, self.end_time)
return self
ALLOWED_MESSAGE_ROLES = ("system", "user", "assistant")
class ChatMessage(CamelModel):
"""OpenAI 兼容消息子集:第一版仅支持文本内容(design.md §6.5)。"""
role: str
content: str = Field(min_length=1)
@field_validator("role")
@classmethod
def _check_role(cls, value):
value = (value or "").strip().lower()
if value not in ALLOWED_MESSAGE_ROLES:
raise ValueError("role 仅支持 system/user/assistant")
return value
class ChatCompletionsRequest(CamelModel):
"""模型调用请求(design.md §7.3)。requestId 取 X-Request-Id Header。
阶段 3 只有 Java SaaS 代理调用(登录上下文 merchantId 必传);
appId 唯一绑定商户省略 merchantId 的解析在阶段 4 RSA 鉴权后开放。
"""
merchant_id: int = Field(ge=1)
store_id: Optional[int] = Field(default=None, ge=1)
employee_id: Optional[int] = Field(default=None, ge=1)
employee_account: Optional[str] = Field(default=None, max_length=100)
business_code: str = Field(min_length=1, max_length=100)
messages: list[ChatMessage] = Field(min_length=1, max_length=50)
content_length: Optional[int] = Field(default=None, ge=0)
@field_validator("business_code")
@classmethod
def _strip_business_code(cls, value):
value = (value or "").strip()
if not value:
raise ValueError("businessCode 不能为空白")
return value
class AvailabilityRequest(CamelModel):
"""可用性检查请求(design.md §7.5),商户解析规则与模型调用一致。"""
merchant_id: int = Field(ge=1)
store_id: Optional[int] = Field(default=None, ge=1)
employee_id: Optional[int] = Field(default=None, ge=1)
class OptGiftCreateRequest(CamelModel):
"""OPT 赠送请求(design.md §7.7/§6.9)。
Java OPT Controller 先完成权限校验(computing_account_gift),再把操作员
ID/名称与幂等号传入;sourceApp 由服务端从 X-App-Id 取,Body 不接受。
"""
merchant_id: int = Field(ge=1)
request_id: str = Field(min_length=1, max_length=100)
points: int = Field(ge=1, le=2147483647, description="申请整数积分,正数")
remark: str = Field(min_length=1, max_length=500, description="赠送原因,必填")
operator_id: Optional[int] = Field(default=None, ge=1)
operator_name: Optional[str] = Field(default=None, max_length=100)
@field_validator("request_id", "remark", "operator_name")
@classmethod
def _strip_nonblank(cls, value):
if value is None:
return None
value = value.strip()
if not value:
raise ValueError("不能为空白字符")
return value
def _validate_window(start_time, end_time) -> None:
"""两端同时提供时必须 start < end(design.md §7.9.1)。"""
if start_time is not None and end_time is not None and start_time >= end_time:
raise ValueError("startTime 必须早于 endTime")
This diff is collapsed. Click to expand it.
This diff is collapsed. Click to expand it.
[build-system]
requires = ["setuptools>=68"]
build-backend = "setuptools.build_meta"
[project]
name = "mei1-computing-service"
version = "0.1.0"
description = "美一独立算力账户、凭证、账务与模型调用服务"
requires-python = ">=3.12"
dependencies = [
"fastapi>=0.115",
"uvicorn>=0.30",
"sqlalchemy>=2.0",
"pymysql>=1.1",
"cryptography>=42",
"httpx>=0.27",
"pydantic>=2.7",
]
[project.optional-dependencies]
dev = [
"pytest>=8",
]
[tool.setuptools.packages.find]
include = ["app", "app.*"]
[tool.pytest.ini_options]
testpaths = ["tests"]
addopts = "-q"
-- P1 source preflight: read-only, run manually on the OLD SaaS schema.
-- Sections returning anomaly rows must be reviewed before migration.
-- Inventory/version/high-water queries intentionally return rows.
-- Do not infer client ownership, missing fees, or gateway acceptance from these checks.
SELECT VERSION() AS mysql_version, DATABASE() AS source_database,
@@session.time_zone AS session_time_zone, @@global.time_zone AS global_time_zone;
SELECT table_name, column_name, column_type, is_nullable, collation_name
FROM information_schema.columns
WHERE table_schema = DATABASE()
AND table_name IN (
't_saas_ai_merchant_computing_account',
't_saas_ai_merchant_computing_point_record',
't_saas_ai_computing_token_consumption_record',
't_saas_ai_computing_call',
't_saas_ai_computing_account_gate',
't_saas_ai_computing_id_segment'
)
ORDER BY table_name, ordinal_position;
SELECT table_name, index_name, non_unique, seq_in_index, column_name
FROM information_schema.statistics
WHERE table_schema = DATABASE() AND table_name IN (
't_saas_ai_merchant_computing_account',
't_saas_ai_merchant_computing_point_record',
't_saas_ai_computing_token_consumption_record',
't_saas_ai_computing_call'
)
ORDER BY table_name, index_name, seq_in_index;
SELECT merchant_id, COUNT(*) AS account_count
FROM t_saas_ai_merchant_computing_account
WHERE account_type = 'MASTER'
GROUP BY merchant_id HAVING COUNT(*) > 1;
SELECT id, merchant_id, account_type
FROM t_saas_ai_merchant_computing_account
WHERE account_type <> 'MASTER' OR parent_account_id IS NOT NULL;
SELECT gateway_token_id, COUNT(*) AS account_count
FROM t_saas_ai_merchant_computing_account
WHERE gateway_token_id IS NOT NULL
GROUP BY gateway_token_id HAVING COUNT(*) > 1;
SELECT p.id, p.account_id, p.merchant_id
FROM t_saas_ai_merchant_computing_point_record p
LEFT JOIN t_saas_ai_merchant_computing_account a ON a.id = p.account_id
WHERE a.id IS NULL OR p.merchant_id <> a.merchant_id;
SELECT c.id, c.account_id, c.merchant_id
FROM t_saas_ai_computing_token_consumption_record c
LEFT JOIN t_saas_ai_merchant_computing_account a ON a.id = c.account_id
WHERE a.id IS NULL OR c.merchant_id <> a.merchant_id;
SELECT c.id, c.account_id, c.merchant_id
FROM t_saas_ai_computing_call c
LEFT JOIN t_saas_ai_merchant_computing_account a ON a.id = c.account_id
WHERE a.id IS NULL OR c.merchant_id <> a.merchant_id;
SELECT source_app, COUNT(*) AS rows_count
FROM t_saas_ai_computing_token_consumption_record GROUP BY source_app;
SELECT source_app, settlement_source, settlement_status, COUNT(*) AS rows_count
FROM t_saas_ai_computing_token_consumption_record
GROUP BY source_app, settlement_source, settlement_status;
SELECT id, source_app, request_id
FROM t_saas_ai_computing_call
WHERE source_app IS NULL OR source_app NOT REGEXP '^[a-z0-9-]{2,50}$'
OR request_id IS NULL OR request_id NOT REGEXP '^[!-~]{8,100}$';
SELECT request_id, account_id, source_app, COUNT(*) AS rows_count
FROM t_saas_ai_computing_token_consumption_record
WHERE request_id IS NOT NULL
GROUP BY request_id, account_id, source_app HAVING COUNT(*) > 1;
SELECT a.id, a.balance_point_units,
COALESCE(p.net_units, 0) AS recorded_net_units,
a.balance_point_units - COALESCE(p.net_units, 0) AS residual_units
FROM t_saas_ai_merchant_computing_account a
LEFT JOIN (
SELECT account_id, SUM(point_units) AS net_units
FROM t_saas_ai_merchant_computing_point_record
WHERE type = 'CONSUME' OR (type = 'GIFT' AND sync_status = 'SUCCESS')
GROUP BY account_id
) p ON p.account_id = a.id
WHERE a.balance_point_units <> COALESCE(p.net_units, 0);
-- Negative consumed_quota must surface here as its own anomaly row; never mask
-- it by clamping to zero before the multiplication check.
SELECT id, account_id, settlement_status, consumed_quota, consumed_point_units
FROM t_saas_ai_computing_token_consumption_record
WHERE settlement_status <> 'SUCCESS' OR consumed_quota IS NULL
OR consumed_point_units IS NULL OR point_units_per_quota IS NULL
OR consumed_quota < 0
OR consumed_point_units <> consumed_quota * point_units_per_quota;
SELECT id, account_id, status, gateway_request_id, reconciled_time
FROM t_saas_ai_computing_call WHERE status NOT IN ('SUCCESS', 'FAILED');
SELECT merchant_id, active_gift_id
FROM t_saas_ai_computing_account_gate WHERE active_gift_id IS NOT NULL;
SELECT generator_key, next_id
FROM t_saas_ai_computing_id_segment;
-- The persisted next_id already covers allocated but unused segments. Never replace it by MAX(id)+1.
SELECT 'invalid_high_water' AS anomaly, generator_key, next_id
FROM t_saas_ai_computing_id_segment
WHERE generator_key <> 'computing-global'
OR next_id < 100000000000000000 OR next_id > 1000000000000000000;
-- P1 target checks: read-only, run manually on the NEW independent schema.
-- No data import, balance adjustment, ID seed/reset, credential grants or gateway operations.
-- Empty anomaly results do not prove source-to-target equality; compare the source inventory separately.
SELECT DATABASE() AS target_database, VERSION() AS mysql_version,
@@session.transaction_isolation AS isolation_level, @@session.time_zone AS time_zone;
SELECT table_name, table_rows
FROM information_schema.tables
WHERE table_schema = DATABASE() AND table_name LIKE 't_computing_%'
ORDER BY table_name;
SELECT id, balance_point_units, COALESCE(p.net_units, 0) AS recorded_net_units,
balance_point_units - COALESCE(p.net_units, 0) AS residual_units
FROM t_computing_account a
LEFT JOIN (
SELECT account_id, SUM(point_units) AS net_units
FROM t_computing_point_record GROUP BY account_id
) p ON p.account_id = a.id
WHERE balance_point_units <> COALESCE(p.net_units, 0);
SELECT 'call_consumption_mismatch' AS anomaly, c.id, c.account_id, c.consumption_record_id
FROM t_computing_call c
LEFT JOIN t_computing_consumption r ON r.id = c.consumption_record_id
WHERE c.billing_status = 'SUCCESS' AND (
r.id IS NULL OR NOT (r.call_id <=> c.id) OR NOT (r.account_id <=> c.account_id)
OR NOT (r.client_id <=> c.client_id) OR NOT (r.settlement_status <=> 'SUCCESS')
OR NOT (r.consumed_quota <=> c.consumed_quota)
OR NOT (r.point_units_per_quota <=> c.point_units_per_quota)
OR NOT (r.consumed_point_units <=> c.consumed_point_units)
);
SELECT 'consumption_call_mismatch' AS anomaly, r.id, r.account_id, r.call_id
FROM t_computing_consumption r
LEFT JOIN t_computing_call c ON c.id = r.call_id
WHERE r.settlement_status = 'SUCCESS' AND r.call_id IS NOT NULL AND (
c.id IS NULL OR NOT (c.consumption_record_id <=> r.id)
OR NOT (c.account_id <=> r.account_id) OR NOT (c.client_id <=> r.client_id)
OR NOT (c.billing_status <=> 'SUCCESS') OR NOT (c.settled_flag <=> 1)
OR NOT (c.consumed_quota <=> r.consumed_quota)
OR NOT (c.point_units_per_quota <=> r.point_units_per_quota)
OR NOT (c.consumed_point_units <=> r.consumed_point_units)
);
SELECT 'consumption_point_mismatch' AS anomaly, r.id, r.consumed_point_units, p.id AS point_record_id
FROM t_computing_consumption r
LEFT JOIN t_computing_point_record p ON p.consumption_record_id = r.id
WHERE r.settlement_status = 'SUCCESS' AND (
(r.consumed_point_units > 0 AND (p.id IS NULL OR p.point_units <> -r.consumed_point_units))
OR (r.consumed_point_units = 0 AND p.id IS NOT NULL)
);
SELECT 'point_consumption_mismatch' AS anomaly, p.id, p.call_id, p.consumption_record_id
FROM t_computing_point_record p
LEFT JOIN t_computing_consumption r ON r.id = p.consumption_record_id
WHERE p.type = 'CONSUME' AND (
r.id IS NULL OR NOT (r.settlement_status <=> 'SUCCESS')
OR NOT (p.account_id <=> r.account_id) OR NOT (p.client_id <=> r.client_id)
OR NOT (p.call_id <=> r.call_id) OR NOT (p.point_units <=> -r.consumed_point_units)
);
SELECT g.id, g.account_id, p.id AS point_record_id
FROM t_computing_gift g
LEFT JOIN t_computing_point_record p ON p.gift_id = g.id
WHERE g.status = 'SUCCESS' AND (p.id IS NULL OR p.point_units <> g.credited_point_units);
-- Reverse direction: a written GIFT point record must sit on a SUCCESS gift
-- with the credited amount; records anchored to non-terminal gifts are anomalies.
SELECT 'gift_point_record_terminal_mismatch' AS anomaly, p.id, p.gift_id, g.status
FROM t_computing_point_record p
LEFT JOIN t_computing_gift g ON g.id = p.gift_id
WHERE p.type = 'GIFT' AND p.gift_id IS NOT NULL AND p.legacy_record = 0
AND (g.id IS NULL OR g.status <> 'SUCCESS' OR p.point_units <> g.credited_point_units);
SELECT a.id, a.gift_gate, g.status
FROM t_computing_account a
LEFT JOIN t_computing_gift g ON g.id = a.gift_gate
WHERE a.gift_gate IS NOT NULL AND (g.id IS NULL OR g.account_id <> a.id OR g.status = 'SUCCESS');
SELECT id, account_id, billing_status, execution_status
FROM t_computing_call WHERE billing_status NOT IN ('SUCCESS', 'FAILED');
SELECT generator_key, next_id FROM t_computing_id_segment;
SELECT 'missing_generator' AS anomaly
WHERE NOT EXISTS (SELECT 1 FROM t_computing_id_segment WHERE generator_key = 'computing-global');
-- Historical IDs at or above the high water mark would collide with future allocations.
SELECT 'high_water_collision' AS anomaly, t.table_name, t.max_id, s.next_id
FROM (
SELECT 't_computing_account' AS table_name, MAX(id) AS max_id FROM t_computing_account
UNION ALL SELECT 't_computing_credential', MAX(id) FROM t_computing_credential
UNION ALL SELECT 't_computing_gift', MAX(id) FROM t_computing_gift
UNION ALL SELECT 't_computing_call', MAX(id) FROM t_computing_call
UNION ALL SELECT 't_computing_consumption', MAX(id) FROM t_computing_consumption
UNION ALL SELECT 't_computing_point_record', MAX(id) FROM t_computing_point_record
UNION ALL SELECT 't_computing_management_request', MAX(id) FROM t_computing_management_request
UNION ALL SELECT 't_computing_operation_audit', MAX(id) FROM t_computing_operation_audit
) t
JOIN t_computing_id_segment s ON s.generator_key = 'computing-global'
WHERE t.max_id IS NOT NULL AND t.max_id >= s.next_id;
SELECT c.id, c.category
FROM t_computing_credential c
LEFT JOIN t_computing_account_client a ON a.account_id = c.account_id AND a.client_id = c.client_id
WHERE c.category = 'CALL' AND (a.account_id IS NULL OR a.status <> 'AUTHORIZED');
-- P5 SaaS 侧交付(2026-09-23):手动执行,不由 Python 服务自动执行。
-- 1) 商户算力服务绑定表:SaaS 商户与 Python 独立算力账户(18 位数值 ID)一对一绑定,
-- 保存开户恢复请求号与 CALL Key AES-GCM 密文(base64,密钥来自部署 Secret,不入库明文)。
-- 2) 美际分析表补充 computing_request_id:稳定业务请求号 meiji-skin-analysis:{analysisId},
-- 用于按原请求号回查 /api/v1/calls/{request_id},不换号重发。
-- 注意:阶段 0 脚本 20260922_computing_stage0_ddl.sql 已含同一列增量;若该脚本已在目标库执行,
-- 本段 ALTER 会因列重复失败,跳过即可(仅执行第 1 段建表)。
-- 遵循算力迁移 ID 规则:业务数值 ID 为 18 位、显式插入、无自增;本表主键沿用 SaaS 既有发号器生成的 BIGINT。
CREATE TABLE `t_saas_ai_computing_service_binding`
(
`id` bigint(20) NOT NULL,
`create_timestamp` bigint(20) DEFAULT NULL,
`last_update_timestamp` bigint(20) DEFAULT NULL,
`merchant_id` bigint(20) NOT NULL COMMENT 'SaaS 商户 ID',
`account_id` bigint(20) DEFAULT NULL COMMENT 'Python 算力账户 18 位数值 ID,开户确认后写入',
`provision_request_id` varchar(100) NOT NULL COMMENT '开户/恢复使用的稳定请求号,响应丢失时按原号查询,不换号重发',
`call_credential_id` bigint(20) DEFAULT NULL COMMENT 'Python CALL 凭证 ID(十进制字符串精确转 Long)',
`call_key_cipher` varchar(512) DEFAULT NULL COMMENT 'CALL Key AES-GCM 密文(base64),密钥由部署 Secret 注入,不落明文',
`call_key_mask` varchar(100) DEFAULT NULL COMMENT 'CALL Key 脱敏展示,仅保留前后片段',
`cipher_key_version` int(11) NOT NULL DEFAULT 1 COMMENT '密文密钥版本,用于密钥轮换',
`provision_status` varchar(30) NOT NULL COMMENT 'PENDING/READY/NEEDS_RECONCILIATION/FAILED',
`provision_error` varchar(500) DEFAULT NULL COMMENT '最近一次开户/发证/激活失败摘要',
`create_time` datetime DEFAULT NULL,
PRIMARY KEY (`id`),
UNIQUE KEY `uk_merchant` (`merchant_id`),
UNIQUE KEY `uk_provision_request_id` (`provision_request_id`),
KEY `idx_account_id` (`account_id`)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_general_ci COMMENT='商户与Python算力服务账户绑定';
ALTER TABLE `t_saas_mb_meiji_skin_ai_analysis`
ADD COLUMN `computing_request_id` varchar(100) DEFAULT NULL COMMENT '算力服务稳定请求号 meiji-skin-analysis:{analysisId},回补与人工核对按原号查询' AFTER `gateway_request_id`,
ADD KEY `idx_computing_request_id` (`computing_request_id`);
import argparse
import json
import os
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from sqlalchemy.orm import sessionmaker
from app.account_bootstrap import bootstrap_client
from app.account_config import load_account_settings
from app.account_models import IdSegment
from app.db import create_mysql_engine
from app.id_generator import SegmentIDGenerator
def write_delivery(path, data):
target = Path(path).resolve()
root = Path(__file__).resolve().parents[1]
if target == root or root in target.parents:
raise ValueError("交付文件必须位于项目目录外")
descriptor = os.open(target, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600)
with os.fdopen(descriptor, "w", encoding="utf-8") as stream:
json.dump(data, stream, ensure_ascii=True)
stream.write("\n")
stream.flush()
os.fsync(stream.fileno())
directory = os.open(target.parent, os.O_RDONLY)
try:
os.fsync(directory)
finally:
os.close(directory)
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--config", required=True)
parser.add_argument("--client-id", required=True)
parser.add_argument("--name", required=True)
parser.add_argument("--category", choices=["INTEGRATION", "PLATFORM"], required=True)
parser.add_argument("--request-id", required=True)
parser.add_argument("--output", required=True)
args = parser.parse_args()
engine = None
try:
settings = load_account_settings(args.config)
engine = create_mysql_engine(settings.database_url, pool_size=settings.db_pool_size,
max_overflow=settings.db_max_overflow)
factory = sessionmaker(engine)
generator = SegmentIDGenerator(factory, segment_model=IdSegment)
result = bootstrap_client(factory, generator, client_id=args.client_id, name=args.name,
category=args.category, request_id=args.request_id,
deliver_secret=lambda data: write_delivery(args.output, data))
print(json.dumps(result))
except Exception:
parser.exit(1, "初始化未确认完成;请保留交付文件并核查原请求,勿换号重试或覆盖文件\n")
finally:
if engine is not None:
engine.dispose()
if __name__ == "__main__":
main()
-- 算力账户只读对账脚本(migration.md 阶段 1 验收判据,design.md §6.11/§7.9)。
--
-- 用途:在联调/测试环境用真实 MySQL 数据校验 Python 只读查询结果。
-- 约束:只读 SELECT,不修正数据;输出差异需人工追溯 source_app + request_id。
-- 执行:mysql -h <host> -u <user> -p <db> < scripts/reconciliation_check.sql
-- 注意:需要先发布阶段 0 DDL(source_app/settlement_source/call_id 等增强列),
-- 旧表结构下部分查询会因缺列失败,属预期。
-- 1. 账户权威余额 vs 流水恒等式:
-- balance_point_units ?= 成功赠送子单位 + CONSUME 流水子单位之和(流水为负,
-- 负余额合法)。差异需逐笔追溯。
SELECT
a.merchant_id,
a.balance_point_units,
COALESCE(g.gifted_units, 0) AS gifted_units,
COALESCE(c.consumed_units, 0) AS consumed_units,
a.balance_point_units
- COALESCE(g.gifted_units, 0)
- COALESCE(c.consumed_units, 0) AS residual_units
FROM t_saas_ai_merchant_computing_account a
LEFT JOIN (
SELECT merchant_id, SUM(point_units) AS gifted_units
FROM t_saas_ai_merchant_computing_point_record
WHERE type = 'GIFT' AND sync_status = 'SUCCESS'
GROUP BY merchant_id
) g ON g.merchant_id = a.merchant_id
LEFT JOIN (
SELECT merchant_id, SUM(point_units) AS consumed_units
FROM t_saas_ai_merchant_computing_point_record
WHERE type = 'CONSUME'
GROUP BY merchant_id
) c ON c.merchant_id = a.merchant_id
WHERE a.account_type = 'MASTER'
HAVING residual_units <> 0;
-- 2. 已结算正消耗必须有唯一对应积分流水(零消耗不写流水,不算缺失)。
SELECT r.id AS consumption_id, r.merchant_id, r.request_id, r.source_app
FROM t_saas_ai_computing_token_consumption_record r
LEFT JOIN t_saas_ai_merchant_computing_point_record p
ON p.consumption_record_id = r.id
WHERE r.settlement_status = 'SUCCESS'
AND r.consumed_point_units > 0
AND p.id IS NULL;
-- 3. 积分流水反向核对:CONSUME 流水必须指向已结算消耗记录。
SELECT p.id AS point_record_id, p.merchant_id, p.request_id, p.source_app
FROM t_saas_ai_merchant_computing_point_record p
LEFT JOIN t_saas_ai_computing_token_consumption_record r
ON r.id = p.consumption_record_id
WHERE p.type = 'CONSUME'
AND (r.id IS NULL OR r.settlement_status <> 'SUCCESS');
-- 4. 消耗记录子单位换算一致性:consumed_point_units ?= max(quota,0) * point_units_per_quota。
SELECT id, merchant_id, request_id, source_app,
consumed_quota, consumed_point_units, point_units_per_quota
FROM t_saas_ai_computing_token_consumption_record
WHERE settlement_status = 'SUCCESS'
AND consumed_point_units <> COALESCE(GREATEST(consumed_quota, 0), 0) * point_units_per_quota;
-- 5. 赠送换算一致性:point_units ?= quota_delta * point_units_per_quota(仅成功同步)。
SELECT id, merchant_id, request_id, source_app,
quota_delta, point_units, point_units_per_quota
FROM t_saas_ai_merchant_computing_point_record
WHERE type = 'GIFT' AND sync_status = 'SUCCESS'
AND point_units <> COALESCE(quota_delta, 0) * point_units_per_quota;
-- 6. 幂等与唯一性抽查:
-- 同一 source_app + request_id 只能有一条消耗记录 / 一条 GIFT 流水。
SELECT source_app, request_id, COUNT(*) AS cnt
FROM t_saas_ai_computing_token_consumption_record
WHERE request_id IS NOT NULL
GROUP BY source_app, request_id
HAVING cnt > 1;
SELECT source_app, request_id, COUNT(*) AS cnt
FROM t_saas_ai_merchant_computing_point_record
WHERE type = 'GIFT' AND request_id IS NOT NULL
GROUP BY source_app, request_id
HAVING cnt > 1;
-- 7. 来源分布与状态分布(人工比对 Python 查询接口输出)。
SELECT source_app, settlement_status, COUNT(*) AS cnt,
SUM(COALESCE(consumed_point_units, 0)) AS total_units
FROM t_saas_ai_computing_token_consumption_record
GROUP BY source_app, settlement_status;
-- 8. 混合来源明细抽样(与 POST /merchant/consumptions/page 输出逐字段比对)。
SELECT id, merchant_id, account_id, store_id, employee_id, employee_account,
business_code, model, request_id, gateway_request_id, source_app,
consumed_quota, consumed_point_units, point_units_per_quota,
settlement_status, settlement_source,
input_tokens, cache_hit_input_tokens, output_tokens, total_tokens, create_time
FROM t_saas_ai_computing_token_consumption_record
ORDER BY create_time DESC, id DESC
LIMIT 50;
import argparse
import sys
from pathlib import Path
from sqlalchemy.dialects.mysql import dialect
from sqlalchemy.schema import CreateIndex, CreateTable
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from app.account_models import Base
HEADER = """-- P1 independent account schema; generated by scripts/render_account_ddl.py.
-- Target: a NEW, explicitly selected computing schema, MySQL >= 8.0.16.
-- Review and execute manually. Do not run on the existing SaaS schema.
-- No CREATE DATABASE/USE, no account grants, no source data changes, no ID seed.
-- No IF NOT EXISTS: existing tables must stop execution; never continue with --force.
-- After import, copy the committed source ID high-water mark before issuing IDs.
-- For a genuinely empty installation only, initialize computing-global at 100000000000000000.
-- Retrying this DDL after partial failure requires inspecting the target, not overwriting it.
-- Rollback: keep the old schema untouched; discard this target only before any new writes.
-- Once IDs or gateway side effects exist, reconcile before any rollback; never rewind IDs.
"""
def render_ddl():
statements = []
mysql = dialect()
for table in Base.metadata.sorted_tables:
statements.append(str(CreateTable(table).compile(dialect=mysql)).strip() + ";")
for index in sorted(table.indexes, key=lambda item: item.name):
statements.append(str(CreateIndex(index).compile(dialect=mysql)).strip() + ";")
ddl = HEADER + "\n\n".join(statements)
return "\n".join(line.rstrip() for line in ddl.splitlines()) + "\n"
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--output", type=Path)
args = parser.parse_args()
ddl = render_ddl()
if args.output:
args.output.write_text(ddl, encoding="utf-8")
else:
print(ddl, end="")
if __name__ == "__main__":
main()
"""mei1-computing-service 测试包。"""
"""pytest 公共夹具:sqlite 内存库 + TestClient。
测试不触发 lifespan(TestClient 不进 with 块),app.state 由夹具手动装配;
真实启动路径(load_settings/validate/init_engine)由 test_auth 单独覆盖。
主数据沿用小 ID;算力业务表使用显式 18 位 ID(design.md §9.6)。
"""
from datetime import date, datetime
import pytest
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import Session, sessionmaker
from sqlalchemy.pool import StaticPool
from app.auth import DisabledAuthProvider
from app.config import Settings
from app.db import get_session
from app.id_generator import SegmentIDGenerator, seed_segment
from app.main import create_app
from app.models import (
Base,
ComputingTokenConsumptionRecord,
Employee,
Merchant,
MerchantComputingAccount,
MerchantComputingPointRecord,
Store,
)
def pytest_addoption(parser):
parser.addoption("--account-mysql-socket", default=None, help="P1 隔离临时 MySQL 的 Unix socket")
AUTH_HEADERS = {"X-App-Id": "mei1-saas", "X-Request-Id": "req-test-1"}
ID_A = 100000000000000011
ID_B = 100000000000000012
ID_P1 = 100000000000000101
ID_P2 = 100000000000000102
ID_P3 = 100000000000000103
ID_P4 = 100000000000000104
ID_C1 = 100000000000000201
ID_C2 = 100000000000000202
ID_C3 = 100000000000000203
def make_settings(**overrides) -> Settings:
values = dict(
app_env="local",
app_name="mei1-computing-service-test",
log_level="WARNING",
auth_mode="disabled",
database_url="sqlite://",
db_pool_size=5,
db_max_overflow=0,
)
values.update(overrides)
return Settings(**values)
def seed(session: Session) -> None:
"""覆盖契约测试所需的最小数据集:3 商户(1003 未开通)、账户、流水、门店、员工。"""
session.add_all([
Merchant(id=1001, merchant_info="M1001", name="美际一店", status=1),
Merchant(id=1002, merchant_info="M1002", name="美际二店", status=1),
Merchant(id=1003, merchant_info="M1003", name="美际三店", status=1),
MerchantComputingAccount(
id=ID_A, merchant_id=1001, account_type="MASTER", account_code="MASTER",
name="MASTER", balance_point_units=5000000, status="ACTIVE", sync_status="SUCCESS",
gateway_token_mask="sk-****abcd",
),
MerchantComputingAccount(
id=ID_B, merchant_id=1002, account_type="MASTER", account_code="MASTER",
name="MASTER", balance_point_units=0, status="ACTIVE", sync_status="SUCCESS",
),
# 商户 1001:成功赠送 999.9978 + octop 成功赠送 1.0000;失败赠送不累计。
MerchantComputingPointRecord(
id=ID_P1, merchant_id=1001, account_id=ID_A, type="GIFT",
requested_points=1000, point_units=9999978,
quota_before=0, quota_delta=68493, quota_after=68493, point_units_per_quota=146,
balance_before_units=0, balance_after_units=9999978,
source_app="mei1-saas", request_id="gift-1", sync_status="SUCCESS",
create_time=datetime(2026, 8, 5, 9, 0, 0),
),
MerchantComputingPointRecord(
id=ID_P2, merchant_id=1001, account_id=ID_A, type="GIFT",
requested_points=1, point_units=9928,
source_app="mei1-saas", request_id="gift-2", sync_status="FAILED", sync_error="网关超时",
create_time=datetime(2026, 8, 6, 9, 0, 0),
),
MerchantComputingPointRecord(
id=ID_P3, merchant_id=1001, account_id=ID_A, type="GIFT",
requested_points=1, point_units=10000,
source_app="octop", request_id="gift-3", sync_status="SUCCESS",
create_time=datetime(2026, 9, 10, 9, 0, 0), remark="活动赠送",
),
MerchantComputingPointRecord(
id=ID_P4, merchant_id=1002, account_id=ID_B, type="CONSUME",
point_units=-14600, consumption_record_id=ID_C2,
source_app="octop", request_id="req-2", sync_status="SUCCESS",
),
Store(id=21, merchant_id=1001, status=1),
Store(id=22, merchant_id=1001, status=2),
Store(id=23, merchant_id=1002, status=1),
Employee(id=31, merchant_id=1001, store_id=21, status=1, position_status=1),
Employee(id=32, merchant_id=1001, store_id=None, status=1, position_status=1),
Employee(id=33, merchant_id=1002, store_id=23, status=1, position_status=1),
Employee(id=34, merchant_id=1001, store_id=21, status=2, position_status=1),
Employee(
id=35, merchant_id=1001, store_id=21, status=1, position_status=2,
dimission_date=date(2026, 1, 1),
),
# 已结算正消耗(368 quota -> 53728 子单位)
ComputingTokenConsumptionRecord(
id=ID_C1, merchant_id=1001, account_id=ID_A, business_code="meiji_skin_report_analysis",
model="gpt-4o", input_tokens=100, cache_hit_input_tokens=20,
output_tokens=200, total_tokens=300,
consumed_quota=368, consumed_point_units=53728, point_units_per_quota=146,
settlement_status="SUCCESS", settlement_source="GATEWAY_LOG",
employee_id=31, employee_account="emp01",
source_app="mei1-saas", request_id="req-1", gateway_request_id="gw-1",
create_time=datetime(2026, 9, 1, 10, 0, 0),
),
# octop 来源、待结算:金额字段必须为 null
ComputingTokenConsumptionRecord(
id=ID_C2, merchant_id=1002, account_id=ID_B, business_code="meiji_skin_report_analysis",
model="gpt-4o-mini", input_tokens=10, output_tokens=20, total_tokens=30,
settlement_status="SETTLE_PENDING",
source_app="octop", request_id="req-2", gateway_request_id="gw-2",
create_time=datetime(2026, 9, 15, 10, 0, 0),
),
# 已结算零消耗:金额 "0"/"0.0000",不与未结算混淆(design.md §7.9.3)
ComputingTokenConsumptionRecord(
id=ID_C3, merchant_id=1001, account_id=ID_A, business_code="meiji_skin_report_analysis",
model="gpt-4o", input_tokens=5, output_tokens=0, total_tokens=5,
consumed_quota=0, consumed_point_units=0, point_units_per_quota=146,
settlement_status="SUCCESS", settlement_source="GATEWAY_LOG",
source_app="mei1-saas", request_id="req-3", gateway_request_id="gw-3",
create_time=datetime(2026, 9, 20, 10, 0, 0),
),
])
seed_segment(session)
session.flush()
@pytest.fixture()
def session():
engine = create_engine(
"sqlite://",
connect_args={"check_same_thread": False},
poolclass=StaticPool,
future=True,
)
Base.metadata.create_all(engine)
session = Session(engine, future=True)
seed(session)
yield session
session.close()
engine.dispose()
@pytest.fixture()
def client(session):
engine = session.get_bind()
app = create_app()
app.state.settings = make_settings()
app.state.auth_provider = DisabledAuthProvider()
app.state.id_generator = SegmentIDGenerator(
sessionmaker(bind=engine, expire_on_commit=False, future=True)
)
def override_session():
yield session
app.dependency_overrides[get_session] = override_session
return TestClient(app)
import hashlib
import json
import traceback
from dataclasses import FrozenInstanceError, replace
from pathlib import Path
import pytest
from app.account_config import AccountConfigError, AccountSettings, load_account_settings
def values(**overrides):
result = {
"database_url": "mysql+pymysql://test_user:db-secret@db.invalid/accounts",
"newapi_base_url": "https://gateway.invalid/",
"newapi_pat": "pat-secret",
"newapi_model": "provider/model-v1",
}
result.update(overrides)
return result
def write_config(tmp_path, data):
path = tmp_path / "account.json"
path.write_text(json.dumps(data), encoding="utf-8")
return path
def assert_redacted(error):
rendered = str(error) + repr(error) + "".join(traceback.format_exception_only(type(error), error))
assert rendered.count("Invalid account configuration") > 0
for secret in ("db-secret", "pat-secret", "gateway.invalid", "test_user", "secret-path"):
assert secret not in rendered
assert error.__suppress_context__
def test_external_json_defaults_and_frozen_redacted_settings(tmp_path):
settings = load_account_settings(write_config(tmp_path, values()))
settings.validate()
assert settings.db_pool_size == 5
assert settings.db_max_overflow == 0
assert settings.management_timeout_seconds == 30.0
assert settings.management_concurrency == 8
assert settings.connect_timeout_seconds == 10.0
assert settings.read_timeout_seconds == 30.0
for secret in values().values():
assert secret not in repr(settings)
assert secret not in str(settings)
with pytest.raises(FrozenInstanceError):
settings.newapi_pat = "changed"
def test_canonical_identity_is_stable_and_not_credential_dependent():
settings = AccountSettings(**values(newapi_base_url="HTTPS://Gateway.INVALID:443/prefix///"))
expected = hashlib.sha256(b"https://gateway.invalid/prefix").hexdigest()
assert settings.gateway_identity == expected
assert replace(settings, newapi_pat="different-pat", database_url="mysql+pymysql:///other").gateway_identity == expected
assert replace(settings, newapi_base_url="https://gateway.invalid/prefix").gateway_identity == expected
assert replace(settings, newapi_base_url="https://gateway.invalid:8443/prefix").gateway_identity != expected
def test_all_explicit_options_and_string_path(tmp_path):
data = values(db_pool_size=1, db_max_overflow=2, management_timeout_seconds=150,
management_concurrency=100, connect_timeout_seconds=10, read_timeout_seconds=120)
assert load_account_settings(str(write_config(tmp_path, data))) == AccountSettings(**data)
@pytest.mark.parametrize("key", ["database_url", "newapi_base_url", "newapi_pat", "newapi_model"])
def test_missing_required_keys_ignore_environment(tmp_path, monkeypatch, key):
data = values()
for name, value in data.items():
monkeypatch.setenv(name.upper(), value)
monkeypatch.setenv(name, value)
del data[key]
with pytest.raises(AccountConfigError) as caught:
load_account_settings(write_config(tmp_path, data))
assert_redacted(caught.value)
def test_environment_cannot_override_file_or_supply_defaults(tmp_path, monkeypatch):
for key in values():
monkeypatch.setenv(key.upper(), "wrong-env-secret")
monkeypatch.setenv("DB_POOL_SIZE", "999")
settings = load_account_settings(write_config(tmp_path, values()))
assert settings == AccountSettings(**values())
@pytest.mark.parametrize("data", [None, [], "pat-secret", 42, True,
values(unknown="pat-secret"), values(auth_mode="disabled")])
def test_reject_nonobjects_and_unknown_keys(tmp_path, data):
with pytest.raises(AccountConfigError) as caught:
load_account_settings(write_config(tmp_path, data))
assert_redacted(caught.value)
@pytest.mark.parametrize("content", [b'{"newapi_pat": "pat-secret",', b"\xffpat-secret", b"NaN",
b'{"newapi_pat":"pat-secret","newapi_pat":"other"}'])
def test_parse_errors_are_redacted(tmp_path, content):
path = tmp_path / "secret-path.json"
path.write_bytes(content)
with pytest.raises(AccountConfigError) as caught:
load_account_settings(path)
assert_redacted(caught.value)
formatted = "".join(traceback.format_exception(type(caught.value), caught.value, caught.value.__traceback__))
assert "JSONDecodeError" not in formatted
assert "UnicodeDecodeError" not in formatted
def test_missing_and_directory_paths_are_redacted(tmp_path):
for path in (tmp_path / "secret-path.json", tmp_path):
with pytest.raises(AccountConfigError) as caught:
load_account_settings(path)
assert_redacted(caught.value)
def test_project_files_and_symlink_targets_are_rejected_before_reading(tmp_path, monkeypatch):
internal = Path(__file__).resolve().parents[1] / "app" / "account_config.py"
link = tmp_path / "account.json"
link.symlink_to(internal)
opened = []
original = Path.open
def track_open(path, *args, **kwargs):
opened.append(path)
return original(path, *args, **kwargs)
monkeypatch.setattr(Path, "open", track_open)
for path in (internal, link):
with pytest.raises(AccountConfigError):
load_account_settings(path)
assert opened == []
@pytest.mark.parametrize("field,bad", [
("database_url", "sqlite:///db-secret"),
("database_url", "mysql://test_user:db-secret@db.invalid/accounts"),
("database_url", "mysql+pymysql:/accounts"),
("database_url", "mysql+pymysql://test_user:db-secret@db.invalid"),
("database_url", "mysql+pymysql://test_user:db-secret@db.invalid/"),
("database_url", "mysql+pymysql://test_user:db-secret@db.invalid/%20"),
("database_url", "mysql+pymysql://test_user:db-secret@db.invalid/a/b"),
("database_url", "mysql+pymysql://test_user:db-secret@db.invalid:bad/accounts"),
("database_url", "mysql+pymysql://test_user:db-secret@db.invalid/accounts#fragment"),
("newapi_base_url", "http://gateway.invalid"),
("newapi_base_url", "https:///gateway.invalid"),
("newapi_base_url", "https://test_user:pat-secret@gateway.invalid"),
("newapi_base_url", "https://@gateway.invalid"),
("newapi_base_url", "https://gateway.invalid?pat-secret"),
("newapi_base_url", "https://gateway.invalid?"),
("newapi_base_url", "https://gateway.invalid#pat-secret"),
("newapi_base_url", "https://gateway.invalid#"),
("newapi_base_url", "https://gateway.invalid:0"),
("newapi_base_url", "https://gateway.invalid:99999"),
("newapi_base_url", "https://gateway.invalid:"),
("newapi_base_url", "https://gate%20way.invalid"),
("newapi_base_url", "https://gate\"way.invalid"),
("newapi_base_url", "https://-gateway.invalid"),
("newapi_base_url", "https://gateway.invalid/../private"),
("newapi_base_url", "https://gateway.invalid/%2e%2e/private"),
("newapi_base_url", "https://[gateway.invalid"),
("newapi_base_url", " https://gateway.invalid"),
("newapi_base_url", "https://gate\nway.invalid"),
("newapi_base_url", "https://gateway.invalid\\other"),
("newapi_pat", ""), ("newapi_pat", " "), ("newapi_pat", "pat-secret\n"),
("newapi_pat", "pat secret"), ("newapi_pat", "pat-密钥"),
("newapi_model", ""), ("newapi_model", " "), ("newapi_model", "model one"),
("newapi_model", "model\n"), ("newapi_model", "one,two"),
("newapi_model", "模型"), ("newapi_model", "m" * 101),
("db_pool_size", 0), ("db_pool_size", -1), ("db_pool_size", 1.0),
("db_max_overflow", -1), ("db_max_overflow", 1.0),
("management_concurrency", 0), ("management_concurrency", 101),
("management_concurrency", 1.0),
("management_timeout_seconds", 0), ("management_timeout_seconds", -1),
("management_timeout_seconds", 150.01),
("connect_timeout_seconds", 0), ("connect_timeout_seconds", 10.01),
("read_timeout_seconds", 0), ("read_timeout_seconds", 120.01),
])
def test_invalid_settings_are_rejected_and_redacted(field, bad):
with pytest.raises(AccountConfigError) as caught:
AccountSettings(**values(**{field: bad})).validate()
assert_redacted(caught.value)
@pytest.mark.parametrize("field", list(values()))
@pytest.mark.parametrize("bad", [None, False, 123])
def test_credentials_must_be_strings(field, bad):
with pytest.raises(AccountConfigError):
AccountSettings(**values(**{field: bad})).validate()
@pytest.mark.parametrize("field", ["db_pool_size", "db_max_overflow", "management_concurrency",
"management_timeout_seconds", "connect_timeout_seconds", "read_timeout_seconds"])
@pytest.mark.parametrize("bad", [True, False, "5", None, float("inf"), float("nan")])
def test_numeric_types_and_finite_bounds(field, bad):
with pytest.raises(AccountConfigError):
AccountSettings(**values(**{field: bad})).validate()
def test_model_length_boundary_and_socket_database_url():
AccountSettings(**values(newapi_model="m" * 100,
database_url="mysql+pymysql://root@/accounts?unix_socket=%2Ftmp%2Ftest.sock")).validate()
def test_identity_validation_never_leaks_credentials():
with pytest.raises(AccountConfigError) as caught:
_ = AccountSettings(**values(newapi_base_url="https://test_user:pat-secret@gateway.invalid")).gateway_identity
assert_redacted(caught.value)
import pytest
from fastapi.testclient import TestClient
from app.account_model_gateway import _valid_params
from tests.test_account_calls import ANSWER, REQUEST_ID, funded, calling
from tests.test_account_foundation import account_engine
from tests.test_account_gifts import gift
from tests.test_account_management import headers, issue, activate, management, open_account
CHAT_PATH = "/v1/chat/completions"
PROMPT = "private-openai-prompt-not-for-storage"
def call_headers(m, *, key=None, request_id=None):
values = {"Authorization": "Bearer " + (key or m.secret)}
if request_id is not None:
values["X-Request-Id"] = request_id
return values
def openai_body(**overrides):
body = {"messages": [{"role": "user", "content": PROMPT}],
"temperature": 0.5, "top_p": 0.9, "max_tokens": 512, "unknown_field": {"x": 1}}
body.update(overrides)
return body
def test_openai_success_shape_usage_and_params_forwarded(funded):
m = funded
response = m.client.post(CHAT_PATH, json=openai_body(), headers=call_headers(m))
assert response.status_code == 200, response.text
data = response.json()
assert data["object"] == "chat.completion" and data["id"].startswith("chatcmpl-")
assert data["created"] is not None and data["model"] == "test-model"
assert data["choices"] == [{"index": 0, "message": {"role": "assistant", "content": ANSWER},
"finish_reason": "stop"}]
assert data["usage"] == {"prompt_tokens": 100, "completion_tokens": 200, "total_tokens": 300}
assert data["mei1"]["billingStatus"] == "SUCCESS" and data["mei1"]["consumedPoints"] is not None
assert len(data["mei1"]["requestId"]) >= 8 and data["mei1"]["callId"].isdigit()
# 白名单采样参数原样透传,未知字段不透传也不报错。
assert m.gateway.chat_params == {"temperature": 0.5, "top_p": 0.9, "max_tokens": 512}
assert m.gateway.dispatches[0]["messages"] == [{"role": "user", "content": PROMPT}]
# 结算落账口径与 /api/v1 相同。
with m.factory() as session:
from sqlalchemy import select
from app.account_models import Call
row = session.scalar(select(Call).order_by(Call.id.desc()))
assert row.billing_status == "SUCCESS" and row.consumed_point_units is not None
def test_openai_invalid_stream_options_rejected(funded):
m = funded
response = m.client.post(CHAT_PATH, json=openai_body(stream=True, stream_options={"include_usage": "yes"}),
headers=call_headers(m))
assert response.status_code == 400, response.text
error = response.json()["error"]
assert error["type"] == "invalid_request_error" and error["code"] == "INVALID_ARGUMENT"
assert not m.gateway.dispatches
def test_openai_developer_role_maps_to_system(funded):
m = funded
messages = [{"role": "developer", "content": "policy"}, {"role": "user", "content": PROMPT}]
response = m.client.post(CHAT_PATH, json={"messages": messages}, headers=call_headers(m))
assert response.status_code == 200, response.text
assert m.gateway.dispatches[0]["messages"] == [{"role": "system", "content": "policy"},
{"role": "user", "content": PROMPT}]
def test_openai_request_id_idempotent_and_fingerprint(funded):
m = funded
first = m.client.post(CHAT_PATH, json=openai_body(), headers=call_headers(m, request_id=REQUEST_ID))
assert first.status_code == 200, first.text
replay = m.client.post(CHAT_PATH, json=openai_body(), headers=call_headers(m, request_id=REQUEST_ID))
assert replay.status_code == 200, replay.text
assert replay.json()["mei1"]["callId"] == first.json()["mei1"]["callId"]
assert len(m.gateway.dispatches) == 1
conflict = m.client.post(CHAT_PATH, json=openai_body(temperature=0.6),
headers=call_headers(m, request_id=REQUEST_ID))
assert conflict.status_code == 409, conflict.text
assert conflict.json()["error"]["code"] == "FINGERPRINT_MISMATCH"
def test_openai_sampling_params_affect_fingerprint(funded):
m = funded
first = m.client.post(CHAT_PATH, json=openai_body(), headers=call_headers(m, request_id="openai-fp-1"))
assert first.status_code == 200, first.text
second = m.client.post(CHAT_PATH, json=openai_body(),
headers=call_headers(m, request_id="openai-fp-2"))
assert second.status_code == 200, second.text
assert m.gateway.chat_params == {"temperature": 0.5, "top_p": 0.9, "max_tokens": 512}
def test_openai_invalid_model_and_validation(funded):
m = funded
response = m.client.post(CHAT_PATH, json=openai_body(model="other-model"), headers=call_headers(m))
assert response.status_code == 400, response.text
assert response.json()["error"]["type"] == "invalid_request_error"
broken = m.client.post(CHAT_PATH, json={"messages": [{"role": "user", "content": ""}]},
headers=call_headers(m))
assert broken.status_code == 400, broken.text
assert broken.json()["error"]["type"] == "invalid_request_error"
def test_openai_auth_and_permission_shapes_openai_error(funded):
m = funded
missing = m.client.post(CHAT_PATH, json=openai_body())
assert missing.status_code == 401 and missing.json()["error"]["code"] == "AUTHENTICATION_FAILED"
management_key = m.client.post(CHAT_PATH, json=openai_body(), headers={"Authorization": "Bearer " + m.keys["client-a"]})
assert management_key.status_code == 403 and management_key.json()["error"]["code"] == "PERMISSION_DENIED"
models = m.client.get("/v1/models", headers=call_headers(m))
assert models.status_code == 200 and models.json()["object"] == "list"
assert models.json()["data"][0]["id"] == m.service.settings.newapi_model
def test_openai_insufficient_balance_maps_to_quota_error(calling):
m = calling
response = m.client.post(CHAT_PATH, json=openai_body(), headers=call_headers(m))
assert response.status_code == 429, response.text
error = response.json()["error"]
assert error["type"] == "insufficient_quota" and error["code"] == "ACCOUNT_BLOCKED"
def test_openai_nonstream_finish_reason_preserved_not_rewritten(funded):
m = funded
from dataclasses import replace
m.gateway.result = replace(m.gateway.result, finish_reason="length")
response = m.client.post(CHAT_PATH, json=openai_body(), headers=call_headers(m))
assert response.status_code == 200, response.text
assert response.json()["choices"][0]["finish_reason"] == "length"
def test_openai_execution_failure_returns_error_shape(funded):
m = funded
from dataclasses import replace
from app.account_model_gateway import ModelResult
m.gateway.result = replace(m.gateway.result, gateway_request_id=None, content=None, model=None,
input_tokens=None, cache_hit_input_tokens=None, output_tokens=None,
total_tokens=None, execution_status="FAILED", http_status=502)
response = m.client.post(CHAT_PATH, json=openai_body(), headers=call_headers(m))
assert 500 <= response.status_code < 600, response.text
error = response.json()["error"]
assert error["type"] == "api_error" and error["mei1"]["executionStatus"] == "FAILED"
def test_openai_nonstream_unknown_partial_content_not_returned_as_success(funded):
m = funded
from dataclasses import replace
m.gateway.result = replace(m.gateway.result, execution_status="UNKNOWN", error_code="HTTP_ERROR",
http_status=503)
response = m.client.post(CHAT_PATH, json=openai_body(), headers=call_headers(m))
assert response.status_code == 503, response.text
error = response.json()["error"]
assert error["type"] == "api_error" and error["code"] == "HTTP_ERROR"
assert error["mei1"]["executionStatus"] == "UNKNOWN"
assert m.gateway.dispatches and m.gateway.queries
def test_openai_nonstream_settle_pending_still_returns_success(funded):
m = funded
m.gateway.logs_visible = False
response = m.client.post(CHAT_PATH, json=openai_body(), headers=call_headers(m))
assert response.status_code == 200, response.text
data = response.json()
assert data["choices"][0]["message"]["content"] == ANSWER
assert data["mei1"]["billingStatus"] == "SETTLE_PENDING"
def test_valid_params_whitelist_bounds():
assert _valid_params({"temperature": 0, "top_p": 1.0, "max_tokens": 10})
assert _valid_params({"stop": "done"})
assert _valid_params({"stop": ["a", "b"], "seed": 7, "presence_penalty": -2, "frequency_penalty": 2})
assert not _valid_params({})
assert not _valid_params({"model": "x", "messages": [], "stream": True})
assert not _valid_params({"temperature": 3})
assert not _valid_params({"max_tokens": -1})
assert not _valid_params({"seed": True})
assert not _valid_params({"user": "u1"})
assert not _valid_params({"stop": ["a"] * 5})
assert not _valid_params({"temperature": "0.5"})
assert not _valid_params({"top_p": float("nan")})
import asyncio
from types import SimpleNamespace
import pytest
from app import account_worker as worker
from app.account_config import AccountSettings
@pytest.mark.parametrize("failure", [None, "database", "settlement", "cleanup"])
def test_worker_single_batch_lifecycle(monkeypatch, failure):
events = []
settings = AccountSettings(database_url="mysql+pymysql://user@localhost/independent",
newapi_base_url="https://gateway.invalid", newapi_pat="test-pat",
newapi_model="test-model")
engine = SimpleNamespace(dispose=lambda: events.append("dispose"))
monkeypatch.setattr(worker, "create_mysql_engine", lambda *args, **kwargs: engine)
def ready(value):
assert value is engine
events.append("ready")
if failure == "database":
raise RuntimeError("unavailable")
class Gateway:
def __init__(self, value):
assert value is settings
events.append("gateway")
async def aclose(self):
events.append("close")
class Settlement:
def __init__(self, accounts):
assert accounts.settings is settings
assert isinstance(accounts.gateway, Gateway)
async def run_once(self, batch_size):
assert batch_size == 20
events.append("settlement")
if failure == "settlement":
raise RuntimeError("unavailable")
return 3
def clear_results(self, batch_size):
assert batch_size == 20
events.append("cleanup")
if failure == "cleanup":
raise RuntimeError("unavailable")
return 2
monkeypatch.setattr(worker, "_database_ready", ready)
monkeypatch.setattr(worker, "ModelGateway", Gateway)
monkeypatch.setattr(worker, "SettlementService", Settlement)
if failure:
with pytest.raises(RuntimeError):
asyncio.run(worker.run(settings, 20))
else:
assert asyncio.run(worker.run(settings, 20)) == (3, 2)
assert events[-1] == "dispose"
if failure != "database":
assert events[-2] == "close"
assert events.count("settlement") <= 1
def test_worker_cli_rejects_unbounded_batch_before_loading_config(monkeypatch):
monkeypatch.setattr("sys.argv", ["account_worker", "--config", "/not/read.json", "--batch-size", "201"])
monkeypatch.setattr(worker, "load_account_settings", lambda path: pytest.fail("Config must not be read"))
with pytest.raises(SystemExit) as result:
worker.main()
assert result.value.code == 2
def test_api_cli_rejects_public_host_before_loading_config(monkeypatch):
from app import account_main
monkeypatch.setattr("sys.argv", ["account_main", "--config", "/not/read.json", "--host", "8.8.8.8"])
monkeypatch.setattr(account_main, "load_account_settings",
lambda path: pytest.fail("Config must not be read"))
with pytest.raises(SystemExit) as result:
account_main.main()
assert result.value.code == 2
def test_api_cli_accepts_container_bind_and_passes_settings(monkeypatch):
from app import account_main
settings = SimpleNamespace()
monkeypatch.setattr("sys.argv", ["account_main", "--config", "/not/read.json",
"--host", "0.0.0.0", "--port", "8000"])
monkeypatch.setattr(account_main, "load_account_settings", lambda path: settings)
captured = {}
def run(app, *, host, port, access_log):
captured.update(app=app, host=host, port=port)
monkeypatch.setattr("uvicorn.run", run)
account_main.main()
assert captured["host"] == "0.0.0.0" and captured["port"] == 8000
"""AuthContext / disabled 提供器 / 启动校验单元测试(design.md §5.1、authentication.md §4)。"""
import pytest
from app.auth import AuthContext, DisabledAuthProvider, get_auth_provider
from app.config import Settings, StartupError, load_settings
from app.errors import AppError
PROVIDER = DisabledAuthProvider()
def test_load_settings_defaults():
settings = load_settings(environ={})
assert settings.app_env == "local"
assert settings.auth_mode == "disabled"
def test_disabled_provider_ok():
auth = PROVIDER.from_headers({"x-app-id": "mei1-saas", "x-request-id": "req-1"})
assert auth.app_id == "mei1-saas"
assert auth.request_id == "req-1"
assert auth.allows_scope("account:read")
assert auth.allows_scope("opt:gift:read")
assert not auth.allows_scope("unknown:scope")
def test_disabled_provider_missing_headers():
with pytest.raises(AppError) as exc:
PROVIDER.from_headers({})
assert exc.value.code == "AUTH_HEADER_MISSING"
assert exc.value.http_status == 401
with pytest.raises(AppError) as exc:
PROVIDER.from_headers({"x-app-id": "mei1-saas"})
assert exc.value.code == "AUTH_HEADER_MISSING"
def test_disabled_provider_invalid_format():
with pytest.raises(AppError) as exc:
PROVIDER.from_headers({"x-app-id": "has space", "x-request-id": "req-1"})
assert exc.value.code == "AUTH_HEADER_INVALID"
with pytest.raises(AppError) as exc:
PROVIDER.from_headers({"x-app-id": "mei1-saas", "x-request-id": "r" * 101})
assert exc.value.code == "AUTH_HEADER_INVALID"
def test_auth_context_allows_merchant():
auth_all = AuthContext(app_id="a", request_id="r")
assert auth_all.allows_merchant(1001)
auth_bound = AuthContext(app_id="a", request_id="r", merchant_scope="BOUND",
bound_merchant_ids=frozenset({1001}))
assert auth_bound.allows_merchant(1001)
assert not auth_bound.allows_merchant(1002)
assert not auth_bound.allows_merchant(None)
def test_auth_context_business_code_whitelist():
"""businessCodes 白名单:None 为未配置(disabled 放行),配置后按白名单校验。"""
auth_open = AuthContext(app_id="a", request_id="r")
assert auth_open.allows_business_code("any_business")
auth_bounded = AuthContext(
app_id="a", request_id="r",
business_codes=frozenset({"meiji_skin_report_analysis"}),
)
assert auth_bounded.allows_business_code("meiji_skin_report_analysis")
assert not auth_bounded.allows_business_code("other_business")
def test_rsa_provider_not_implemented():
with pytest.raises(RuntimeError, match="阶段 4"):
get_auth_provider("rsa")
def _settings(app_env="local", auth_mode="disabled") -> Settings:
return Settings(
app_env=app_env, app_name="t", log_level="INFO", auth_mode=auth_mode,
database_url="sqlite://", db_pool_size=1, db_max_overflow=0,
)
def test_validate_local_disabled_ok():
_settings().validate()
def test_validate_test_disabled_ok():
_settings(app_env="test").validate()
def test_validate_rsa_ok_in_prod():
_settings(app_env="prod", auth_mode="rsa").validate()
@pytest.mark.parametrize("app_env", ["test2", "test3", "test4", "prod", "staging", "uat"])
def test_validate_rejects_disabled_outside_whitelist(app_env):
"""disabled 白名单仅 local/test;未枚举环境一律拒绝(design.md §12)。"""
with pytest.raises(StartupError):
_settings(app_env=app_env, auth_mode="disabled").validate()
def test_validate_rejects_unknown_mode():
with pytest.raises(StartupError):
_settings(auth_mode="none").validate()
"""固定计费基线单元测试(design.md §6.6/§6.9 固定向量)。"""
import pytest
from decimal import Decimal
from app.billing import (
EXPECTED_QUOTA_PER_UNIT,
POINT_SCALE,
POINT_UNITS_PER_QUOTA,
format_points,
point_units_to_quota_half_up,
quota_to_point_units,
requested_points_to_units,
require_bigint,
units_json,
units_to_points,
)
def test_fixed_constants():
assert POINT_SCALE == 10000
assert POINT_UNITS_PER_QUOTA == 146
assert EXPECTED_QUOTA_PER_UNIT == 500000
@pytest.mark.parametrize("quota,units,points", [
(0, 0, "0.0000"),
(-5, 0, "0.0000"),
(1, 146, "0.0146"),
(368, 53728, "5.3728"),
(999, 145854, "14.5854"),
(1000, 146000, "14.6000"),
(1001, 146146, "14.6146"),
])
def test_quota_to_point_units_vectors(quota, units, points):
assert quota_to_point_units(quota) == units
assert format_points(units) == points
def test_no_per_call_ceiling():
"""无逐笔整分向上取整:1 quota 扣 0.0146 而不是 1 分。"""
assert quota_to_point_units(1) == 146
assert format_points(quota_to_point_units(1)) == "0.0146"
def test_gift_half_up_vectors():
assert requested_points_to_units(1000) == 10000000
assert point_units_to_quota_half_up(10000000) == 68493
assert 68493 * 146 == 9999978
assert format_points(9999978) == "999.9978"
assert requested_points_to_units(1) == 10000
assert point_units_to_quota_half_up(10000) == 68
assert 68 * 146 == 9928
@pytest.mark.parametrize("units,quota", [(0, 0), (72, 0), (73, 1), (74, 1), (145, 1), (146, 1), (147, 1), (219, 2), (220, 2)])
def test_half_up_boundaries(units, quota):
assert point_units_to_quota_half_up(units) == quota
def test_units_to_points_decimal():
assert units_to_points(-46) == Decimal("-0.0046")
assert format_points(-46) == "-0.0046"
def test_units_json_and_null():
assert units_json(53728) == "53728"
assert units_json(0) == "0"
assert units_json(None) is None
assert format_points(None) is None
def test_require_bigint_bounds():
assert require_bigint(2**63 - 1) == 2**63 - 1
with pytest.raises(OverflowError):
require_bigint(2**63)
with pytest.raises(OverflowError):
require_bigint(-(2**63) - 1)
"""18 位号段 ID 生成器单元测试(design.md §9.6)。"""
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import Session
from sqlalchemy.pool import StaticPool
from app.errors import AppError
from app.id_generator import (
GENERATOR_KEY,
ID_HIGH_EXCLUSIVE,
ID_LOW,
SEGMENT_SIZE,
SegmentIDGenerator,
seed_segment,
)
from app.models import Base, IdSegment
@pytest.fixture()
def engine():
engine = create_engine(
"sqlite://",
connect_args={"check_same_thread": False},
poolclass=StaticPool,
future=True,
)
Base.metadata.create_all(engine)
with Session(engine, future=True) as session:
seed_segment(session)
session.commit()
yield engine
engine.dispose()
def _high_watermark(engine):
with Session(engine, future=True) as session:
return session.get(IdSegment, GENERATOR_KEY).next_id
def test_ids_are_fixed_18_digits(engine):
gen = SegmentIDGenerator(lambda: Session(engine, future=True))
ids = [gen.next_id() for _ in range(10)]
assert ids[0] == ID_LOW
assert all(isinstance(v, int) for v in ids)
assert all(10**17 <= v < 10**18 for v in ids)
assert ids == sorted(ids)
assert len(set(ids)) == 10
def test_batch_allocation_advances_watermark(engine):
gen = SegmentIDGenerator(lambda: Session(engine, future=True))
gen.next_id()
assert _high_watermark(engine) == ID_LOW + SEGMENT_SIZE
def test_next_generator_continues_without_overlap(engine):
"""新实例(模拟重启/第二个 worker)从数据库高水位继续,不与旧候选段重叠。"""
gen1 = SegmentIDGenerator(lambda: Session(engine, future=True))
first = gen1.next_id()
gen2 = SegmentIDGenerator(lambda: Session(engine, future=True))
second = gen2.next_id()
assert second == first + SEGMENT_SIZE
assert second >= ID_LOW + SEGMENT_SIZE
def test_last_segment_partial_and_exhaustion(engine):
"""高水位接近上限时只发剩余 ID;耗尽哨兵不可作为 ID 发出。"""
with Session(engine, future=True) as session:
row = session.get(IdSegment, GENERATOR_KEY)
row.next_id = ID_HIGH_EXCLUSIVE - 3
session.commit()
gen = SegmentIDGenerator(lambda: Session(engine, future=True))
ids = [gen.next_id() for _ in range(3)]
assert ids == [ID_HIGH_EXCLUSIVE - 3, ID_HIGH_EXCLUSIVE - 2, ID_HIGH_EXCLUSIVE - 1]
with pytest.raises(AppError) as exc:
gen.next_id()
assert exc.value.code == "SERVICE_UNAVAILABLE"
def test_exhausted_watermark_rejects(engine):
with Session(engine, future=True) as session:
row = session.get(IdSegment, GENERATOR_KEY)
row.next_id = ID_HIGH_EXCLUSIVE
session.commit()
gen = SegmentIDGenerator(lambda: Session(engine, future=True))
with pytest.raises(AppError) as exc:
gen.next_id()
assert exc.value.code == "SERVICE_UNAVAILABLE"
def test_allocation_failure_discards_and_raises(engine):
"""提交失败废弃候选段:发号器不得返回未确认 ID。"""
attempts = {"n": 0}
def flaky_factory():
attempts["n"] += 1
raise RuntimeError("db down")
gen = SegmentIDGenerator(flaky_factory)
with pytest.raises(AppError) as exc:
gen.next_id()
assert exc.value.code == "SERVICE_UNAVAILABLE"
assert attempts["n"] == 3
assert gen._cursor is None
def test_retry_after_transient_failure_uses_committed_state(engine):
"""一段连续失败转 SERVICE_UNAVAILABLE;恢复后按数据库已提交高水位重新申请。"""
calls = {"n": 0}
def flaky_factory():
calls["n"] += 1
if calls["n"] <= 3:
raise RuntimeError("transient")
return Session(engine, future=True)
gen = SegmentIDGenerator(flaky_factory)
with pytest.raises(AppError) as exc:
gen.next_id()
assert exc.value.code == "SERVICE_UNAVAILABLE"
assert calls["n"] == 3
value = gen.next_id()
assert value == ID_LOW
def test_fork_discards_inherited_segment(engine):
"""fork 后子进程丢弃继承候选段,按数据库高水位重新申请(design.md §9.6)。"""
gen = SegmentIDGenerator(lambda: Session(engine, future=True))
parent_id = gen.next_id()
watermark = _high_watermark(engine)
gen._pid = gen._pid + 1 # 模拟 fork 后子进程 pid 变化
child_id = gen.next_id()
assert child_id >= watermark
assert child_id != parent_id
def test_ids_unique_across_generators(engine):
gen1 = SegmentIDGenerator(lambda: Session(engine, future=True))
gen2 = SegmentIDGenerator(lambda: Session(engine, future=True))
ids = {gen1.next_id() for _ in range(50)} | {gen2.next_id() for _ in range(50)}
assert len(ids) == 100
"""归属解析单元测试(design.md §6.3):逐级降级、永不抛错、不猜测商户。"""
from datetime import date
from app.resolution import (
EMPLOYEE_CROSS_MERCHANT,
EMPLOYEE_INACTIVE,
EMPLOYEE_NOT_FOUND,
EMPLOYEE_NO_STORE,
STORE_CROSS_MERCHANT,
STORE_NOT_FOUND,
STORE_NOT_OPEN,
EmployeeSnapshot,
StoreSnapshot,
resolve_identity,
)
def _employee(**overrides):
values = dict(id=31, merchant_id=1001, status=1, position_status=1,
dimission_date=None, store_id=21, delete_flag=False)
values.update(overrides)
return EmployeeSnapshot(**values)
def _store(**overrides):
values = dict(id=21, merchant_id=1001, status=1)
values.update(overrides)
return StoreSnapshot(**values)
def test_employee_overrides_request_store():
resolved = resolve_identity(1001, 99, 31, "emp01",
lambda _id: _employee(), lambda _id: _store())
assert resolved.merchant_id == 1001
assert resolved.store_id == 21 # 可信员工的门店覆盖请求候选门店
assert resolved.employee_id == 31
assert resolved.employee_account == "emp01"
assert resolved.degradations == ()
def test_employee_without_store_keeps_request_store_and_warns():
resolved = resolve_identity(1001, 99, 32, None,
lambda _id: _employee(id=32, store_id=None), lambda sid: _store(id=sid))
# 总部员工无门店是正常场景:保留请求候选门店 + 记 EMPLOYEE_NO_STORE 降级
assert resolved.store_id == 99
assert resolved.employee_id == 32
assert resolved.degradations == (EMPLOYEE_NO_STORE,)
def test_employee_missing_degrades_and_drops_employee():
resolved = resolve_identity(1001, 99, 404, "x", lambda _id: None, lambda _id: _store())
assert resolved.employee_id is None
assert resolved.store_id == 99
assert resolved.employee_account == "x"
assert resolved.degradations == (EMPLOYEE_NOT_FOUND,)
def test_employee_deleted_treated_as_not_found():
resolved = resolve_identity(1001, None, 31, None,
lambda _id: _employee(delete_flag=True), lambda _id: None)
assert resolved.employee_id is None
assert resolved.degradations == (EMPLOYEE_NOT_FOUND,)
def test_employee_cross_merchant_dropped():
resolved = resolve_identity(1001, None, 33, "x",
lambda _id: _employee(id=33, merchant_id=1002, store_id=23),
lambda _id: None)
assert resolved.employee_id is None
assert resolved.store_id is None
assert resolved.degradations == (EMPLOYEE_CROSS_MERCHANT,)
def test_employee_inactive_variants():
for employee in (
_employee(status=2),
_employee(position_status=2),
_employee(dimission_date=date(2026, 1, 1)),
):
resolved = resolve_identity(1001, None, 31, None, lambda _id: employee, lambda _id: None)
assert resolved.degradations == (EMPLOYEE_INACTIVE,)
assert resolved.employee_id is None
def test_store_cross_merchant_degrades():
resolved = resolve_identity(1001, 23, None, None, lambda _id: None,
lambda _id: _store(id=23, merchant_id=1002))
assert resolved.store_id is None
assert resolved.degradations == (STORE_CROSS_MERCHANT,)
def test_store_not_found_degrades():
resolved = resolve_identity(1001, 404, None, None, lambda _id: None, lambda _id: None)
assert resolved.store_id is None
assert resolved.degradations == (STORE_NOT_FOUND,)
def test_store_not_open_degrades():
resolved = resolve_identity(1001, 22, None, None, lambda _id: None,
lambda _id: _store(id=22, status=2))
assert resolved.store_id is None
assert resolved.degradations == (STORE_NOT_OPEN,)
def test_employee_no_store_then_request_store_not_open():
resolved = resolve_identity(1001, 22, 32, None,
lambda _id: _employee(id=32, store_id=None),
lambda _id: _store(id=22, status=2))
assert resolved.store_id is None
assert resolved.employee_id == 32
assert resolved.degradations == (EMPLOYEE_NO_STORE, STORE_NOT_OPEN)
def test_empty_input_resolves_master_account_only():
resolved = resolve_identity(1001, None, None, None, lambda _id: None, lambda _id: None)
assert resolved.merchant_id == 1001
assert resolved.store_id is None
assert resolved.employee_id is None
assert resolved.employee_account is None
assert resolved.degradations == ()
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment