feat(db): 数据库层 MySQL 适配(方言钩子/同步驱动映射/种子方言分支)
- 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 说明
This commit is contained in:
@@ -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=
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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="")
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
+5
-5
@@ -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]:
|
||||
|
||||
@@ -29,6 +29,7 @@ dependencies = [
|
||||
"transformers>=5.15.1",
|
||||
"alembic>=1.19.1",
|
||||
"aiofiles>=24.1.0",
|
||||
"pymysql>=1.2.0",
|
||||
]
|
||||
|
||||
[dependency-groups]
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user