110 lines
3.6 KiB
Python
110 lines
3.6 KiB
Python
import functools
|
|||
|
|
import json
|
||
|
|
from functools import cached_property
|
||
|
|
from typing import List, Optional, TYPE_CHECKING
|
||
|
|
|
||
|
|
try:
|
||
|
|
import ujson as json
|
||
|
|
except ImportError:
|
||
|
|
import json
|
||
|
|
try:
|
||
|
|
import jieba_fast as jieba
|
||
|
|
import jieba_fast.analyse as jieba_analyse
|
||
|
|
except ImportError:
|
||
|
|
import jieba
|
||
|
|
import jieba.analyse as jieba_analyse
|
||
|
|
|
||
|
|
from nonebot import require
|
||
|
|
require("nonebot_plugin_orm")
|
||
|
|
|
||
|
|
from nonebot_plugin_orm import Model
|
||
|
|
from sqlalchemy import BigInteger, Boolean, Integer, Text, ForeignKey, Index
|
||
|
|
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||
|
|
from sqlalchemy.types import JSON
|
||
|
|
|
||
|
|
from .config import config_manager
|
||
|
|
|
||
|
|
config = config_manager.config
|
||
|
|
|
||
|
|
JSON_DUMPS = functools.partial(json.dumps, ensure_ascii=False)
|
||
|
|
jieba.setLogLevel(jieba.logging.INFO)
|
||
|
|
jieba.load_userdict(config.dictionary)
|
||
|
|
|
||
|
|
|
||
|
|
class ChatMessage(Model):
|
||
|
|
__tablename__ = "learning_chat_message"
|
||
|
|
__table_args__ = (
|
||
|
|
Index("ix_message_group_time", "group_id", "time"),
|
||
|
|
)
|
||
|
|
|
||
|
|
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||
|
|
group_id: Mapped[int] = mapped_column(BigInteger)
|
||
|
|
user_id: Mapped[int] = mapped_column(BigInteger)
|
||
|
|
message_id: Mapped[int] = mapped_column(BigInteger)
|
||
|
|
message: Mapped[str] = mapped_column(Text)
|
||
|
|
raw_message: Mapped[str] = mapped_column(Text)
|
||
|
|
plain_text: Mapped[str] = mapped_column(Text)
|
||
|
|
time: Mapped[int] = mapped_column(Integer)
|
||
|
|
|
||
|
|
@cached_property
|
||
|
|
def is_plain_text(self) -> bool:
|
||
|
|
return "[CQ:" not in self.message
|
||
|
|
|
||
|
|
@cached_property
|
||
|
|
def keyword_list(self) -> List[str]:
|
||
|
|
if not self.is_plain_text and not len(self.plain_text):
|
||
|
|
return []
|
||
|
|
return jieba_analyse.extract_tags(self.plain_text, topK=config.KEYWORDS_SIZE)
|
||
|
|
|
||
|
|
@cached_property
|
||
|
|
def keywords(self) -> str:
|
||
|
|
if not self.is_plain_text and not len(self.plain_text):
|
||
|
|
return self.message
|
||
|
|
return (
|
||
|
|
self.message if len(self.keyword_list) < 2 else " ".join(self.keyword_list)
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
class ChatContext(Model):
|
||
|
|
__tablename__ = "learning_chat_context"
|
||
|
|
__table_args__ = (
|
||
|
|
Index("ix_context_keywords_time", "keywords", "time"),
|
||
|
|
)
|
||
|
|
|
||
|
|
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||
|
|
keywords: Mapped[str] = mapped_column(Text)
|
||
|
|
time: Mapped[int] = mapped_column(Integer)
|
||
|
|
count: Mapped[int] = mapped_column(Integer, default=1)
|
||
|
|
|
||
|
|
answers: Mapped[List["ChatAnswer"]] = relationship(back_populates="context")
|
||
|
|
|
||
|
|
|
||
|
|
class ChatAnswer(Model):
|
||
|
|
__tablename__ = "learning_chat_answer"
|
||
|
|
__table_args__ = (
|
||
|
|
Index("ix_answer_keywords_time", "keywords", "time"),
|
||
|
|
)
|
||
|
|
|
||
|
|
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||
|
|
keywords: Mapped[str] = mapped_column(Text)
|
||
|
|
group_id: Mapped[int] = mapped_column(BigInteger)
|
||
|
|
count: Mapped[int] = mapped_column(Integer, default=1)
|
||
|
|
time: Mapped[int] = mapped_column(Integer)
|
||
|
|
messages: Mapped[list] = mapped_column(JSON, default=list)
|
||
|
|
|
||
|
|
context_id: Mapped[Optional[int]] = mapped_column(
|
||
|
|
Integer, ForeignKey("learning_chat_context.id"), nullable=True
|
||
|
|
)
|
||
|
|
context: Mapped[Optional[ChatContext]] = relationship(back_populates="answers")
|
||
|
|
|
||
|
|
|
||
|
|
class ChatBlackList(Model):
|
||
|
|
__tablename__ = "learning_chat_blacklist"
|
||
|
|
__table_args__ = (
|
||
|
|
Index("ix_blacklist_keywords", "keywords"),
|
||
|
|
)
|
||
|
|
|
||
|
|
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||
|
|
keywords: Mapped[str] = mapped_column(Text)
|
||
|
|
global_ban: Mapped[bool] = mapped_column(Boolean, default=False)
|
||
|
|
ban_group_id: Mapped[list] = mapped_column(JSON, default=list)
|