c22669c31d
- 新增迁移 000031-000034(登录日志/角色/模型配额/记忆) - 新增包: license/memory/modelquota/assistant,扫描引擎 - 全部功能后端+前端+端到端验证通过(25 包单测)
294 lines
9.9 KiB
Go
294 lines
9.9 KiB
Go
// Package memory 实现多层记忆管理:用户个人/部门/全局记忆集合,
|
|
// 向量化语义召回(复用 Ollama embedding),支持向用户授权与衰减清理。
|
|
package memory
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"strconv"
|
|
"fmt"
|
|
"strings"
|
|
"time"
|
|
|
|
platformid "aigateway.local/core/internal/platform/id"
|
|
"github.com/jackc/pgx/v5/pgxpool"
|
|
)
|
|
|
|
var (
|
|
ErrNotFound = errors.New("记忆不存在")
|
|
ErrUnavailable = errors.New("记忆服务不可用")
|
|
)
|
|
|
|
// OwnerKind 记忆归属层级。
|
|
type OwnerKind string
|
|
|
|
const (
|
|
OwnerUser OwnerKind = "user"
|
|
OwnerDepartment OwnerKind = "department"
|
|
OwnerGlobal OwnerKind = "global"
|
|
)
|
|
|
|
// Entry 是一条记忆。
|
|
type Entry struct {
|
|
ID string `json:"id"`
|
|
OwnerKind OwnerKind `json:"owner_kind"`
|
|
OwnerID string `json:"owner_id"`
|
|
Category string `json:"category"`
|
|
Content string `json:"content"`
|
|
Importance int `json:"importance"`
|
|
SharedWith []string `json:"shared_with"`
|
|
Source string `json:"source"`
|
|
LastAccessedAt *time.Time `json:"last_accessed_at,omitempty"`
|
|
CreatedBy *string `json:"created_by,omitempty"`
|
|
CreatedAt time.Time `json:"created_at"`
|
|
UpdatedAt time.Time `json:"updated_at"`
|
|
}
|
|
|
|
// Embedder 生成文本向量(复用知识库的 Ollama embedder)。
|
|
type Embedder interface {
|
|
Embed(ctx context.Context, texts []string) ([][]float32, error)
|
|
}
|
|
|
|
// Service 记忆管理服务。
|
|
type Service struct {
|
|
pool *pgxpool.Pool
|
|
embedder Embedder
|
|
}
|
|
|
|
func NewService(pool *pgxpool.Pool, embedder Embedder) *Service {
|
|
return &Service{pool: pool, embedder: embedder}
|
|
}
|
|
|
|
func (s *Service) SetEmbedder(embedder Embedder) { s.embedder = embedder }
|
|
|
|
const entrySelect = `SELECT id::text,owner_kind,owner_id,category,content,importance,shared_with::text[],source,last_accessed_at,created_by::text,created_at,updated_at FROM gateway.memory_entries`
|
|
|
|
func (s *Service) scanEntry(row interface{ Scan(dest ...any) error }) (Entry, error) {
|
|
var e Entry
|
|
var shared []string
|
|
var createdBy *string
|
|
err := row.Scan(&e.ID, &e.OwnerKind, &e.OwnerID, &e.Category, &e.Content, &e.Importance, &shared, &e.Source, &e.LastAccessedAt, &createdBy, &e.CreatedAt, &e.UpdatedAt)
|
|
if err != nil {
|
|
if strings.Contains(err.Error(), "no rows") {
|
|
return Entry{}, ErrNotFound
|
|
}
|
|
return Entry{}, err
|
|
}
|
|
e.SharedWith = shared
|
|
e.CreatedBy = createdBy
|
|
return e, nil
|
|
}
|
|
|
|
// Save 创建或更新一条记忆。ownerID 为空时按 kind 处理(global 无归属)。
|
|
func (s *Service) Save(ctx context.Context, kind OwnerKind, ownerID, id, category, content, source string, importance int, sharedWith []string, actorID string) (Entry, error) {
|
|
if s == nil || s.pool == nil {
|
|
return Entry{}, ErrUnavailable
|
|
}
|
|
content = strings.TrimSpace(content)
|
|
category = strings.TrimSpace(category)
|
|
if content == "" || len(content) > 8000 {
|
|
return Entry{}, errors.New("记忆内容必须为 1-8000 字符")
|
|
}
|
|
if category == "" {
|
|
category = "general"
|
|
}
|
|
if len(category) > 64 || len(source) > 128 {
|
|
return Entry{}, errors.New("分类或来源过长")
|
|
}
|
|
if importance < 1 {
|
|
importance = 5
|
|
}
|
|
if importance > 10 {
|
|
importance = 10
|
|
}
|
|
if sharedWith == nil {
|
|
sharedWith = []string{} // NOT NULL 列:空授权为显式空数组
|
|
}
|
|
var err error
|
|
var embedding string
|
|
if s.embedder != nil {
|
|
vectors, err := s.embedder.Embed(ctx, []string{content})
|
|
if err == nil && len(vectors) == 1 && len(vectors[0]) == 1024 {
|
|
// pgx 不识别 vector 类型的二进制编码:与知识库一致,用文本格式
|
|
// "[0.1,0.2,...]" 配合 ::vector 转换。
|
|
embedding = "[" + strings.Trim(strings.Join(joinFloats(vectors[0]), ","), " ") + "]"
|
|
}
|
|
}
|
|
// 向量维度必须为 1024 与列匹配。
|
|
validVector := embedding != "" && strings.HasPrefix(embedding, "[")
|
|
|
|
if id == "" {
|
|
var newID string
|
|
newID, err = platformid.NewUUID()
|
|
if err != nil {
|
|
return Entry{}, err
|
|
}
|
|
id = newID
|
|
if validVector {
|
|
_, err = s.pool.Exec(ctx, `INSERT INTO gateway.memory_entries(id,owner_kind,owner_id,category,content,importance,embedding,shared_with,source,created_by) VALUES($1,$2,$3,$4,$5,$6,$7::vector,$8,$9,$10)`, id, kind, ownerID, category, content, importance, embedding, sharedWith, source, actorID)
|
|
} else {
|
|
_, err = s.pool.Exec(ctx, `INSERT INTO gateway.memory_entries(id,owner_kind,owner_id,category,content,importance,shared_with,source,created_by) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9)`, id, kind, ownerID, category, content, importance, sharedWith, source, actorID)
|
|
}
|
|
if err != nil {
|
|
return Entry{}, err
|
|
}
|
|
} else {
|
|
var tag interface{ RowsAffected() int64 }
|
|
if validVector {
|
|
tag, err = s.pool.Exec(ctx, `UPDATE gateway.memory_entries SET category=$3,content=$4,importance=$5,embedding=$6::vector,shared_with=$7,source=$8,updated_at=clock_timestamp() WHERE id=$1 AND owner_kind=$2`, id, kind, category, content, importance, embedding, sharedWith, source)
|
|
} else {
|
|
tag, err = s.pool.Exec(ctx, `UPDATE gateway.memory_entries SET category=$3,content=$4,importance=$5,shared_with=$6,source=$7,updated_at=clock_timestamp() WHERE id=$1 AND owner_kind=$2`, id, kind, category, content, importance, sharedWith, source)
|
|
}
|
|
if err != nil {
|
|
return Entry{}, err
|
|
}
|
|
if tag.RowsAffected() == 0 {
|
|
return Entry{}, ErrNotFound
|
|
}
|
|
}
|
|
return s.Get(ctx, id)
|
|
}
|
|
|
|
func (s *Service) Get(ctx context.Context, id string) (Entry, error) {
|
|
if s == nil || s.pool == nil {
|
|
return Entry{}, ErrUnavailable
|
|
}
|
|
return s.scanEntry(s.pool.QueryRow(ctx, entrySelect+` WHERE id=$1`, id))
|
|
}
|
|
|
|
// List 列出归属下的记忆(global 与 department 可见性由调用方合并)。
|
|
func (s *Service) List(ctx context.Context, kind OwnerKind, ownerID string, limit int) ([]Entry, error) {
|
|
if s == nil || s.pool == nil {
|
|
return nil, ErrUnavailable
|
|
}
|
|
if limit < 1 || limit > 500 {
|
|
limit = 100
|
|
}
|
|
rows, err := s.pool.Query(ctx, entrySelect+` WHERE owner_kind=$1 AND owner_id=$2 ORDER BY importance DESC,updated_at DESC LIMIT $3`, kind, ownerID, limit)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
items := []Entry{}
|
|
for rows.Next() {
|
|
e, err := s.scanEntry(rows)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
items = append(items, e)
|
|
}
|
|
return items, rows.Err()
|
|
}
|
|
|
|
// Delete 删除记忆(global 允许任意管理员)。
|
|
func (s *Service) Delete(ctx context.Context, id string) error {
|
|
if s == nil || s.pool == nil {
|
|
return ErrUnavailable
|
|
}
|
|
tag, err := s.pool.Exec(ctx, `DELETE FROM gateway.memory_entries WHERE id=$1`, id)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if tag.RowsAffected() == 0 {
|
|
return ErrNotFound
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// Recall 语义召回:按向量相似度返回与 query 最相关的记忆。
|
|
// scopes 决定搜索范围(user 自己的 + shared_with 含用户的 + department + global)。
|
|
func (s *Service) Recall(ctx context.Context, userID string, departmentID *string, query string, limit int) ([]Entry, error) {
|
|
if s == nil || s.pool == nil {
|
|
return nil, ErrUnavailable
|
|
}
|
|
if limit < 1 || limit > 20 {
|
|
limit = 5
|
|
}
|
|
query = strings.TrimSpace(query)
|
|
if query == "" {
|
|
return nil, errors.New("查询内容不能为空")
|
|
}
|
|
var embedding string
|
|
if s.embedder != nil {
|
|
vectors, err := s.embedder.Embed(ctx, []string{query})
|
|
if err == nil && len(vectors) == 1 && len(vectors[0]) == 1024 {
|
|
embedding = "[" + strings.Trim(strings.Join(joinFloats(vectors[0]), ","), " ") + "]"
|
|
}
|
|
}
|
|
if embedding == "" {
|
|
// 向量不可用(embedding 关闭/Ollama 故障):按关键字+重要度召回。
|
|
return s.recallFallback(ctx, userID, departmentID, query, limit)
|
|
}
|
|
scope := `(owner_kind='global' OR (owner_kind='user' AND owner_id=$1) OR (owner_kind='department' AND owner_id=$2) OR $1::uuid = ANY(shared_with))`
|
|
rows, err := s.pool.Query(ctx, entrySelect+` WHERE `+scope+` AND embedding IS NOT NULL ORDER BY embedding <=> $3::vector LIMIT $4`, userID, departmentID, embedding, limit)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
items := []Entry{}
|
|
for rows.Next() {
|
|
e, err := s.scanEntry(rows)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
items = append(items, e)
|
|
}
|
|
// 记录访问时间(衰减依据)。
|
|
if len(items) > 0 {
|
|
ids := make([]string, 0, len(items))
|
|
for _, e := range items {
|
|
ids = append(ids, e.ID)
|
|
}
|
|
_, _ = s.pool.Exec(ctx, `UPDATE gateway.memory_entries SET last_accessed_at=clock_timestamp() WHERE id = ANY($1::uuid[])`, ids)
|
|
}
|
|
return items, rows.Err()
|
|
}
|
|
|
|
func (s *Service) recallFallback(ctx context.Context, userID string, departmentID *string, query string, limit int) ([]Entry, error) {
|
|
scope := `(owner_kind='global' OR (owner_kind='user' AND owner_id=$1) OR (owner_kind='department' AND owner_id=$2) OR $1::uuid = ANY(shared_with))`
|
|
rows, err := s.pool.Query(ctx, entrySelect+` WHERE `+scope+` AND (content ILIKE '%'||$3||'%' OR to_tsvector('simple', content) @@ plainto_tsquery('simple', $3)) ORDER BY importance DESC,updated_at DESC LIMIT $4`, userID, departmentID, query, limit)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
items := []Entry{}
|
|
for rows.Next() {
|
|
e, err := s.scanEntry(rows)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
items = append(items, e)
|
|
}
|
|
return items, rows.Err()
|
|
}
|
|
|
|
// joinFloats 把 float32 切片格式化为 pgvector 文本。
|
|
func joinFloats(values []float32) []string {
|
|
out := make([]string, len(values))
|
|
for i, v := range values {
|
|
out[i] = strconv.FormatFloat(float64(v), 'f', -1, 32)
|
|
}
|
|
return out
|
|
}
|
|
|
|
// Decay 衰减清理:低重要度且长期未访问的记忆降权并最终删除
|
|
// (由 maintenance worker 定期调用)。
|
|
func (s *Service) Decay(ctx context.Context, now time.Time, inactiveDays int) (int64, error) {
|
|
if s == nil || s.pool == nil {
|
|
return 0, ErrUnavailable
|
|
}
|
|
if inactiveDays < 7 {
|
|
inactiveDays = 30
|
|
}
|
|
tag, err := s.pool.Exec(ctx, `DELETE FROM gateway.memory_entries
|
|
WHERE importance <= 3 AND (last_accessed_at IS NULL OR last_accessed_at < $1::timestamptz)`,
|
|
now.AddDate(0, 0, -inactiveDays))
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
return tag.RowsAffected(), nil
|
|
}
|
|
|
|
// String 便捷格式化。
|
|
func (s *Service) String() string { return fmt.Sprintf("memory-service") }
|