DMP游戏内AI聊天的想法

8471 字
42 分钟
DMP游戏内AI聊天的想法

事情大概是这样的

闲得蛋疼发呆时,突然想到,关于饥荒这款游戏,里面是没有任何新手引导的,就等着你死一次长一智,死一次长一智

新手项快速成长,无非就是三种方式,多死找大佬一起查资料

如果是查资料的话,就要切出去,查看多个网页,我觉的挺割裂的,就想在游戏中做一个智能助理,专门解决游戏问题

说做就做,捋一下思路:

  1. 实时获取玩家聊天记录,需要毫秒级
  2. 判断聊天记录类型,玩家闲聊还是向AI提问
  3. 获取每个玩家的问题并分别保存,让AI有记忆且不混乱
  4. 调用大语言模型拿到答案
  5. 将答案发送到公屏
  6. 设置聊天记录定时清理,防止内存无限增长

有了这个思路,就开撸,下面是这一版游戏中的表现

游戏内聊天
游戏内聊天

符合预期的:

  1. 不会回答饥荒以外的东西
  2. 能够正确识别玩家的意图
  3. 聊天响应时间较快(小于1秒)

不符合预期的:

  1. 回答并不全对,回答的内容基本是靠AI训练时的数据,且污染严重
  2. 有很多数据没有,也是因为AI训练时,没有相关的游戏资料

那么接下来就是改进方案了:

  1. 构建一个权威的知识库(RAG),AI只能从知识库中获取资料并回答
  2. 知识库为MD文档
  3. 构建关键词索引和向量数据
  4. 搜索机制采用关键词搜索和向量搜索
    1. 查看是否配置了embedding模型,有的话就向量搜索
    2. 没有embedding就回退至关键词搜索

哈哈哈哈,说了这么多,卡在第1步了,我没有游戏的资料

我自己搞了点测试数据,试了一下

就比如

:火炬如何制作

AI:火炬是游戏中可携带的光源,需要2个采下的草和2个树枝

网上爬是不可能的,没必要为了这个功能产生法务风险

TIP

期间也是发现了,某些人,用着开源的东西,却把自己有的东西,死死捏住不放,还阴阳怪气,跟狗一样

放弃了

下面是写好的代码,就留在这里,哪天有资料了,再捡起来

