400 lines
14 KiB
Python
400 lines
14 KiB
Python
"""数字分身管理服务层 — 同步连接数字分身应用的 SQLite 数据库"""
|
|
import json
|
|
import os
|
|
from datetime import datetime, timedelta
|
|
from typing import Optional, Tuple
|
|
|
|
from sqlalchemy import create_engine, select, text
|
|
from sqlalchemy.orm import sessionmaker, Session
|
|
|
|
from app.core.config import settings
|
|
from app.core.logger import logger
|
|
from app.models import UserPersonality, VirtualUser
|
|
|
|
|
|
_engine = None
|
|
_SessionLocal: Optional[sessionmaker] = None
|
|
|
|
AVATAR_ACCOUNT_PREFIX = "__avatar__:"
|
|
SQUARE_INTERACTION_PERMISSION = "interact"
|
|
SQUARE_INTERACTION_ACTIONS = frozenset({"like", "collect", "comment", "reply"})
|
|
|
|
|
|
def _get_engine_and_session():
|
|
global _engine, _SessionLocal
|
|
if _engine is None:
|
|
db_path = settings.AVATAR_DB_PATH
|
|
if not db_path:
|
|
# 默认路径:从 backend/app/core/ 向上三级
|
|
base = os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))))
|
|
db_path = os.path.join(base, "digital-avatar-app", "backend", "avatar.db")
|
|
if not os.path.isabs(db_path):
|
|
db_path = os.path.abspath(db_path)
|
|
if not os.path.exists(db_path):
|
|
# 数据库不存在时返回 None,由调用方处理
|
|
return None, None
|
|
_engine = create_engine(
|
|
f"sqlite:///{db_path}",
|
|
connect_args={"check_same_thread": False},
|
|
)
|
|
_SessionLocal = sessionmaker(bind=_engine, autoflush=False, expire_on_commit=False)
|
|
return _engine, _SessionLocal()
|
|
|
|
|
|
def get_session() -> Optional[Session]:
|
|
_, session = _get_engine_and_session()
|
|
return session
|
|
|
|
|
|
def is_available() -> bool:
|
|
"""检查数字分身数据库是否可用"""
|
|
engine, _ = _get_engine_and_session()
|
|
return engine is not None
|
|
|
|
|
|
def _resolve_photo_url(photo_url: str) -> str:
|
|
"""将相对路径的头像 URL 补全为绝对路径"""
|
|
if not photo_url:
|
|
return ""
|
|
if photo_url.startswith(("http://", "https://")):
|
|
return photo_url
|
|
base = settings.AVATAR_BACKEND_URL
|
|
if base:
|
|
base = base.rstrip("/")
|
|
return f"{base}{photo_url}"
|
|
return photo_url
|
|
|
|
|
|
def _get_global_token_balance(db: Session) -> int:
|
|
"""获取全局 token_account 余额(单行表)"""
|
|
try:
|
|
result = db.execute(text("SELECT balance FROM token_account LIMIT 1")).fetchone()
|
|
return result.balance if result and result.balance else 0
|
|
except Exception:
|
|
return 0
|
|
|
|
|
|
def _decode_config(value) -> dict:
|
|
if isinstance(value, dict):
|
|
return value
|
|
if isinstance(value, str):
|
|
try:
|
|
decoded = json.loads(value)
|
|
return decoded if isinstance(decoded, dict) else {}
|
|
except (json.JSONDecodeError, ValueError):
|
|
return {}
|
|
return {}
|
|
|
|
|
|
def is_delegated_avatar_user(user: VirtualUser | None) -> bool:
|
|
return bool(user and (user.account or "").startswith(AVATAR_ACCOUNT_PREFIX))
|
|
|
|
|
|
def delegated_avatar_id(user: VirtualUser | None) -> str:
|
|
if not is_delegated_avatar_user(user):
|
|
return ""
|
|
return (user.account or "")[len(AVATAR_ACCOUNT_PREFIX):]
|
|
|
|
|
|
def _list_square_interaction_authorizations(db: Session) -> list[dict]:
|
|
"""读取已明确授权分身参与广场互动的身份与会会令牌。"""
|
|
rows = db.execute(text("""
|
|
SELECT
|
|
a.id AS avatar_id,
|
|
a.name AS avatar_name,
|
|
a.display_name AS avatar_display_name,
|
|
a.description AS avatar_description,
|
|
a.photo_url AS avatar_photo_url,
|
|
a.config AS avatar_config,
|
|
u.huihui_user_id,
|
|
u.nickname AS owner_nickname,
|
|
u.avatar_url AS owner_avatar_url,
|
|
u.huihui_token
|
|
FROM avatars a
|
|
JOIN users u ON u.huihui_user_id = a.owner_id
|
|
WHERE a.status = 'active'
|
|
""")).fetchall()
|
|
|
|
authorized = []
|
|
for row in rows:
|
|
config = _decode_config(row.avatar_config)
|
|
permissions = config.get("authorizationPermissions", [])
|
|
if not isinstance(permissions, list) or SQUARE_INTERACTION_PERMISSION not in permissions:
|
|
continue
|
|
platform_uid = str(row.huihui_user_id or "").strip()
|
|
token = str(row.huihui_token or "").strip()
|
|
if not platform_uid or not token:
|
|
continue
|
|
authorized.append({
|
|
"avatar_id": str(row.avatar_id),
|
|
"avatar_name": row.avatar_display_name or row.avatar_name or row.owner_nickname or "数字分身",
|
|
"avatar_description": row.avatar_description or "",
|
|
"avatar_url": _resolve_photo_url(row.avatar_photo_url or row.owner_avatar_url or ""),
|
|
"config": config,
|
|
"platform_uid": platform_uid,
|
|
"token": token,
|
|
})
|
|
return authorized
|
|
|
|
|
|
def get_square_interaction_permissions(avatar_id: str) -> frozenset[str]:
|
|
"""实时复核授权;数据库不可用、令牌失效或撤权时一律拒绝执行。"""
|
|
avatar_db = get_session()
|
|
if avatar_db is None:
|
|
return frozenset()
|
|
try:
|
|
authorized_ids = {
|
|
item["avatar_id"] for item in _list_square_interaction_authorizations(avatar_db)
|
|
}
|
|
return SQUARE_INTERACTION_ACTIONS if avatar_id in authorized_ids else frozenset()
|
|
except Exception as exc:
|
|
logger.error(f"读取数字分身广场互动授权失败: {exc}")
|
|
return frozenset()
|
|
finally:
|
|
avatar_db.close()
|
|
|
|
|
|
def _word_count_range(config: dict) -> tuple[int, int]:
|
|
ranges = {
|
|
"short": (10, 35),
|
|
"medium": (20, 60),
|
|
"long": (30, 80),
|
|
}
|
|
return ranges.get(str(config.get("responseLength") or "medium"), (20, 60))
|
|
|
|
|
|
async def sync_square_interaction_users(db) -> set[str]:
|
|
"""把已授权分身同步为调度器身份,并刷新其会会会话。"""
|
|
avatar_db = get_session()
|
|
if avatar_db is None:
|
|
logger.warning("数字分身数据库不可用,跳过广场互动授权同步")
|
|
return set()
|
|
try:
|
|
authorized = _list_square_interaction_authorizations(avatar_db)
|
|
except Exception as exc:
|
|
logger.error(f"同步数字分身广场互动授权失败: {exc}")
|
|
return set()
|
|
finally:
|
|
avatar_db.close()
|
|
|
|
from app.core.redis_client import delete_session, set_session
|
|
|
|
result = await db.execute(
|
|
select(VirtualUser).where(VirtualUser.account.like(f"{AVATAR_ACCOUNT_PREFIX}%"))
|
|
)
|
|
existing_users = {delegated_avatar_id(user): user for user in result.scalars().all()}
|
|
authorized_ids = {item["avatar_id"] for item in authorized}
|
|
|
|
for avatar_id, user in existing_users.items():
|
|
if avatar_id not in authorized_ids:
|
|
user.is_enabled = 0
|
|
user.status = 0
|
|
user.session_token = None
|
|
user.session_expires_at = None
|
|
await delete_session(user.id)
|
|
|
|
for item in authorized:
|
|
avatar_id = item["avatar_id"]
|
|
user = existing_users.get(avatar_id)
|
|
if user is None:
|
|
user = VirtualUser(
|
|
nickname=item["avatar_name"],
|
|
account=f"{AVATAR_ACCOUNT_PREFIX}{avatar_id}",
|
|
password_enc="",
|
|
status=2,
|
|
is_enabled=1,
|
|
platform_uid=item["platform_uid"],
|
|
remark="用户授权的数字分身广场互动身份",
|
|
)
|
|
db.add(user)
|
|
await db.flush()
|
|
|
|
expires_at = datetime.now() + timedelta(days=1)
|
|
user.nickname = item["avatar_name"]
|
|
user.real_name = item["avatar_name"]
|
|
user.avatar_url = item["avatar_url"]
|
|
user.platform_uid = item["platform_uid"]
|
|
user.session_token = item["token"]
|
|
user.session_expires_at = expires_at
|
|
user.last_login_at = datetime.now()
|
|
user.status = 2
|
|
user.is_enabled = 1
|
|
|
|
config = item["config"]
|
|
personality_result = await db.execute(
|
|
select(UserPersonality).where(UserPersonality.user_id == user.id)
|
|
)
|
|
personality = personality_result.scalar_one_or_none()
|
|
word_min, word_max = _word_count_range(config)
|
|
prompt_parts = [item["avatar_description"], str(config.get("systemPrompt") or "")]
|
|
style_prompt = "\n".join(part.strip() for part in prompt_parts if part and part.strip())
|
|
if personality is None:
|
|
personality = UserPersonality(user_id=user.id)
|
|
db.add(personality)
|
|
personality.language_style = str(config.get("replyStyle") or "professional")
|
|
personality.personality_desc = item["avatar_description"]
|
|
personality.comment_style_prompt = style_prompt
|
|
personality.word_count_min = word_min
|
|
personality.word_count_max = word_max
|
|
|
|
await set_session(user.id, {
|
|
"token": item["token"],
|
|
"session_id": f"avatar:{avatar_id}",
|
|
"platform_uid": item["platform_uid"],
|
|
"org_id": "",
|
|
"login_time": datetime.now().isoformat(),
|
|
"nickname": item["avatar_name"],
|
|
"real_name": item["avatar_name"],
|
|
"avatar": item["avatar_url"],
|
|
"delegated_avatar_id": avatar_id,
|
|
}, expire=86400)
|
|
|
|
await db.commit()
|
|
return authorized_ids
|
|
|
|
|
|
class AvatarService:
|
|
|
|
@staticmethod
|
|
def list_avatars(
|
|
db: Session,
|
|
page: int = 1,
|
|
page_size: int = 20,
|
|
keyword: Optional[str] = None,
|
|
status: Optional[str] = None,
|
|
) -> Tuple[int, list]:
|
|
"""分页查询所有数字分身(含归属用户信息)"""
|
|
where_clauses = []
|
|
params = {}
|
|
|
|
if keyword:
|
|
where_clauses.append(
|
|
"(a.name LIKE :kw OR a.display_name LIKE :kw)"
|
|
)
|
|
params["kw"] = f"%{keyword}%"
|
|
|
|
if status:
|
|
where_clauses.append("a.status = :status")
|
|
params["status"] = status
|
|
|
|
where_sql = ""
|
|
if where_clauses:
|
|
where_sql = "WHERE " + " AND ".join(where_clauses)
|
|
|
|
# 计数
|
|
count_sql = f"SELECT COUNT(*) FROM avatars a {where_sql}"
|
|
total = db.execute(text(count_sql), params).scalar() or 0
|
|
|
|
# 分页查询
|
|
offset = (page - 1) * page_size
|
|
params["limit"] = page_size
|
|
params["offset"] = offset
|
|
|
|
query = text(f"""
|
|
SELECT
|
|
a.id, a.name, a.display_name, a.description,
|
|
a.photo_url, a.emoji, a.status, a.token_balance,
|
|
a.config, a.created_at, a.updated_at,
|
|
u.nickname AS owner_nickname,
|
|
u.phone AS owner_phone
|
|
FROM avatars a
|
|
LEFT JOIN users u ON a.owner_id = u.huihui_user_id
|
|
{where_sql}
|
|
ORDER BY a.created_at DESC
|
|
LIMIT :limit OFFSET :offset
|
|
""")
|
|
rows = db.execute(query, params).fetchall()
|
|
|
|
items = []
|
|
global_token = _get_global_token_balance(db)
|
|
for row in rows:
|
|
config = {}
|
|
if row.config:
|
|
if isinstance(row.config, str):
|
|
import json
|
|
try:
|
|
config = json.loads(row.config)
|
|
except (json.JSONDecodeError, ValueError):
|
|
config = {}
|
|
elif isinstance(row.config, dict):
|
|
config = row.config
|
|
items.append({
|
|
"id": row.id,
|
|
"name": row.name,
|
|
"display_name": row.display_name,
|
|
"description": row.description or "",
|
|
"photo_url": _resolve_photo_url(row.photo_url),
|
|
"emoji": row.emoji or "🤖",
|
|
"status": row.status or "active",
|
|
"token_balance": global_token or (row.token_balance or 0),
|
|
"config": config,
|
|
"owner_nickname": row.owner_nickname or "",
|
|
"owner_phone": row.owner_phone or "",
|
|
"created_at": str(row.created_at) if row.created_at else "",
|
|
"updated_at": str(row.updated_at) if row.updated_at else "",
|
|
})
|
|
|
|
return total, items
|
|
|
|
@staticmethod
|
|
def get_avatar(db: Session, avatar_id: str) -> Optional[dict]:
|
|
"""获取单个数字分身详情"""
|
|
query = text("""
|
|
SELECT
|
|
a.id, a.name, a.display_name, a.description,
|
|
a.photo_url, a.emoji, a.status, a.token_balance,
|
|
a.config, a.created_at, a.updated_at,
|
|
u.nickname AS owner_nickname,
|
|
u.phone AS owner_phone
|
|
FROM avatars a
|
|
LEFT JOIN users u ON a.owner_id = u.huihui_user_id
|
|
WHERE a.id = :avatar_id
|
|
""")
|
|
row = db.execute(query, {"avatar_id": avatar_id}).fetchone()
|
|
if not row:
|
|
return None
|
|
|
|
config = {}
|
|
if row.config:
|
|
if isinstance(row.config, str):
|
|
import json
|
|
try:
|
|
config = json.loads(row.config)
|
|
except (json.JSONDecodeError, ValueError):
|
|
config = {}
|
|
elif isinstance(row.config, dict):
|
|
config = row.config
|
|
|
|
global_token = _get_global_token_balance(db)
|
|
return {
|
|
"id": row.id,
|
|
"name": row.name,
|
|
"display_name": row.display_name,
|
|
"description": row.description or "",
|
|
"photo_url": _resolve_photo_url(row.photo_url),
|
|
"emoji": row.emoji or "🤖",
|
|
"status": row.status or "active",
|
|
"token_balance": global_token or (row.token_balance or 0),
|
|
"config": config,
|
|
"owner_nickname": row.owner_nickname or "",
|
|
"owner_phone": row.owner_phone or "",
|
|
"created_at": str(row.created_at) if row.created_at else "",
|
|
"updated_at": str(row.updated_at) if row.updated_at else "",
|
|
}
|
|
|
|
@staticmethod
|
|
def update_status(db: Session, avatar_id: str, status: str) -> Optional[dict]:
|
|
"""更新数字分身状态(开关机)"""
|
|
query = text("""
|
|
UPDATE avatars SET status = :status, updated_at = datetime('now')
|
|
WHERE id = :avatar_id
|
|
""")
|
|
result = db.execute(query, {"status": status, "avatar_id": avatar_id})
|
|
db.commit()
|
|
if result.rowcount == 0:
|
|
return None
|
|
return AvatarService.get_avatar(db, avatar_id)
|
|
|
|
|
|
avatar_service = AvatarService()
|