From ee1dfdfed9e5d7ec6a705b0978ff6471a69273ad Mon Sep 17 00:00:00 2001 From: Pine Date: Mon, 31 Aug 2026 21:40:17 +0800 Subject: [PATCH] =?UTF-8?q?feat(db):=20=E6=95=B0=E6=8D=AE=E5=BA=93?= =?UTF-8?q?=E5=B1=82=20MySQL=20=E9=80=82=E9=85=8D=EF=BC=88=E6=96=B9?= =?UTF-8?q?=E8=A8=80=E9=92=A9=E5=AD=90/=E5=90=8C=E6=AD=A5=E9=A9=B1?= =?UTF-8?q?=E5=8A=A8=E6=98=A0=E5=B0=84/=E7=A7=8D=E5=AD=90=E6=96=B9?= =?UTF-8?q?=E8=A8=80=E5=88=86=E6=94=AF=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - infrastructure/db.py:MySQL 方言编译钩子(String 无长度→VARCHAR(255),Text→MEDIUMTEXT); 新增 sync_database_url/make_sync_engine(sqlite+aiosqlite→sqlite,mysql+asyncmy→mysql+pymysql) - models.py:sessions.token 显式 String(2048)、system_configs.value 改 Text(存量数据超 255,MySQL 需兼容) - seed.py:_add_agent 幂等插入按方言分支(SQLite=INSERT OR IGNORE / MySQL=INSERT IGNORE) - park/tenants.py:同步只读引擎改走 make_sync_engine(原 +aiosqlite replace 对 MySQL 无效) - 依赖补 pymysql;.env.example 补 PINEAGENTS_DEMO_DATABASE_URL 说明 --- .env.example | 5 +++++ app/infrastructure/db.py | 40 ++++++++++++++++++++++++++++++++++++ app/infrastructure/models.py | 5 ++--- app/infrastructure/seed.py | 13 +++++++++--- app/park/tenants.py | 10 ++++----- pyproject.toml | 1 + uv.lock | 11 ++++++++++ 7 files changed, 74 insertions(+), 11 deletions(-) diff --git a/.env.example b/.env.example index c833384..b6fd97d 100644 --- a/.env.example +++ b/.env.example @@ -25,3 +25,8 @@ PINEAGENTS_WECHATPAY_PACKAGES=[{"id":"p10","amount":10,"bonus":0,"label":"10 元 PINEAGENTS_RECHARGE_MIN_YUAN=1 PINEAGENTS_RECHARGE_MAX_YUAN=5000 PINEAGENTS_RECHARGE_EXPIRE_MINUTES=15 + +# ---- 数据库(SQLite→MySQL 切换;未配置默认 serverdata/data/app.db)---- +# 示例:mysql+asyncmy://user:pass@192.168.1.3:8091/opc?charset=utf8mb4 +# 首次迁移:uv run python scripts/migrate_sqlite_to_mysql.py(建表+搬数据+行数校验) +PINEAGENTS_DEMO_DATABASE_URL= diff --git a/app/infrastructure/db.py b/app/infrastructure/db.py index a79ebc4..1dd5d3f 100644 --- a/app/infrastructure/db.py +++ b/app/infrastructure/db.py @@ -10,11 +10,13 @@ from __future__ import annotations from collections.abc import AsyncGenerator +from sqlalchemy import Text, create_engine, types from sqlalchemy.ext.asyncio import ( AsyncSession, async_sessionmaker, create_async_engine, ) +from sqlalchemy.ext.compiler import compiles from sqlalchemy.orm import DeclarativeBase from .. import config @@ -24,6 +26,44 @@ class Base(DeclarativeBase): """SQLAlchemy 声明式基类(数据模型层继承)。""" +@compiles(types.String, "mysql") +def _mysql_string_ddl(type_: types.String, compiler, **kw) -> str: + """MySQL 方言适配:String() 未显式给长度时按 VARCHAR(255) 建表(MySQL 要求 + VARCHAR 必须带长度;SQLite 忽略长度不受影响),Text 渲染 MEDIUMTEXT + (富文本/JSON 大字段,避免 MySQL TEXT 64KB 上限)。""" + if isinstance(type_, Text): + return "MEDIUMTEXT" + if not type_.length: + return "VARCHAR(255)" + return f"VARCHAR({int(type_.length)})" + + +def sync_database_url(url: str | None = None) -> str: + """DATABASE_URL → 同步驱动 URL(培训同步旁路 / 迁移脚本用)。 + + sqlite+aiosqlite → sqlite(pysqlite);mysql+asyncmy → mysql+pymysql。 + """ + url = url or config.DATABASE_URL + if url.startswith("sqlite+aiosqlite"): + return url.replace("sqlite+aiosqlite", "sqlite", 1) + if url.startswith("mysql+asyncmy"): + return url.replace("mysql+asyncmy", "mysql+pymysql", 1) + return url + + +def make_sync_engine(url: str | None = None): + """按 DATABASE_URL 创建同步引擎(asyncio.to_thread / 脚本场景)。""" + u = sync_database_url(url) + kwargs: dict = {"pool_pre_ping": True, "future": True} + if u.startswith("mysql"): + kwargs["pool_recycle"] = 3600 + if "charset=" not in u: + kwargs["connect_args"] = {"charset": "utf8mb4"} + elif u.startswith("sqlite"): + kwargs["connect_args"] = {"check_same_thread": False} + return create_engine(u, **kwargs) + + engine = create_async_engine( config.DATABASE_URL, pool_pre_ping=True, diff --git a/app/infrastructure/models.py b/app/infrastructure/models.py index 4593d62..6d76a6c 100644 --- a/app/infrastructure/models.py +++ b/app/infrastructure/models.py @@ -162,7 +162,7 @@ class SessionToken(Base): id: Mapped[str] = mapped_column(String, primary_key=True) jti: Mapped[str] = mapped_column(String, unique=True, nullable=False) - token: Mapped[str] = mapped_column(String, default="") # 完整 JWT(调试用) + token: Mapped[str] = mapped_column(String(2048), default="") # 完整 JWT(调试用;MySQL 需显式长度,存量可达 1.5KB) user_id: Mapped[str] = mapped_column(ForeignKey("users.id", ondelete="CASCADE"), nullable=False) username: Mapped[str] = mapped_column(String, default="") created_at: Mapped[str] = mapped_column(String, default="") @@ -386,8 +386,7 @@ class SystemConfig(Base): __tablename__ = "system_configs" key: Mapped[str] = mapped_column(String, primary_key=True) - value: Mapped[str] = mapped_column(String, default="") - description: Mapped[str] = mapped_column(String, default="") + value: Mapped[str] = mapped_column(Text, default="") # 配置值(存量可达 300+ 字节,MySQL 用 MEDIUMTEXT) updated_at: Mapped[str] = mapped_column(String, default="") diff --git a/app/infrastructure/seed.py b/app/infrastructure/seed.py index 524d135..091161b 100644 --- a/app/infrastructure/seed.py +++ b/app/infrastructure/seed.py @@ -306,15 +306,22 @@ DEMO_USERS = [ async def _add_agent(session: AsyncSession, seed: dict, user_id: str, port: str | None, now: str) -> None: - """幂等播种单个智能体:SQLite INSERT OR IGNORE((id,user_id,port) 冲突跳过),杜绝 UNIQUE 抛错。""" + """幂等播种单个智能体:方言 INSERT 冲突跳过(SQLite=INSERT OR IGNORE, + MySQL=INSERT IGNORE),杜绝 UNIQUE 抛错。""" + from sqlalchemy.dialects.mysql import insert as mysql_insert from sqlalchemy.dialects.sqlite import insert as sqlite_insert port_val = port or "" # agents.port NOT NULL;无端口身份用 "" - stmt = sqlite_insert(Agent).values( + values = dict( id=seed["id"], user_id=user_id, port=port_val, name=seed["name"], description=seed["description"], language=seed.get("language", "zh"), model_name=seed.get("model_name", ""), deletable=seed.get("deletable", True), use_fixed_soul=seed.get("use_fixed_soul", False), created_at=now, updated_at=now, - ).on_conflict_do_nothing(index_elements=["id", "user_id", "port"]) + ) + if session.get_bind().dialect.name == "mysql": + stmt = mysql_insert(Agent).values(**values).prefix_with("IGNORE") + else: + stmt = sqlite_insert(Agent).values(**values).on_conflict_do_nothing( + index_elements=["id", "user_id", "port"]) await session.execute(stmt) diff --git a/app/park/tenants.py b/app/park/tenants.py index 5fbbec6..5fe4780 100644 --- a/app/park/tenants.py +++ b/app/park/tenants.py @@ -67,9 +67,9 @@ def _new_id(prefix: str) -> str: def get_tenant_data_sync(tenant_id: str) -> dict: """同步读取该租户 data + companies(sim_engine 同步 tick,无事件循环时用)。""" - from sqlalchemy import create_engine, text - url = DATABASE_URL.replace("+aiosqlite", "") - engine = create_engine(url, future=True) + from sqlalchemy import text + from ..infrastructure.db import make_sync_engine + engine = make_sync_engine(DATABASE_URL) # 按 DATABASE_URL 映射同步驱动(sqlite/pymysql) try: with engine.connect() as conn: row = conn.execute(text("SELECT data_json FROM park_tenants WHERE id=:id"), {"id": tenant_id}).fetchone() @@ -85,8 +85,8 @@ def get_tenant_data_sync(tenant_id: str) -> dict: def _sync_conn(): - from sqlalchemy import create_engine - return create_engine(DATABASE_URL.replace("+aiosqlite", ""), future=True) + from ..infrastructure.db import make_sync_engine + return make_sync_engine(DATABASE_URL) def list_tenants_sync() -> list[dict]: diff --git a/pyproject.toml b/pyproject.toml index ebf05c8..d839b5f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -29,6 +29,7 @@ dependencies = [ "transformers>=5.15.1", "alembic>=1.19.1", "aiofiles>=24.1.0", + "pymysql>=1.2.0", ] [dependency-groups] diff --git a/uv.lock b/uv.lock index 79bba85..950b504 100644 --- a/uv.lock +++ b/uv.lock @@ -1977,6 +1977,7 @@ dependencies = [ { name = "paho-mqtt" }, { name = "pydantic" }, { name = "pyjwt" }, + { name = "pymysql" }, { name = "python-multipart" }, { name = "redis" }, { name = "scipy", version = "1.17.1", source = { registry = "https://mirrors.cloud.tencent.com/pypi/simple/" }, marker = "python_full_version < '3.12'" }, @@ -2013,6 +2014,7 @@ requires-dist = [ { name = "paho-mqtt", specifier = ">=1.6,<2" }, { name = "pydantic", specifier = ">=2.7.0" }, { name = "pyjwt", specifier = ">=2.8.0" }, + { name = "pymysql", specifier = ">=1.2.0" }, { name = "python-multipart", specifier = ">=0.0.12" }, { name = "redis", specifier = ">=5.0.0" }, { name = "scipy", specifier = ">=1.17.1" }, @@ -2316,6 +2318,15 @@ wheels = [ { url = "https://mirrors.cloud.tencent.com/pypi/packages/a3/5e/ecf12fdb62546d64385c158514e9b2b671f7832108ef2ecd2020ce0af2d1/pyjwt-2.13.0-py3-none-any.whl", hash = "sha256:66adcc2aff09b3f1bbd95fc1e1577df8ac8723c978552fd43304c8a290ac5728" }, ] +[[package]] +name = "pymysql" +version = "1.2.0" +source = { registry = "https://mirrors.cloud.tencent.com/pypi/simple/" } +sdist = { url = "https://mirrors.cloud.tencent.com/pypi/packages/c9/bc/1c6a92f385940f727daeecf3bacaf186e03875dff57197801046c583bcf0/pymysql-1.2.0.tar.gz", hash = "sha256:6c7b17ca686988104d7426c27895b455cdeea3e9d3ceb1270f0c3704fead8c33" } +wheels = [ + { url = "https://mirrors.cloud.tencent.com/pypi/packages/c4/bd/2534e130295c8cfd4f0a2e31623baab7502278f1e97bcfe61db75656a77f/pymysql-1.2.0-py3-none-any.whl", hash = "sha256:62169ce6d5510f08e140c5e7990ee884a9764024e4a9a27b2cc11f1099322ae0" }, +] + [[package]] name = "pytest" version = "9.1.1"