aichat.go
package aichat
import (
"context"
"dst-management-platform-api/database/dao"
"dst-management-platform-api/database/models"
"sync"
)
// ================================================================
// 导出类型
// ================================================================
// Manager AI 对话服务管理器,是 aichat 包的核心入口。
// 负责管理游戏内 AI 聊天的工作协程、Wiki 知识库搜索(关键词+向量)。
type Manager struct {
roomAISettingDao *dao.RoomAISettingDAO
client *aiClient
ctx context.Context
cancel context.CancelFunc
lifecycle sync.Mutex
mu sync.Mutex
workers map[int]*roomWorker
active bool
closed bool
keywordSearcher *keywordWikiSearcher
embedSearcher *embeddingWikiSearcher
embedSearcherMu sync.Mutex
lastEmbedConfig string
}
// EmbeddingConfig 向量嵌入模型配置
type EmbeddingConfig struct {
APIURL string
APIKey string
Model string
Dimensions int
}
// EmbeddingStats 向量索引统计信息
type EmbeddingStats struct {
TotalDocs int `json:"totalDocs"`
TotalCategories int `json:"totalCategories"`
VectorDim int `json:"vectorDim"`
Categories map[string]int `json:"categories"`
}
// NewManager 创建 AI 对话服务管理器
func NewManager(roomAISettingDao *dao.RoomAISettingDAO, pluginDao *dao.PluginDAO) *Manager {
return newManager(roomAISettingDao, pluginDao)
}
// Start 启动所有已启用房间的 AI 对话监听
func (m *Manager) Start() error {
return m.start()
}
// StopAll 停止所有房间的 AI 对话监听
func (m *Manager) StopAll() {
m.stopAll()
}
// StopRoom 停止指定房间的 AI 对话监听
func (m *Manager) StopRoom(roomID int) {
m.lifecycle.Lock()
defer m.lifecycle.Unlock()
m.stopRoom(roomID)
}
// Reload 重新加载指定房间的 AI 配置并重启监听
func (m *Manager) Reload(roomID int) error {
return m.reload(roomID)
}
// Close 关闭 AI 对话服务,释放所有资源
func (m *Manager) Close() {
m.closeManager()
}
// BuildEmbeddingIndex 手动构建向量索引。config 指定 Embedding API 配置,force 为 true 时强制全量重建。
// 构建过程可能耗时较长。
func (m *Manager) BuildEmbeddingIndex(config EmbeddingConfig, force bool) error {
return m.buildEmbeddingIndex(config, force)
}
// GetEmbeddingStats 获取向量索引统计信息
func (m *Manager) GetEmbeddingStats(config EmbeddingConfig) EmbeddingStats {
return m.getEmbeddingStats(config)
}
// BuildKeywordIndex 手动构建关键词搜索索引。force 为 true 时强制重建。
func (m *Manager) BuildKeywordIndex(force bool) error {
return m.buildKeywordIndex(force)
}
// UnloadKeywordIndex 释放关键词搜索索引占用的内存。下次 Search 时会自动从磁盘重新加载。
func (m *Manager) UnloadKeywordIndex() {
m.unloadKeywordIndex()
}
// ValidateModelConfig 校验 AI 模型配置参数
func ValidateModelConfig(config *models.AIModelConfig) error {
return validateModelConfig(config)
}
// ValidateRoomSetting 校验房间 AI 设置参数
func ValidateRoomSetting(setting *models.RoomAISetting) error {
return validateRoomSetting(setting)
}
client.go
package aichat
import (
"bytes"
"context"
"dst-management-platform-api/database/models"
"dst-management-platform-api/utils"
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"strings"
"time"
)
type chatMessage struct {
Role string `json:"role"`
Content string `json:"content"`
}
type chatCompletionRequest struct {
Model string `json:"model"`
Messages []chatMessage `json:"messages"`
Temperature float64 `json:"temperature"`
MaxTokens int `json:"max_tokens"`
Stream bool `json:"stream"`
}
type chatCompletionResponse struct {
Choices []struct {
Message chatMessage `json:"message"`
} `json:"choices"`
Error *struct {
Message string `json:"message"`
} `json:"error"`
}
type aiClient struct {
httpClient *http.Client
}
func newClient() *aiClient {
return &aiClient{
httpClient: &http.Client{
CheckRedirect: func(_ *http.Request, _ []*http.Request) error {
return http.ErrUseLastResponse
},
},
}
}
func (c *aiClient) complete(ctx context.Context, config models.AIModelConfig, messages []chatMessage) (string, error) {
requestBody := chatCompletionRequest{
Model: config.ChatModel,
Messages: messages,
Temperature: config.Temperature,
MaxTokens: config.MaxTokens,
Stream: false,
}
body, err := json.Marshal(requestBody)
if err != nil {
return "", fmt.Errorf("序列化大模型请求失败: %w", err)
}
endpoint, err := chatCompletionsEndpoint(config.ChatBaseURL)
if err != nil {
return "", err
}
requestCtx, cancel := context.WithTimeout(ctx, time.Duration(config.RequestTimeoutSeconds)*time.Second)
defer cancel()
req, err := http.NewRequestWithContext(requestCtx, http.MethodPost, endpoint, bytes.NewReader(body))
if err != nil {
return "", fmt.Errorf("创建大模型请求失败: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("User-Agent", fmt.Sprintf("DMP-AI/%s", utils.Version))
if config.ChatApiKey != "" {
req.Header.Set("Authorization", "Bearer "+config.ChatApiKey)
}
resp, err := c.httpClient.Do(req)
if err != nil {
return "", fmt.Errorf("请求大模型失败: %w", err)
}
defer resp.Body.Close()
responseBody, err := io.ReadAll(io.LimitReader(resp.Body, 2*1024*1024))
if err != nil {
return "", fmt.Errorf("读取大模型响应失败: %w", err)
}
var result chatCompletionResponse
if len(responseBody) > 0 {
if err = json.Unmarshal(responseBody, &result); err != nil {
if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
return "", fmt.Errorf("大模型响应异常 HTTP %d", resp.StatusCode)
}
return "", fmt.Errorf("解析大模型响应失败: %w", err)
}
}
if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
if result.Error != nil && result.Error.Message != "" {
return "", fmt.Errorf("大模型响应异常 HTTP %d: %s", resp.StatusCode, result.Error.Message)
}
return "", fmt.Errorf("大模型响应异常 HTTP %d", resp.StatusCode)
}
if result.Error != nil && result.Error.Message != "" {
return "", fmt.Errorf("大模型响应异常: %s", result.Error.Message)
}
if len(result.Choices) == 0 {
return "", fmt.Errorf("大模型响应中没有候选答案")
}
answer := strings.TrimSpace(result.Choices[0].Message.Content)
if answer == "" {
return "", fmt.Errorf("大模型返回了空答案")
}
return answer, nil
}
func chatCompletionsEndpoint(baseURL string) (string, error) {
u, err := url.Parse(strings.TrimSpace(baseURL))
if err != nil {
return "", fmt.Errorf("解析大模型 Base URL 失败: %w", err)
}
path := strings.TrimRight(u.Path, "/")
if !strings.HasSuffix(path, "/chat/completions") {
if path == "" {
path = "/chat/completions"
} else {
path += "/chat/completions"
}
}
u.Path = path
return u.String(), nil
}
config.go
package aichat
import (
"dst-management-platform-api/database/models"
"dst-management-platform-api/utils"
"fmt"
"strings"
"unicode/utf8"
)
func validateModelConfig(config *models.AIModelConfig) error {
config.ChatBaseURL = strings.TrimSpace(config.ChatBaseURL)
config.ChatModel = strings.TrimSpace(config.ChatModel)
config.EmbeddingModel = strings.TrimSpace(config.EmbeddingModel)
if config.ChatBaseURL == "" || !utils.IsValidURL(config.ChatBaseURL) {
return fmt.Errorf("大模型 Base URL 不合法")
}
if config.ChatModel == "" {
return fmt.Errorf("大模型名称不能为空")
}
if utf8.RuneCountInString(config.ChatModel) > 256 {
return fmt.Errorf("大模型名称过长")
}
if len(config.ChatApiKey) > 16*1024 {
return fmt.Errorf("API Key 过长")
}
if config.EmbeddingModel != "" {
if utf8.RuneCountInString(config.EmbeddingModel) > 256 {
return fmt.Errorf("嵌入模型名称过长")
}
if len(config.EmbeddingApiKey) > 16*1024 {
return fmt.Errorf("嵌入 API Key 过长")
}
}
if utf8.RuneCountInString(config.EmbeddingBaseURL) > 8000 {
return fmt.Errorf("系统提示词过长")
}
if config.Temperature < 0 || config.Temperature > 2 {
return fmt.Errorf("temperature 必须在 0 到 2 之间")
}
if config.MaxTokens <= 0 || config.MaxTokens > 32768 {
return fmt.Errorf("maxTokens 必须在 1 到 32768 之间")
}
if config.RequestTimeoutSeconds <= 0 || config.RequestTimeoutSeconds > 300 {
return fmt.Errorf("请求超时时间必须在 1 到 300 秒之间")
}
return nil
}
func validateRoomSetting(setting *models.RoomAISetting) error {
if setting.RoomID <= 0 {
return fmt.Errorf("房间 ID 不合法")
}
if strings.ContainsAny(setting.Prefix, "\r\n") || utf8.RuneCountInString(setting.Prefix) > 64 {
return fmt.Errorf("AI 对话前缀不合法")
}
if setting.ContextMaxMessages < 2 || setting.ContextMaxMessages > 100 {
return fmt.Errorf("上下文消息数量必须在 2 到 100 之间")
}
if setting.ContextTTLMinutes <= 0 || setting.ContextTTLMinutes > 10080 {
return fmt.Errorf("上下文有效期必须在 1 到 10080 分钟之间")
}
return ValidateModelConfig(&setting.AIModelConfig)
}
manager.go
package aichat
import (
"context"
"dst-management-platform-api/database/dao"
"dst-management-platform-api/database/models"
"dst-management-platform-api/dst"
"dst-management-platform-api/logger"
"errors"
"fmt"
"os"
"strings"
"time"
"unicode/utf8"
"gorm.io/gorm"
)
const (
chatLogBufferSize = 256
maxQuestionRunes = 2000
maxReplyRunes = 180
)
type chatSession struct {
messages []chatMessage
lastActive time.Time
}
type roomWorker struct {
cancel context.CancelFunc
}
func newManager(roomAISettingDao *dao.RoomAISettingDAO, pluginDao *dao.PluginDAO) *Manager {
ctx, cancel := context.WithCancel(context.Background())
aiManager := &Manager{
roomAISettingDao: roomAISettingDao,
client: newClient(),
ctx: ctx,
cancel: cancel,
workers: make(map[int]*roomWorker),
keywordSearcher: newKeywordWikiSearcher(wikiPagesDir, wikiIndexFile),
}
chatPlugin, err := pluginDao.GetPluginByPluginName(models.PluginChat)
if err != nil {
logger.Logger.Errorf("获取 %s 插件状态失败, err: %v", models.PluginChat, err)
} else if chatPlugin.Status {
if err = aiManager.start(); err != nil {
logger.Logger.Errorf("启动游戏内 AI 对话服务失败, err: %v", err)
}
}
return aiManager
}
func (m *Manager) start() error {
m.lifecycle.Lock()
if m.closed {
m.lifecycle.Unlock()
return context.Canceled
}
m.active = true
m.lifecycle.Unlock()
settings, err := m.roomAISettingDao.ListEnabled()
if err != nil {
m.stopAll()
return err
}
for _, setting := range settings {
if err = m.reload(setting.RoomID); err != nil {
logger.Logger.Errorf("启动房间 AI 对话监听失败, roomID: %d, err: %v", setting.RoomID, err)
}
}
return nil
}
func (m *Manager) reload(roomID int) error {
m.lifecycle.Lock()
defer m.lifecycle.Unlock()
m.stopRoom(roomID)
if !m.active || m.closed {
return nil
}
setting, err := m.roomAISettingDao.GetByRoomID(roomID)
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil
}
if err != nil {
return err
}
if !setting.Enabled {
return nil
}
if err = validateRoomSetting(setting); err != nil {
return err
}
room, worlds, roomSetting, err := dao.FetchGameInfo(roomID)
if err != nil {
return err
}
game := dst.NewGameController(room, worlds, roomSetting, "zh")
workerCtx, cancel := context.WithCancel(m.ctx)
m.mu.Lock()
m.workers[roomID] = &roomWorker{cancel: cancel}
m.mu.Unlock()
settingCopy := *setting
go m.runRoom(workerCtx, game, settingCopy)
return nil
}
func (m *Manager) stopAll() {
m.lifecycle.Lock()
defer m.lifecycle.Unlock()
m.active = false
m.stopAllWorkers()
}
func (m *Manager) stopRoom(roomID int) {
m.mu.Lock()
worker, ok := m.workers[roomID]
if ok {
delete(m.workers, roomID)
}
m.mu.Unlock()
if ok {
worker.cancel()
}
}
func (m *Manager) closeManager() {
m.lifecycle.Lock()
defer m.lifecycle.Unlock()
m.active = false
m.closed = true
m.cancel()
m.stopAllWorkers()
}
func (m *Manager) stopAllWorkers() {
m.mu.Lock()
workers := m.workers
m.workers = make(map[int]*roomWorker)
m.mu.Unlock()
for _, worker := range workers {
worker.cancel()
}
}
func (m *Manager) runRoom(ctx context.Context, game *dst.Game, setting models.RoomAISetting) {
lines := make(chan string, chatLogBufferSize)
go m.watchChatLog(ctx, game, setting.RoomID, lines)
sessions := make(map[string]*chatSession)
cleanupTicker := time.NewTicker(time.Minute)
defer cleanupTicker.Stop()
for {
select {
case <-ctx.Done():
return
case <-cleanupTicker.C:
cleanupSessions(sessions, time.Duration(setting.ContextTTLMinutes)*time.Minute)
case line := <-lines:
event, ok := parseChatEvent(line)
if !ok {
continue
}
question, ok := matchQuestion(setting.Prefix, event.Message)
if !ok {
continue
}
m.answer(ctx, game, setting, sessions, event, question)
}
}
}
func (m *Manager) watchChatLog(ctx context.Context, game *dst.Game, roomID int, lines chan<- string) {
for {
err := game.TailChatLog(ctx, 0, lines)
if ctx.Err() != nil {
return
}
logger.Logger.Errorf("房间聊天日志监听异常, roomID: %d, err: %v", roomID, err)
timer := time.NewTimer(2 * time.Second)
select {
case <-ctx.Done():
timer.Stop()
return
case <-timer.C:
}
}
}
// isEmbeddingConfigured 检查是否配置了嵌入模型
func (m *Manager) isEmbeddingConfigured(setting models.RoomAISetting) bool {
return setting.EmbeddingModel != "" && setting.ChatBaseURL != ""
}
// getEmbeddingSearcher 获取或创建向量搜索引擎
// 仅创建 searcher 实例并缓存,不自动构建索引(构建需通过 BuildEmbeddingIndex 手动触发)
func (m *Manager) getEmbeddingSearcher(setting models.RoomAISetting) *embeddingWikiSearcher {
// EmbeddingApiKey 如果单独配置了则优先使用,否则复用 ChatApiKey
apiKey := setting.ChatApiKey
if setting.EmbeddingApiKey != "" {
apiKey = setting.EmbeddingApiKey
}
configKey := setting.ChatBaseURL + "|" + apiKey + "|" + setting.EmbeddingModel
m.embedSearcherMu.Lock()
defer m.embedSearcherMu.Unlock()
// 配置未变,返回缓存的 searcher
if m.embedSearcher != nil && m.lastEmbedConfig == configKey {
return m.embedSearcher
}
// 配置变更或首次调用 — 创建新的 searcher(不自动构建索引)
if m.embedSearcher != nil {
m.embedSearcher.stopIdleTimer()
}
embedConfig := EmbeddingConfig{
APIURL: setting.ChatBaseURL,
APIKey: apiKey,
Model: setting.EmbeddingModel,
Dimensions: 1024,
}
searcher := newEmbeddingWikiSearcher(wikiPagesDir, embedConfig)
m.embedSearcher = searcher
m.lastEmbedConfig = configKey
return m.embedSearcher
}
// BuildEmbeddingIndex 手动构建向量索引。
// config 指定 Embedding API 配置,force 为 true 时强制全量重建。
// 构建过程可能耗时较长,建议在后台执行。
func (m *Manager) buildEmbeddingIndex(config EmbeddingConfig, force bool) error {
searcher := newEmbeddingWikiSearcher(wikiPagesDir, config)
if !force && !searcher.needsSetup() {
logger.Logger.Infof("向量索引已存在,无需重建。使用 force=true 强制重建。")
return nil
}
logger.Logger.Infof("开始手动构建向量索引 (force=%v)...", force)
if err := searcher.buildIndex(force); err != nil {
return fmt.Errorf("构建向量索引失败: %w", err)
}
// 更新缓存的 searcher
m.embedSearcherMu.Lock()
if m.embedSearcher != nil {
m.embedSearcher.stopIdleTimer()
}
apiKey := config.APIKey
configKey := config.APIURL + "|" + apiKey + "|" + config.Model
m.embedSearcher = searcher
m.lastEmbedConfig = configKey
m.embedSearcherMu.Unlock()
logger.Logger.Infof("向量索引构建完成")
return nil
}
// GetEmbeddingStats 获取向量索引统计信息
func (m *Manager) getEmbeddingStats(config EmbeddingConfig) EmbeddingStats {
searcher := newEmbeddingWikiSearcher(wikiPagesDir, config)
defer searcher.stopIdleTimer()
return searcher.getStats()
}
// BuildKeywordIndex 手动构建关键词搜索索引。
// force 为 true 时强制重建,否则索引已存在则跳过。
func (m *Manager) buildKeywordIndex(force bool) error {
// 停止旧 searcher 的空闲计时器
if m.keywordSearcher != nil {
m.keywordSearcher.stopIdleTimer()
}
searcher := newKeywordWikiSearcher(wikiPagesDir, wikiIndexFile)
if !force {
if _, err := os.Stat(wikiIndexFile); err == nil {
if loadErr := searcher.load(); loadErr == nil {
logger.Logger.Infof("关键词索引已存在,无需重建。使用 force=true 强制重建。")
m.keywordSearcher = searcher
return nil
}
}
}
logger.Logger.Infof("开始构建关键词搜索索引 (force=%v)...", force)
if err := searcher.buildIndex(force); err != nil {
return fmt.Errorf("构建关键词索引失败: %w", err)
}
m.keywordSearcher = searcher
logger.Logger.Infof("关键词搜索索引构建完成")
return nil
}
// UnloadKeywordIndex 释放关键词搜索索引占用的内存。
// 下次 Search 时会自动从磁盘重新加载。
func (m *Manager) unloadKeywordIndex() {
if m.keywordSearcher != nil {
m.keywordSearcher.unload()
}
}
// searchWiki 搜索 Wiki 知识库,返回格式化的参考上下文
func (m *Manager) searchWiki(setting models.RoomAISetting, question string) string {
// 优先使用向量搜索
if m.isEmbeddingConfigured(setting) {
if searcher := m.getEmbeddingSearcher(setting); searcher != nil {
results, err := searcher.search(question, nil, 3, 0.3)
if err != nil {
logger.Logger.Warnf("向量搜索 Wiki 失败: %v,回退到关键词搜索", err)
} else if len(results) > 0 {
return formatWikiContext(results, maxContextTokens)
}
}
}
// 回退到关键词搜索
if m.keywordSearcher != nil {
results, err := m.keywordSearcher.search(question, 3)
if err != nil {
logger.Logger.Warnf("关键词搜索 Wiki 失败: %v", err)
return ""
}
if len(results) > 0 {
return formatWikiContext(results, maxContextTokens)
}
}
return ""
}
// buildSystemPrompt 构建系统提示词(Wiki 上下文 + 用户设定的提示词)
func buildSystemPrompt(setting models.RoomAISetting, wikiContext string) string {
var parts []string
// 默认提示词
defaultPrompt := "你是饥荒联机版游戏内的 AI 助手。请根据以上参考文档,用中文回答玩家的问题。回答应简洁、准确,适合在游戏聊天框中显示,坚决不能使用使用 Markdown 格式,回答不能超过30个字。"
// Wiki 参考文档
if wikiContext != "" {
parts = append(parts, wikiContext)
parts = append(parts, "")
parts = append(parts, defaultPrompt)
return strings.Join(parts, "\n")
}
// 用户设定的系统提示词(存储在 EmbeddingBaseURL 字段中)
systemPrompt := strings.TrimSpace(setting.EmbeddingBaseURL)
if systemPrompt != "" {
return systemPrompt
}
return defaultPrompt
}
func (m *Manager) answer(ctx context.Context, game *dst.Game, setting models.RoomAISetting, sessions map[string]*chatSession, event chatEvent, question string) {
now := time.Now()
ttl := time.Duration(setting.ContextTTLMinutes) * time.Minute
session := sessions[event.UID]
if session == nil || now.Sub(session.lastActive) >= ttl {
session = &chatSession{}
sessions[event.UID] = session
}
// 搜索 Wiki 知识库获取参考上下文
wikiContext := m.searchWiki(setting, question)
userMessage := chatMessage{Role: "user", Content: truncateRunes(question, maxQuestionRunes)}
messages := make([]chatMessage, 0, len(session.messages)+3)
// 构建系统提示词(Wiki 上下文 + 用户设定的提示词)
systemContent := buildSystemPrompt(setting, wikiContext)
if systemContent != "" {
messages = append(messages, chatMessage{Role: "system", Content: systemContent})
}
messages = append(messages, session.messages...)
messages = append(messages, userMessage)
answer, err := m.client.complete(ctx, setting.AIModelConfig, messages)
if err != nil {
if ctx.Err() != nil {
return
}
logger.Logger.Errorf("游戏内 AI 回答失败, roomID: %d, uid: %s, err: %v", setting.RoomID, event.UID, err)
_ = sendGameReply(game, event.Nickname, "暂时无法回答,请稍后再试。")
return
}
session.messages = append(session.messages, userMessage, chatMessage{Role: "assistant", Content: answer})
maxMessages := setting.ContextMaxMessages
if maxMessages%2 != 0 {
maxMessages--
}
if len(session.messages) > maxMessages {
session.messages = append([]chatMessage(nil), session.messages[len(session.messages)-maxMessages:]...)
}
session.lastActive = time.Now()
if err = sendGameReply(game, event.Nickname, answer); err != nil {
logger.Logger.Errorf("发送游戏内 AI 回答失败, roomID: %d, uid: %s, err: %v", setting.RoomID, event.UID, err)
}
}
func cleanupSessions(sessions map[string]*chatSession, ttl time.Duration) {
now := time.Now()
for uid, session := range sessions {
if now.Sub(session.lastActive) >= ttl {
delete(sessions, uid)
}
}
}
func matchQuestion(prefix, message string) (string, bool) {
message = strings.TrimSpace(message)
if prefix != "" {
if !strings.HasPrefix(message, prefix) {
return "", false
}
message = strings.TrimSpace(strings.TrimPrefix(message, prefix))
}
return message, message != ""
}
func sendGameReply(game *dst.Game, nickname, answer string) error {
answer = strings.Join(strings.Fields(answer), " ")
if answer == "" {
return fmt.Errorf("AI 回答为空")
}
prefix := fmt.Sprintf("[AI] %s: ", nickname)
for _, part := range splitRunes(answer, maxReplyRunes-utf8.RuneCountInString(prefix)) {
if err := game.SystemMsg(prefix + part); err != nil {
return err
}
}
return nil
}
func splitRunes(value string, size int) []string {
if size <= 0 {
size = maxReplyRunes
}
runes := []rune(value)
parts := make([]string, 0, (len(runes)+size-1)/size)
for len(runes) > 0 {
n := size
if len(runes) < n {
n = len(runes)
}
parts = append(parts, string(runes[:n]))
runes = runes[n:]
}
return parts
}
func truncateRunes(value string, size int) string {
runes := []rune(value)
if len(runes) <= size {
return value
}
return string(runes[:size])
}
parser.go
package aichat
import (
"regexp"
"strings"
)
var sayLogPattern = regexp.MustCompile(`^\[[^\]]+\]:\s*\[Say\]\s*\(([^)]+)\)\s+(.+?):\s*(.*)$`)
type chatEvent struct {
UID string
Nickname string
Message string
}
func parseChatEvent(line string) (chatEvent, bool) {
matches := sayLogPattern.FindStringSubmatch(strings.TrimSpace(line))
if len(matches) != 4 {
return chatEvent{}, false
}
event := chatEvent{
UID: strings.TrimSpace(matches[1]),
Nickname: strings.TrimSpace(matches[2]),
Message: strings.TrimSpace(matches[3]),
}
if event.UID == "" || event.Nickname == "" || event.Message == "" {
return chatEvent{}, false
}
return event, true
}
wiki_embedding.go
package aichat
import (
"bytes"
"dst-management-platform-api/logger"
"dst-management-platform-api/utils"
"encoding/json"
"fmt"
"io"
"math"
"net/http"
"os"
"path/filepath"
"regexp"
"sort"
"strings"
"sync"
"time"
"unicode/utf8"
)
// ========== 配置 ==========
const embeddingDir = utils.PluginAiChatSearchDataPath + "/embeddings"
// ========== 构建索引常量 ==========
const (
embedBatchSize = 10
embedMaxChars = 4000
embedRequestInterval = 300 * time.Millisecond
embedMinBatchSize = 5
embedMinContentLength = 50
embeddingIdleTimeout = 5 * time.Minute
)
// ========== 预处理正则 ==========
var (
embedMdTitleRe = regexp.MustCompile(`(?m)^#\s+.+$`)
embedMdCategoryRe = regexp.MustCompile(`(?m)^\*\*分类\*\*:.*$`)
embedMdHrRe = regexp.MustCompile(`(?m)^---$`)
embedMdLinkRe = regexp.MustCompile(`\[([^\]]+)\]\([^)]+\)`)
embedMdImgRe = regexp.MustCompile(`!\[.*?\]\(.*?\)`)
embedMdFmtRe = regexp.MustCompile(`[*#>` + "`" + `|]`)
embedWsRe = regexp.MustCompile(`\s+`)
embedCatLinkRe = regexp.MustCompile(`\[([^\]]+)\]\(`)
)
// ========== 元数据结构 ==========
type wikiMeta struct {
Title string `json:"title"`
Categories []string `json:"categories"`
Content string `json:"content"`
ContentLength int `json:"content_length"`
Filename string `json:"filename"`
}
// ========== 向量搜索引擎 ==========
// EmbeddingWikiSearcher 基于向量相似度的 Wiki 搜索引擎
type embeddingWikiSearcher struct {
pagesDir string
embeddingDir string
apiURL string
apiKey string
model string
dimensions int
httpClient *http.Client
mu sync.RWMutex
metadata map[string]wikiMeta // filename -> meta
embeddings map[string][]float64 // filename -> vector
loaded bool
idleTimer *time.Timer // 空闲自动释放计时器
}
// newEmbeddingWikiSearcher 创建向量搜索引擎
func newEmbeddingWikiSearcher(pagesDir string, config EmbeddingConfig) *embeddingWikiSearcher {
if config.Dimensions <= 0 {
config.Dimensions = 1024
}
return &embeddingWikiSearcher{
pagesDir: pagesDir,
embeddingDir: embeddingDir,
apiURL: strings.TrimRight(config.APIURL, "/"),
apiKey: config.APIKey,
model: config.Model,
dimensions: config.Dimensions,
httpClient: &http.Client{
Timeout: 120 * time.Second,
},
}
}
// NeedsSetup 是否需要构建索引
func (s *embeddingWikiSearcher) needsSetup() bool {
s.mu.RLock()
defer s.mu.RUnlock()
if !s.loaded {
s.mu.RUnlock()
s.load()
s.mu.RLock()
}
return len(s.embeddings) == 0
}
// Search 向量语义搜索
func (s *embeddingWikiSearcher) search(query string, categories []string, maxResults int, minScore float64) ([]wikiResult, error) {
s.mu.RLock()
if !s.loaded {
s.mu.RUnlock()
s.load()
s.mu.RLock()
}
embeddings := s.embeddings
metadata := s.metadata
s.mu.RUnlock()
// 每次搜索重置空闲计时器
s.mu.Lock()
s.resetIdleTimer()
s.mu.Unlock()
if len(embeddings) == 0 {
return nil, fmt.Errorf("向量索引为空")
}
// 对查询文本做 embedding
queryVector, err := s.embedSingle(query)
if err != nil {
return nil, fmt.Errorf("查询文本 embedding 失败: %w", err)
}
// 计算余弦相似度
type scoredDoc struct {
filename string
score float64
}
var scores []scoredDoc
for filename, docVec := range embeddings {
meta, ok := metadata[filename]
if !ok || meta.ContentLength < embedMinContentLength {
continue
}
if len(categories) > 0 {
catSet := make(map[string]struct{}, len(categories))
for _, c := range categories {
catSet[c] = struct{}{}
}
hasCat := false
for _, c := range meta.Categories {
if _, ok := catSet[c]; ok {
hasCat = true
break
}
}
if !hasCat {
continue
}
}
sim := cosineSimilarity(queryVector, docVec)
if sim >= minScore {
scores = append(scores, scoredDoc{filename, sim})
}
}
// 排序取 Top N
sort.Slice(scores, func(i, j int) bool {
return scores[i].score > scores[j].score
})
results := make([]wikiResult, 0, maxResults)
for i := 0; i < len(scores) && i < maxResults; i++ {
meta := metadata[scores[i].filename]
results = append(results, wikiResult{
Title: meta.Title,
Filename: meta.Filename,
Categories: meta.Categories,
Content: meta.Content,
ContentLen: meta.ContentLength,
Score: roundFloat(scores[i].score, 4),
})
}
return results, nil
}
// ========== 内部方法 ==========
func (s *embeddingWikiSearcher) load() {
s.mu.Lock()
defer s.mu.Unlock()
s.loadLocked()
}
// loadLocked 不加锁地从磁盘加载(调用方需持有 s.mu 写锁)
func (s *embeddingWikiSearcher) loadLocked() {
if s.loaded {
return
}
metadataPath := filepath.Join(s.embeddingDir, "metadata.json")
embeddingsPath := filepath.Join(s.embeddingDir, "embeddings.json")
if data, err := os.ReadFile(metadataPath); err == nil {
err = json.Unmarshal(data, &s.metadata)
if err != nil {
logger.Logger.Errorf("载入向量模型失败: %v", err)
}
}
if s.metadata == nil {
s.metadata = make(map[string]wikiMeta)
}
if data, err := os.ReadFile(embeddingsPath); err == nil {
err = json.Unmarshal(data, &s.embeddings)
if err != nil {
logger.Logger.Errorf("载入向量模型失败: %v", err)
}
}
if s.embeddings == nil {
s.embeddings = make(map[string][]float64)
}
s.loaded = true
s.resetIdleTimer()
logger.Logger.Infof("已加载向量索引: %d 篇文档", len(s.embeddings))
}
// unload 释放向量索引占用的内存。下次 search 时会自动从磁盘重新加载。
func (s *embeddingWikiSearcher) unload() {
if s.embeddings != nil {
logger.Logger.Infof("释放向量索引内存 (%d 篇文档)", len(s.embeddings))
}
s.metadata = nil
s.embeddings = nil
s.loaded = false
s.stopIdleTimer()
}
// resetIdleTimer 重置空闲计时器。调用方需持有 s.mu 写锁。
func (s *embeddingWikiSearcher) resetIdleTimer() {
if s.idleTimer != nil {
s.idleTimer.Stop()
}
s.idleTimer = time.AfterFunc(embeddingIdleTimeout, func() {
s.mu.Lock()
//logger.Logger.Infof("向量索引空闲超过 %v,自动释放内存", embeddingIdleTimeout)
s.unload()
s.mu.Unlock()
})
}
// stopIdleTimer 停止空闲计时器
func (s *embeddingWikiSearcher) stopIdleTimer() {
if s.idleTimer != nil {
s.idleTimer.Stop()
s.idleTimer = nil
}
}
func (s *embeddingWikiSearcher) embedSingle(text string) ([]float64, error) {
vectors, err := s.embedBatch([]string{text})
if err != nil {
return nil, err
}
if len(vectors) == 0 {
return nil, fmt.Errorf("embedding 返回为空")
}
return vectors[0], nil
}
func (s *embeddingWikiSearcher) embedBatch(texts []string) ([][]float64, error) {
// 构建请求
reqBody := map[string]interface{}{
"model": s.model,
"input": texts,
"encoding_format": "float",
}
if s.dimensions > 0 {
reqBody["dimensions"] = s.dimensions
}
body, err := json.Marshal(reqBody)
if err != nil {
return nil, fmt.Errorf("序列化 embedding 请求失败: %w", err)
}
endpoint := s.apiURL + "/embeddings"
req, err := http.NewRequest(http.MethodPost, endpoint, bytes.NewReader(body))
if err != nil {
return nil, fmt.Errorf("创建 embedding 请求失败: %w", err)
}
req.Header.Set("Content-Type", "application/json")
if s.apiKey != "" {
req.Header.Set("Authorization", "Bearer "+s.apiKey)
}
resp, err := s.httpClient.Do(req)
if err != nil {
return nil, fmt.Errorf("embedding API 请求失败: %w", err)
}
defer resp.Body.Close()
respBody, err := io.ReadAll(io.LimitReader(resp.Body, 32*1024*1024))
if err != nil {
return nil, fmt.Errorf("读取 embedding 响应失败: %w", err)
}
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("embedding API 返回 HTTP %d: %s", resp.StatusCode, string(respBody))
}
var result struct {
Data []struct {
Embedding []float64 `json:"embedding"`
Index int `json:"index"`
} `json:"data"`
Error *struct {
Message string `json:"message"`
} `json:"error"`
}
if err := json.Unmarshal(respBody, &result); err != nil {
return nil, fmt.Errorf("解析 embedding 响应失败: %w", err)
}
if result.Error != nil && result.Error.Message != "" {
return nil, fmt.Errorf("embedding API 错误: %s", result.Error.Message)
}
// 按 index 排序
sort.Slice(result.Data, func(i, j int) bool {
return result.Data[i].Index < result.Data[j].Index
})
vectors := make([][]float64, len(result.Data))
for i, d := range result.Data {
vectors[i] = d.Embedding
}
return vectors, nil
}
func cosineSimilarity(a, b []float64) float64 {
if len(a) != len(b) || len(a) == 0 {
return 0
}
var dot, normA, normB float64
for i := range a {
dot += a[i] * b[i]
normA += a[i] * a[i]
normB += b[i] * b[i]
}
if normA == 0 || normB == 0 {
return 0
}
return dot / (math.Sqrt(normA) * math.Sqrt(normB))
}
// ================================================================
// 构建向量索引
// ================================================================
// save 不加锁保存(调用方需持有 s.mu 读锁或写锁)
func (s *embeddingWikiSearcher) save() error {
if err := os.MkdirAll(s.embeddingDir, 0755); err != nil {
return fmt.Errorf("创建 embedding 目录失败: %w", err)
}
metadataPath := filepath.Join(s.embeddingDir, "metadata.json")
embeddingsPath := filepath.Join(s.embeddingDir, "embeddings.json")
data, err := json.Marshal(s.metadata)
if err != nil {
return fmt.Errorf("序列化 metadata 失败: %w", err)
}
if err := os.WriteFile(metadataPath, data, 0644); err != nil {
return fmt.Errorf("保存 metadata 失败: %w", err)
}
data, err = json.Marshal(s.embeddings)
if err != nil {
return fmt.Errorf("序列化 embeddings 失败: %w", err)
}
if err := os.WriteFile(embeddingsPath, data, 0644); err != nil {
return fmt.Errorf("保存 embeddings 失败: %w", err)
}
return nil
}
// BuildIndex 构建向量索引。只跑一次(或 force=true 强制重建)。
//
// 流程:
// - 扫描所有 .md 页面
// - 文本预处理
// - 分批调用 Embedding API
// - 每批完成后保存(支持断点续传)
// - 连接错误自动重试,严重时自动降级 batch_size
func (s *embeddingWikiSearcher) buildIndex(force bool) error {
s.mu.Lock()
defer s.mu.Unlock()
// 确保已加载已有数据
s.loadLocked()
// 扫描所有 .md 页面
entries, err := os.ReadDir(s.pagesDir)
if err != nil {
return fmt.Errorf("扫描 Wiki 页面目录失败: %w", err)
}
var mdFiles []string
for _, entry := range entries {
if !entry.IsDir() && filepath.Ext(entry.Name()) == ".md" {
mdFiles = append(mdFiles, filepath.Join(s.pagesDir, entry.Name()))
}
}
sort.Strings(mdFiles)
total := len(mdFiles)
// 增量构建 vs 全量重建
if !force {
var newFiles []string
for _, f := range mdFiles {
if _, ok := s.embeddings[filepath.Base(f)]; !ok {
newFiles = append(newFiles, f)
}
}
if len(newFiles) == 0 {
logger.Logger.Infof("所有 %d 篇文档已有向量,无需重建。用 force=true 强制重建。", total)
return nil
}
logger.Logger.Infof("增量构建: %d/%d 篇新文档需要 embedding", len(newFiles), total)
mdFiles = newFiles
} else {
s.embeddings = make(map[string][]float64)
s.metadata = make(map[string]wikiMeta)
}
// 分批处理
currentBatchSize := embedBatchSize
batches, totalBatches := makeBatches(mdFiles, currentBatchSize)
logger.Logger.Infof("共 %d 篇文档,分 %d 批 (batch_size=%d)", len(mdFiles), totalBatches, currentBatchSize)
logger.Logger.Infof("API: %s 模型: %s", s.apiURL, s.model)
var failedFiles []string
consecutiveFailures := 0
batchIdx := 0
for batchIdx < totalBatches {
batchFiles := batches[batchIdx]
// 准备这批的文本
texts, metas, prepErr := s.prepareBatch(batchFiles)
if prepErr != nil {
logger.Logger.Errorf("准备第 %d 批文档失败: %v", batchIdx+1, prepErr)
for _, f := range batchFiles {
failedFiles = append(failedFiles, filepath.Base(f))
}
batchIdx++
continue
}
// 调用 API
vectors, apiErr := s.embedBatch(texts)
if apiErr == nil {
// 成功
for i, fp := range batchFiles {
s.embeddings[filepath.Base(fp)] = vectors[i]
}
for k, v := range metas {
s.metadata[k] = v
}
consecutiveFailures = 0
if saveErr := s.save(); saveErr != nil {
logger.Logger.Errorf("保存 embedding 失败: %v", saveErr)
}
logger.Logger.Infof("第 %d/%d 批完成 (%d 篇)", batchIdx+1, totalBatches, len(s.embeddings))
if batchIdx < totalBatches-1 {
time.Sleep(embedRequestInterval)
}
batchIdx++
continue
}
// --- API 调用失败,错误恢复 ---
consecutiveFailures++
errorStr := strings.ToLower(apiErr.Error())
logger.Logger.Errorf("第 %d/%d 批失败: %v", batchIdx+1, totalBatches, apiErr)
// batch size 超限 → 立刻调整为 10
if strings.Contains(errorStr, "batch size") && currentBatchSize > 10 {
currentBatchSize = 10
logger.Logger.Infof("[自动调整] batch_size 调整为 %d(API 限制)", currentBatchSize)
batches, totalBatches = rebuildBatches(mdFiles, s.embeddings, currentBatchSize)
batchIdx = 0
consecutiveFailures = 0
continue
}
// 连续失败 3 次,降级 batch_size
if consecutiveFailures >= 3 && currentBatchSize > embedMinBatchSize {
currentBatchSize = max(embedMinBatchSize, currentBatchSize/2)
logger.Logger.Infof("[自动调整] batch_size 降为 %d,重新分批...", currentBatchSize)
batches, totalBatches = rebuildBatches(mdFiles, s.embeddings, currentBatchSize)
batchIdx = 0
consecutiveFailures = 0
continue
}
// 把这一批拆成单篇重试
if currentBatchSize > 1 && len(batchFiles) > 1 {
logger.Logger.Infof("尝试逐篇处理这一批...")
for _, fp := range batchFiles {
singleTexts, singleMetas, prepErr := s.prepareBatch([]string{fp})
if prepErr != nil {
logger.Logger.Errorf("[跳过] %s: %v", filepath.Base(fp), prepErr)
failedFiles = append(failedFiles, filepath.Base(fp))
continue
}
vec, embErr := s.embedBatch(singleTexts)
if embErr != nil {
logger.Logger.Errorf("[跳过] %s: %v", filepath.Base(fp), embErr)
failedFiles = append(failedFiles, filepath.Base(fp))
continue
}
s.embeddings[filepath.Base(fp)] = vec[0]
for k, v := range singleMetas {
s.metadata[k] = v
}
time.Sleep(500 * time.Millisecond)
}
if saveErr := s.save(); saveErr != nil {
logger.Logger.Errorf("保存 embedding 失败: %v", saveErr)
}
batchIdx++
consecutiveFailures = 0
continue
}
// 单篇也失败,跳过
logger.Logger.Errorf("[跳过] 这 %d 篇 embed 失败", len(batchFiles))
for _, f := range batchFiles {
failedFiles = append(failedFiles, filepath.Base(f))
}
if saveErr := s.save(); saveErr != nil {
logger.Logger.Errorf("保存 embedding 失败: %v", saveErr)
}
batchIdx++
}
// 完成
s.resetIdleTimer()
logger.Logger.Infof("索引构建完成! 共 %d 篇文档。", len(s.embeddings))
if len(failedFiles) > 0 {
logger.Logger.Warnf("跳过 %d 篇:", len(failedFiles))
for i, f := range failedFiles {
if i >= 10 {
logger.Logger.Warnf("... 等 %d 篇", len(failedFiles)-10)
break
}
logger.Logger.Warnf(" - %s", f)
}
}
return nil
}
// ========== 批次准备 ==========
// prepareBatch 批量读取并预处理文档
func (s *embeddingWikiSearcher) prepareBatch(files []string) (texts []string, metas map[string]wikiMeta, err error) {
texts = make([]string, 0, len(files))
metas = make(map[string]wikiMeta, len(files))
for _, fp := range files {
content, readErr := os.ReadFile(fp)
if readErr != nil {
return nil, nil, fmt.Errorf("读取 %s 失败: %w", filepath.Base(fp), readErr)
}
rawContent := string(content)
title := extractWikiTitle(rawContent, filepath.Base(fp))
categories := extractWikiCategories(rawContent)
cleanText := preprocessWikiDoc(rawContent, title)
strippedContent := stripWikiMeta(rawContent)
texts = append(texts, cleanText)
metas[filepath.Base(fp)] = wikiMeta{
Title: title,
Categories: categories,
Content: strippedContent,
ContentLength: utf8.RuneCountInString(cleanText),
Filename: filepath.Base(fp),
}
}
return texts, metas, nil
}
// ========== 预处理 ==========
// extractWikiTitle 提取文档标题(markdown 一级标题)
func extractWikiTitle(rawContent, fallback string) string {
match := embedMdTitleRe.FindStringSubmatch(rawContent)
if len(match) >= 1 {
return strings.TrimSpace(strings.TrimPrefix(match[0], "# "))
}
// 回退:文件名去扩展名
name := fallback
if ext := filepath.Ext(name); ext != "" {
name = strings.TrimSuffix(name, ext)
}
return name
}
// extractWikiCategories 提取文档分类
func extractWikiCategories(rawContent string) []string {
match := embedMdCategoryRe.FindStringSubmatch(rawContent)
if len(match) < 1 {
return nil
}
linkMatches := embedCatLinkRe.FindAllStringSubmatch(match[0], -1)
cats := make([]string, 0, len(linkMatches))
for _, m := range linkMatches {
if len(m) >= 2 {
cats = append(cats, m[1])
}
}
return cats
}
// preprocessWikiDoc 预处理文档文本,使其适合 embedding
//
// 处理步骤:
// - 去 markdown 语法
// - 标题前置(标题是最重要的语义信号)
// - 截断到 embedMaxChars
func preprocessWikiDoc(rawContent, title string) string {
text := rawContent
// 去标题行、分类行、水平线
text = embedMdTitleRe.ReplaceAllString(text, "")
text = embedMdCategoryRe.ReplaceAllString(text, "")
text = embedMdHrRe.ReplaceAllString(text, "")
// 去 markdown 格式字符
text = embedMdLinkRe.ReplaceAllString(text, "$1") // [text](url) → text
text = embedMdImgRe.ReplaceAllString(text, " ")
text = embedMdFmtRe.ReplaceAllString(text, " ")
text = embedWsRe.ReplaceAllString(text, " ")
text = strings.TrimSpace(text)
// 标题前置
text = fmt.Sprintf("标题: %s\n正文: %s", title, text)
// 截断
runes := []rune(text)
if len(runes) > embedMaxChars {
text = string(runes[:embedMaxChars])
}
return text
}
// stripWikiMeta 去除元数据行,保留纯正文
func stripWikiMeta(rawContent string) string {
text := rawContent
text = embedMdTitleRe.ReplaceAllString(text, "")
text = embedMdCategoryRe.ReplaceAllString(text, "")
text = embedMdHrRe.ReplaceAllString(text, "")
return strings.TrimSpace(text)
}
// ========== 辅助函数 ==========
func makeBatches(files []string, size int) (batches [][]string, total int) {
for i := 0; i < len(files); i += size {
end := i + size
if end > len(files) {
end = len(files)
}
batches = append(batches, files[i:end])
}
return batches, len(batches)
}
func rebuildBatches(allFiles []string, embedded map[string][]float64, size int) ([][]string, int) {
var remaining []string
for _, f := range allFiles {
if _, ok := embedded[filepath.Base(f)]; !ok {
remaining = append(remaining, f)
}
}
logger.Logger.Infof("剩余 %d 篇,分为 %d 批", len(remaining), (len(remaining)+size-1)/size)
return makeBatches(remaining, size)
}
// ================================================================
// 索引统计
// ================================================================
// GetStats 获取索引统计信息
func (s *embeddingWikiSearcher) getStats() EmbeddingStats {
s.mu.RLock()
defer s.mu.RUnlock()
if !s.loaded {
s.mu.RUnlock()
s.load()
s.mu.RLock()
}
stats := EmbeddingStats{
TotalDocs: len(s.embeddings),
Categories: make(map[string]int),
}
for _, meta := range s.metadata {
for _, c := range meta.Categories {
stats.Categories[c]++
}
}
stats.TotalCategories = len(stats.Categories)
if len(s.embeddings) > 0 {
for _, v := range s.embeddings {
stats.VectorDim = len(v)
break
}
}
return stats
}
wiki_keyword.go
package aichat
import (
"dst-management-platform-api/logger"
"dst-management-platform-api/utils"
"encoding/json"
"fmt"
"os"
"path/filepath"
"regexp"
"sort"
"strings"
"sync"
"time"
"unicode"
"unicode/utf8"
)
// ========== 配置 ==========
const (
wikiPagesDir = utils.PluginAiChatWikiPath + "/pages"
wikiIndexFile = utils.PluginAiChatSearchDataPath + "/search_index.json"
keywordIndexIdleTimeout = 5 * time.Minute
)
// ========== 数据结构 ==========
// wikiResult 单条搜索结果
type wikiResult struct {
Title string `json:"title"`
Filename string `json:"filename"`
Categories []string `json:"categories"`
Content string `json:"content"`
ContentLen int `json:"content_length"`
Score float64 `json:"score"`
}
type wikiPageInfo struct {
Title string `json:"title"`
Categories []string `json:"categories"`
Content string `json:"content"`
ContentLength int `json:"content_length"`
Links []string `json:"links"`
Filename string `json:"filename"`
Terms []string `json:"terms"`
}
type wikiIndex struct {
Pages map[string]wikiPageInfo `json:"pages"`
CategoryIndex map[string][]string `json:"category_index"`
TermIndex map[string]map[string]float64 `json:"term_index"`
}
// ========== 关键词搜索引擎 ==========
// KeywordWikiSearcher 基于关键词的 Wiki 搜索引擎
type keywordWikiSearcher struct {
pagesDir string
indexPath string
mu sync.RWMutex
index *wikiIndex
idleTimer *time.Timer // 空闲自动释放计时器
}
// NewKeywordWikiSearcher 创建关键词搜索引擎
func newKeywordWikiSearcher(pagesDir, indexPath string) *keywordWikiSearcher {
return &keywordWikiSearcher{
pagesDir: pagesDir,
indexPath: indexPath,
}
}
// Load 从磁盘加载搜索索引
func (s *keywordWikiSearcher) load() error {
s.mu.Lock()
defer s.mu.Unlock()
return s.loadLocked()
}
// Unload 释放搜索索引占用的内存。下次 Search 时会自动重新从磁盘加载。
func (s *keywordWikiSearcher) unload() {
s.mu.Lock()
defer s.mu.Unlock()
if s.index != nil {
logger.Logger.Infof("释放关键词搜索索引内存 (%d 篇文档)", len(s.index.Pages))
s.index = nil
}
s.stopIdleTimer()
}
// IsLoaded 检查索引是否已加载到内存
func (s *keywordWikiSearcher) isLoaded() bool {
s.mu.RLock()
defer s.mu.RUnlock()
return s.index != nil
}
// resetIdleTimer 重置空闲计时器。调用方需持有 s.mu 写锁。
// 每次 Search 调用时重置,超时后自动 Unload。
func (s *keywordWikiSearcher) resetIdleTimer() {
if s.idleTimer != nil {
s.idleTimer.Stop()
}
s.idleTimer = time.AfterFunc(keywordIndexIdleTimeout, func() {
//logger.Logger.Infof("关键词搜索索引空闲超过 %v,自动释放内存", keywordIndexIdleTimeout)
s.mu.Lock()
s.unload()
s.mu.Unlock()
})
}
// stopIdleTimer 停止空闲计时器
func (s *keywordWikiSearcher) stopIdleTimer() {
if s.idleTimer != nil {
s.idleTimer.Stop()
s.idleTimer = nil
}
}
// loadLocked 不加锁加载(调用方需持有 s.mu 写锁)
func (s *keywordWikiSearcher) loadLocked() error {
if s.index != nil {
return nil
}
data, err := os.ReadFile(s.indexPath)
if err != nil {
return fmt.Errorf("搜索索引文件不存在: %w", err)
}
var idx wikiIndex
if err := json.Unmarshal(data, &idx); err != nil {
return fmt.Errorf("解析搜索索引失败: %w", err)
}
s.index = &idx
s.resetIdleTimer()
logger.Logger.Infof("已加载关键词搜索索引: %d 篇文档, %d 个词条", len(idx.Pages), len(idx.TermIndex))
return nil
}
// Search 搜索 Wiki 文档。索引未构建时返回 error。
func (s *keywordWikiSearcher) search(query string, maxResults int) ([]wikiResult, error) {
s.mu.Lock()
if s.index == nil {
if err := s.loadLocked(); err != nil {
s.mu.Unlock()
return nil, fmt.Errorf("搜索索引未构建,请先构建索引: %w", err)
}
}
idx := s.index
s.resetIdleTimer() // 每次 Search 重置空闲计时器
s.mu.Unlock()
queryTerms := tokenizeQuery(query)
scores := make(map[string]float64)
matchedTerms := make(map[string][]string)
queryLower := strings.ToLower(strings.TrimSpace(query))
// TF 评分
for term := range queryTerms {
if docs, ok := idx.TermIndex[term]; ok {
for filename, tfScore := range docs {
scores[filename] += tfScore
matchedTerms[filename] = append(matchedTerms[filename], term)
}
}
}
// 标题匹配加权
for filename, info := range idx.Pages {
titleLower := strings.ToLower(info.Title)
if strings.Contains(titleLower, queryLower) {
scores[filename] += 5.0
matchedTerms[filename] = append(matchedTerms[filename], "标题精确匹配")
}
for term := range queryTerms {
if utf8.RuneCountInString(term) >= 3 && strings.Contains(titleLower, term) {
scores[filename] += 1.0
}
}
}
// 排序
type scoredFile struct {
filename string
score float64
}
ranked := make([]scoredFile, 0, len(scores))
for filename, score := range scores {
ranked = append(ranked, scoredFile{filename, score})
}
sort.Slice(ranked, func(i, j int) bool {
return ranked[i].score > ranked[j].score
})
// 取 Top N
results := make([]wikiResult, 0, maxResults)
for _, sf := range ranked {
if sf.score <= 0 {
continue
}
info := idx.Pages[sf.filename]
if info.ContentLength < 50 {
continue
}
results = append(results, wikiResult{
Title: info.Title,
Filename: info.Filename,
Categories: info.Categories,
Content: info.Content,
ContentLen: info.ContentLength,
Score: roundFloat(sf.score, 2),
})
if len(results) >= maxResults {
break
}
}
return results, nil
}
// ========== 分词 ==========
func tokenizeQuery(query string) map[string]struct{} {
terms := make(map[string]struct{})
// 中文连续 2-4 字
for _, chunk := range findChineseChunks(query) {
runes := []rune(chunk)
for i := 0; i < len(runes); i++ {
for length := 2; length <= 4; length++ {
if i+length <= len(runes) {
terms[string(runes[i:i+length])] = struct{}{}
}
}
}
}
// 英文词
for _, word := range regexp.MustCompile(`[a-zA-Z0-9]{2,}`).FindAllString(strings.ToLower(query), -1) {
terms[word] = struct{}{}
}
// 完整查询词
terms[strings.ToLower(strings.TrimSpace(query))] = struct{}{}
return terms
}
func findChineseChunks(text string) []string {
var chunks []string
var current []rune
for _, r := range text {
if unicode.Is(unicode.Han, r) {
current = append(current, r)
} else {
if len(current) > 0 {
chunks = append(chunks, string(current))
current = nil
}
}
}
if len(current) > 0 {
chunks = append(chunks, string(current))
}
return chunks
}
// ========== 上下文格式化 ==========
// maxContextTokens 默认最大上下文 token 数
const maxContextTokens = 8000
// formatWikiContext 将搜索结果格式化为 AI 聊天的参考上下文
func formatWikiContext(results []wikiResult, maxTokens int) string {
if len(results) == 0 {
return ""
}
if maxTokens <= 0 {
maxTokens = maxContextTokens
}
var parts []string
parts = append(parts, "# 饥荒 Wiki 参考文档")
totalChars := 0
maxChars := maxTokens * 2 // 粗略按字符数/2 估算 token
for i, r := range results {
header := formatResultHeader(i+1, r)
body := r.Content
entryChars := utf8.RuneCountInString(header) + utf8.RuneCountInString(body) + 10
if totalChars+entryChars > maxChars {
remaining := maxChars - totalChars - utf8.RuneCountInString(header) - 50
if remaining > 200 {
runes := []rune(body)
if len(runes) > remaining {
body = string(runes[:remaining]) + "\n\n(内容已截断...)"
}
} else {
break
}
}
parts = append(parts, "")
parts = append(parts, header)
parts = append(parts, "")
parts = append(parts, body)
totalChars += entryChars
}
return strings.Join(parts, "\n")
}
func formatResultHeader(index int, r wikiResult) string {
h := "## " + itoa(index) + ". " + r.Title
if len(r.Categories) > 0 {
h += " [" + strings.Join(r.Categories, ", ") + "]"
}
return h
}
func itoa(n int) string {
if n <= 0 {
return "0"
}
digits := make([]byte, 0, 10)
for n > 0 {
digits = append([]byte{byte('0' + n%10)}, digits...)
n /= 10
}
return string(digits)
}
func roundFloat(f float64, decimals int) float64 {
pow := 1.0
for i := 0; i < decimals; i++ {
pow *= 10
}
return float64(int(f*pow+0.5)) / pow
}
// ================================================================
// 构建关键词搜索索引
// ================================================================
// BuildIndex 构建关键词搜索索引。扫描所有 .md 文件,构建 TF 评分索引。
// force 为 true 时强制重建,否则索引已存在则跳过。
func (s *keywordWikiSearcher) buildIndex(force bool) error {
s.mu.Lock()
defer s.mu.Unlock()
// 非强制模式下,如果索引文件已存在,直接加载
if !force {
if _, err := os.Stat(s.indexPath); err == nil {
logger.Logger.Infof("关键词索引已存在,加载中...")
return s.loadLocked()
}
}
logger.Logger.Infof("正在构建关键词搜索索引...")
start := time.Now()
idx := &wikiIndex{
Pages: make(map[string]wikiPageInfo),
CategoryIndex: make(map[string][]string),
TermIndex: make(map[string]map[string]float64),
}
entries, err := os.ReadDir(s.pagesDir)
if err != nil {
return fmt.Errorf("扫描 Wiki 页面目录失败: %w", err)
}
var mdFiles []string
for _, entry := range entries {
if !entry.IsDir() && filepath.Ext(entry.Name()) == ".md" {
mdFiles = append(mdFiles, filepath.Join(s.pagesDir, entry.Name()))
}
}
sort.Strings(mdFiles)
total := len(mdFiles)
for i, fp := range mdFiles {
if i%200 == 0 {
logger.Logger.Infof("处理中... %d/%d", i, total)
}
filename := filepath.Base(fp)
content, readErr := os.ReadFile(fp)
if readErr != nil {
logger.Logger.Warnf("读取 %s 失败: %v", filename, readErr)
continue
}
rawContent := string(content)
// 提取标题
title := extractWikiTitle(rawContent, filename)
// 提取分类
categories := extractWikiCategories(rawContent)
// 提取正文
body := stripWikiMeta(rawContent)
// 提取内部链接
links := extractWikiLinks(body)
// 分词
textForSearch := title + " " + body
textForSearch = cleanForTokenize(textForSearch)
allTerms := tokenizeText(textForSearch)
// 记录页面信息
idx.Pages[filename] = wikiPageInfo{
Title: title,
Categories: categories,
Content: body,
ContentLength: utf8.RuneCountInString(body),
Links: links,
Filename: filename,
Terms: sortedTermSlice(allTerms),
}
// 分类索引
for _, cat := range categories {
idx.CategoryIndex[cat] = append(idx.CategoryIndex[cat], filename)
}
// 词索引 (TF scoring: 词在文档中出现越多,分越高)
textLower := strings.ToLower(textForSearch)
termCounts := make(map[string]int)
for term := range allTerms {
termCounts[term] = strings.Count(textLower, strings.ToLower(term))
}
maxCount := 1
for _, count := range termCounts {
if count > maxCount {
maxCount = count
}
}
for term, count := range termCounts {
// 归一化 TF + 标题中出现加权
score := float64(count) / float64(maxCount)
if strings.Contains(strings.ToLower(title), strings.ToLower(term)) {
score *= 2.0
}
if idx.TermIndex[term] == nil {
idx.TermIndex[term] = make(map[string]float64)
}
idx.TermIndex[term][filename] = score
}
}
// 保存索引
if err := saveKeywordIndex(s.indexPath, idx); err != nil {
return err
}
s.index = idx
s.resetIdleTimer()
elapsed := time.Since(start)
logger.Logger.Infof("索引构建完成! %d 个页面, %d 个词条, 耗时 %.1fs", total, len(idx.TermIndex), elapsed.Seconds())
return nil
}
// ========== 索引构建辅助函数 ==========
// tokenizeText 对文本全文分词(用于构建索引,不含完整查询词)
func tokenizeText(text string) map[string]struct{} {
terms := make(map[string]struct{})
// 中文连续 2-4 字
for _, chunk := range findChineseChunks(text) {
runes := []rune(chunk)
for j := 0; j < len(runes); j++ {
for length := 2; length <= 4; length++ {
if j+length <= len(runes) {
terms[string(runes[j:j+length])] = struct{}{}
}
}
}
}
// 英文词
for _, word := range regexp.MustCompile(`[a-zA-Z0-9]{2,}`).FindAllString(strings.ToLower(text), -1) {
terms[word] = struct{}{}
}
return terms
}
// cleanForTokenize 去除 markdown 格式字符,压缩空白
func cleanForTokenize(text string) string {
text = regexp.MustCompile(`[*#>`+"`"+`\[\]()!_~|]`).ReplaceAllString(text, " ")
text = regexp.MustCompile(`\s+`).ReplaceAllString(text, " ")
return text
}
// extractWikiLinks 提取 Wiki 内部链接的显示文本
var wikiLinkRe = regexp.MustCompile(`\[([^\]]+)\]\(([^)]+\.md)\)`)
func extractWikiLinks(body string) []string {
matches := wikiLinkRe.FindAllStringSubmatch(body, -1)
links := make([]string, 0, len(matches))
for _, m := range matches {
if len(m) >= 2 {
links = append(links, m[1])
}
}
return links
}
// sortedTermSlice 将 term set 转为排序的 slice
func sortedTermSlice(terms map[string]struct{}) []string {
result := make([]string, 0, len(terms))
for t := range terms {
result = append(result, t)
}
sort.Strings(result)
return result
}
// saveKeywordIndex 保存搜索索引到磁盘
func saveKeywordIndex(indexPath string, idx *wikiIndex) error {
dir := filepath.Dir(indexPath)
if err := os.MkdirAll(dir, 0755); err != nil {
return fmt.Errorf("创建索引目录失败: %w", err)
}
data, err := json.Marshal(idx)
if err != nil {
return fmt.Errorf("序列化搜索索引失败: %w", err)
}
if err := os.WriteFile(indexPath, data, 0644); err != nil {
return fmt.Errorf("保存搜索索引失败: %w", err)
}
return nil
}

文章分享

如果这篇文章对你有帮助,欢迎分享给更多人!

DMP游戏内AI聊天的想法
https://blog.miraclesses.top/posts/rubbish/dmp_ai_chat/
作者
Miracle
发布于
2026-07-20
许可协议
CC BY-NC-SA 4.0

评论区

Profile Image of the Author
Miracle
不倒点垃圾再走吗
最新动态
分类
标签
站点统计
文章
4
分类
2
标签
8
总字数
2,275
运行时长
0
最后活动
0 天前

当前页面没有目录