168 lines
5.5 KiB
Python
168 lines
5.5 KiB
Python
"""
|
||
角色卡格式转换器
|
||
支持 SillyTavern V2/V3 格式与内部格式的双向转换
|
||
"""
|
||
import json
|
||
import base64
|
||
from typing import Optional, Dict, Any
|
||
from pathlib import Path
|
||
from PIL import Image
|
||
import io
|
||
|
||
try:
|
||
from backend.models.internal import CharacterCard
|
||
except ImportError:
|
||
from models.internal import CharacterCard
|
||
|
||
|
||
class CharacterCardConverter:
|
||
"""角色卡格式转换器"""
|
||
|
||
@staticmethod
|
||
def st_to_internal(st_data: dict, avatar_path: Optional[str] = None) -> CharacterCard:
|
||
"""
|
||
SillyTavern 格式 → 内部格式
|
||
|
||
Args:
|
||
st_data: SillyTavern 角色卡数据(V2/V3)
|
||
avatar_path: 头像路径(可选)
|
||
|
||
Returns:
|
||
CharacterCard 对象
|
||
"""
|
||
import uuid
|
||
from datetime import datetime
|
||
|
||
# 兼容两种传入方式:完整ST格式或直接data
|
||
if 'spec' in st_data:
|
||
data = st_data.get('data', {})
|
||
else:
|
||
data = st_data
|
||
|
||
extensions = data.get('extensions', {})
|
||
|
||
return CharacterCard(
|
||
id=str(uuid.uuid4()),
|
||
name=data['name'],
|
||
description=data.get('description', ''),
|
||
personality=data.get('personality', ''),
|
||
scenario=data.get('scenario', ''),
|
||
first_mes=data.get('first_mes', ''),
|
||
mes_example=data.get('mes_example', ''),
|
||
categories=[], # ST没有categories
|
||
tableHeaders=[], # ST没有tableHeaders
|
||
worldInfoId=extensions.get('world'),
|
||
outputSchema=None, # ST不支持结构化输出
|
||
avatarPath=avatar_path,
|
||
alternate_greetings=data.get('alternate_greetings', []),
|
||
tags=data.get('tags', []),
|
||
createdAt=int(datetime.now().timestamp()),
|
||
updatedAt=int(datetime.now().timestamp()),
|
||
lastChatAt=None,
|
||
isFavorite=extensions.get('fav', False),
|
||
version=1
|
||
)
|
||
|
||
@staticmethod
|
||
def internal_to_st(character: CharacterCard) -> dict:
|
||
"""
|
||
内部格式 → SillyTavern V3 格式
|
||
|
||
Args:
|
||
character: CharacterCard 对象
|
||
|
||
Returns:
|
||
SillyTavern V3 格式字典
|
||
"""
|
||
return {
|
||
"spec": "chara_card_v3",
|
||
"spec_version": "3.0",
|
||
"data": {
|
||
"name": character.name,
|
||
"description": character.description,
|
||
"personality": character.personality,
|
||
"scenario": character.scenario,
|
||
"first_mes": character.first_mes,
|
||
"mes_example": character.mes_example,
|
||
"alternate_greetings": character.alternate_greetings or [],
|
||
"tags": character.tags or [],
|
||
"creator_notes": "",
|
||
"system_prompt": "",
|
||
"post_history_instructions": "",
|
||
"extensions": {
|
||
"world": character.worldInfoId,
|
||
"talkativeness": 0.5,
|
||
"fav": character.isFavorite
|
||
}
|
||
}
|
||
}
|
||
|
||
@staticmethod
|
||
def export_as_png(character: CharacterCard, avatar_path: Optional[str] = None, use_default_avatar: bool = False) -> bytes:
|
||
"""
|
||
导出为 SillyTavern PNG 格式
|
||
|
||
Args:
|
||
character: CharacterCard 对象
|
||
avatar_path: 头像图片路径(可选)
|
||
use_default_avatar: 是否使用默认头像(不嵌入JSON数据)
|
||
|
||
Returns:
|
||
PNG 文件的二进制数据
|
||
"""
|
||
# 1. 创建/加载图片
|
||
if avatar_path and Path(avatar_path).exists():
|
||
img = Image.open(avatar_path)
|
||
else:
|
||
# 创建默认图片(400x600像素,灰色背景)
|
||
img = Image.new('RGB', (400, 600), color=(73, 109, 137))
|
||
|
||
# 确保是 RGBA 模式
|
||
if img.mode != 'RGBA':
|
||
img = img.convert('RGBA')
|
||
|
||
# 2. 如果不是默认头像,才嵌入 JSON 数据
|
||
if not use_default_avatar:
|
||
st_data = CharacterCardConverter.internal_to_st(character)
|
||
json_str = json.dumps(st_data, ensure_ascii=False)
|
||
base64_data = base64.b64encode(json_str.encode('utf-8')).decode('ascii')
|
||
img.text['ccv3'] = base64_data
|
||
|
||
# 3. 保存到字节流
|
||
buffer = io.BytesIO()
|
||
img.save(buffer, format='PNG')
|
||
buffer.seek(0)
|
||
|
||
return buffer.read()
|
||
|
||
@staticmethod
|
||
def extract_from_png(png_data: bytes) -> Optional[dict]:
|
||
"""
|
||
从 PNG 文件中提取嵌入的角色数据
|
||
|
||
Args:
|
||
png_data: PNG 文件的二进制数据
|
||
|
||
Returns:
|
||
SillyTavern 格式字典,如果没有嵌入数据则返回 None
|
||
"""
|
||
try:
|
||
img = Image.open(io.BytesIO(png_data))
|
||
|
||
# 尝试 V3 格式 (ccv3)
|
||
if 'ccv3' in img.text:
|
||
json_str = base64.b64decode(img.text['ccv3']).decode('utf-8')
|
||
return json.loads(json_str)
|
||
|
||
# 尝试 V2 格式 (chara)
|
||
elif 'chara' in img.text:
|
||
json_str = base64.b64decode(img.text['chara']).decode('utf-8')
|
||
return json.loads(json_str)
|
||
|
||
# 没有嵌入数据
|
||
return None
|
||
|
||
except Exception as e:
|
||
print(f"解析PNG失败: {e}")
|
||
return None
|