diff --git a/CLAUDE.md b/CLAUDE.md index 7246938..d07dc8f 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -62,6 +62,7 @@ ## 数据访问规范(硬性要求) - **事务**:涉及多张表的增删改操作必须包数据库事务,禁止逐表裸调用。事务放 dao 层方法内,service 层负责编排;`tx.Begin` 后必须用 `defer` 防护已提交后的二次 Rollback +- **SQL 单表约束**:每个 SQL 只允许访问一张表,禁止 JOIN 与跨表子查询(`IN (SELECT ...)` / `EXISTS`);跨表数据一律拆为多条单表 SQL + 应用层内存组装——先取外键 id 列表,再对目标表 `IN` 查询;`IN` 参数须按 ≤100 分批(SQLite 变量数上限 999) - **禁止 N+1 查询**:禁止在循环中逐条查库。循环场景一律改为批处理——一次 `ListByXxx` 取回后按外键在内存分组 - **缓存一致性**:DAO 查询走缓存(TTL 来自 `database.cache.ttl`),写操作后必须清对应缓存 - **批处理 SQL**:批量写入用 `InsertAll` 类方法,批量删除用 `IN` 子句,禁止循环单条 INSERT/DELETE diff --git a/data/business.db b/data/business.db index 77a9571..dd5b700 100644 Binary files a/data/business.db and b/data/business.db differ diff --git a/kb/consts/consts.go b/kb/consts/consts.go index 5eb092e..b77b662 100644 --- a/kb/consts/consts.go +++ b/kb/consts/consts.go @@ -38,6 +38,11 @@ const ( AnnoMaxCandidates = 60 // 多数据集融合后的候选上限(喂给 LLM 判定) AnnoMaxClauseChars = 2000 // 合同条款全文上限(超长截断,控制 prompt) + // 判定 prompt 上下文预算(本地 4B 模型 context 8192 token,须留生成余量) + AnnoJudgeMaxTokens = 1024 // 判定输出 token 上限(JSON 结果含 ≤3 条风险;须显式限制,防推理模型长思考烧光上下文) + AnnoJudgeCandidateChars = 300 // 判定 prompt 单候选条文截断字数(精确条文文本由 ContentFull 抽取,不依赖 prompt 全文) + AnnoJudgePromptBudget = 1500 // 判定 prompt 候选块总字数预算(候选按 RRF 相关度降序贪心填充,超预算截断后续候选) + // 合同风险识别 RiskLevelHigh = "high" // 高风险(违反强制性规定、可能导致合同无效/赔偿) RiskLevelMid = "mid" // 中风险(约定与法律不符但可补救) diff --git a/kb/controller/document_controller.go b/kb/controller/document_controller.go index ea99697..f683e8d 100644 --- a/kb/controller/document_controller.go +++ b/kb/controller/document_controller.go @@ -3,6 +3,7 @@ package controller import ( "context" + "rag-local/kb/consts" "rag-local/kb/model/dto" "rag-local/kb/service" ) @@ -48,9 +49,9 @@ func (c *document) Delete(ctx context.Context, req *dto.DeleteDocumentReq) (*dto } func (c *document) Reembed(ctx context.Context, req *dto.ReembedDocumentReq) (*dto.ReembedDocumentRes, error) { - n, err := service.DocumentService.Reembed(ctx, req.Id) - if err != nil { + // 入队由轮询器串行消费,避免同步重算阻塞 HTTP 且无法展示过程状态;已排队/处理中会跳过 + if err := service.ParseTaskService.Enqueue(ctx, req.Id, consts.TaskTypeReembed); err != nil { return nil, err } - return &dto.ReembedDocumentRes{Count: n}, nil + return &dto.ReembedDocumentRes{}, nil } diff --git a/kb/controller/kg_entity_controller.go b/kb/controller/kg_entity_controller.go index cfabd9b..e4a7be2 100644 --- a/kb/controller/kg_entity_controller.go +++ b/kb/controller/kg_entity_controller.go @@ -23,3 +23,15 @@ func (c *kgEntity) List(ctx context.Context, req *dto.ListKgEntityReq) (*dto.Lis PageSize: req.PageSize, }, nil } + +func (c *kgEntity) Sources(ctx context.Context, req *dto.SourcesKgEntityReq) (*dto.SourcesKgEntityRes, error) { + list, err := service.KgEntityService.Sources(ctx, req.DatasetId) + if err != nil { + return nil, err + } + out := make([]*dto.KgEntitySourceItem, 0, len(list)) + for _, s := range list { + out = append(out, &dto.KgEntitySourceItem{Name: s.Name, Files: s.Files}) + } + return &dto.SourcesKgEntityRes{List: out}, nil +} diff --git a/kb/dao/chunk_dao.go b/kb/dao/chunk_dao.go index 425e1e9..3fb43af 100644 --- a/kb/dao/chunk_dao.go +++ b/kb/dao/chunk_dao.go @@ -222,25 +222,86 @@ func (d *chunkDao) CountByDocument(ctx context.Context, documentId int64) (int, return g.DB(consts.DbGroupDefault).Model(consts.TableNameChunk).Ctx(ctx).Where("document_id", documentId).Count() } -// VecSearch 向量 KNN 检索:vec0 取最近 topK*4 后按数据集过滤(vec0 无 dataset 列) +// ListByIds 批量按主键查分块(单表约束:IN ≤100 分批;调用方保证 ids 非空) +func (d *chunkDao) ListByIds(ctx context.Context, ids []int64) ([]*entity.Chunk, error) { + var list []*entity.Chunk + for start := 0; start < len(ids); start += 100 { + end := min(start+100, len(ids)) + var part []*entity.Chunk + if err := g.DB(consts.DbGroupDefault).Model(consts.TableNameChunk).Ctx(ctx). + WhereIn("id", ids[start:end]).Scan(&part); err != nil { + return nil, err + } + list = append(list, part...) + } + if list == nil { + list = make([]*entity.Chunk, 0) + } + return list, nil +} + +// ListDocumentIdsByChunkIds 批量查 chunk 归属文档(单表约束:IN ≤100 分批;调用方保证 chunkIds 非空) +func (d *chunkDao) ListDocumentIdsByChunkIds(ctx context.Context, chunkIds []int64) (map[int64]int64, error) { + out := make(map[int64]int64, len(chunkIds)) + for start := 0; start < len(chunkIds); start += 100 { + end := min(start+100, len(chunkIds)) + var rows []*entity.Chunk + if err := g.DB(consts.DbGroupDefault).Model(consts.TableNameChunk).Ctx(ctx). + Fields("id, document_id").WhereIn("id", chunkIds[start:end]).Scan(&rows); err != nil { + return nil, err + } + for _, r := range rows { + out[r.Id] = r.DocumentId + } + } + return out, nil +} + +// VecSearch 向量 KNN 检索:vec0 取最近 topK*4 候选,再按数据集过滤(单表约束:vec0 无 dataset 列, +// 拆两条单表查询 + 内存过滤;候选已按距离升序,过滤后取前 topK 即等价原 JOIN 语义) func (d *chunkDao) VecSearch(ctx context.Context, datasetId int64, vecJson string, topK int) ([]domain.VecHit, error) { r, err := g.DB(consts.DbGroupDefault).Ctx(ctx).Raw( - `SELECT v.chunk_id, v.distance FROM ( - SELECT chunk_id, distance FROM `+consts.TableNameChunkVec+` - WHERE embedding MATCH ? ORDER BY distance LIMIT ? - ) v INNER JOIN `+consts.TableNameChunk+` c ON c.id = v.chunk_id - WHERE c.dataset_id = ? ORDER BY v.distance LIMIT ?`, - vecJson, topK*4, datasetId, topK, + `SELECT chunk_id, distance FROM `+consts.TableNameChunkVec+ + ` WHERE embedding MATCH ? ORDER BY distance LIMIT ?`, + vecJson, topK*4, ).All() if err != nil { return nil, err } - hits := make([]domain.VecHit, 0, len(r)) + if len(r) == 0 { + return nil, nil + } + cand := make([]domain.VecHit, 0, len(r)) + ids := make([]int64, 0, len(r)) for _, row := range r { - hits = append(hits, domain.VecHit{ - ChunkId: row["chunk_id"].Int64(), - Distance: row["distance"].Float64(), - }) + id := row["chunk_id"].Int64() + cand = append(cand, domain.VecHit{ChunkId: id, Distance: row["distance"].Float64()}) + ids = append(ids, id) + } + // 候选 id ≤ topK*4(默认 80),远低于 SQLite 变量上限,无需分批 + placeholders := strings.TrimSuffix(strings.Repeat("?,", len(ids)), ",") + args := make([]any, 0, len(ids)+1) + for _, id := range ids { + args = append(args, id) + } + args = append(args, datasetId) + chunks, err := g.DB(consts.DbGroupDefault).Ctx(ctx).Raw( + "SELECT id FROM "+consts.TableNameChunk+" WHERE id IN ("+placeholders+") AND dataset_id = ?", args...).All() + if err != nil { + return nil, err + } + allowed := make(map[int64]bool, len(chunks)) + for _, row := range chunks { + allowed[row["id"].Int64()] = true + } + hits := make([]domain.VecHit, 0, len(cand)) + for _, h := range cand { + if allowed[h.ChunkId] { + hits = append(hits, h) + if len(hits) >= topK { + break + } + } } return hits, nil } diff --git a/kb/dao/contract_clause_dao.go b/kb/dao/contract_clause_dao.go index 3ce5936..3362943 100644 --- a/kb/dao/contract_clause_dao.go +++ b/kb/dao/contract_clause_dao.go @@ -36,31 +36,51 @@ func init() { } func (d *contractClauseDao) InsertAll(ctx context.Context, taskId int64, clauses []entity.ContractClause) error { + if len(clauses) == 0 { + return nil + } + now := gtime.Now().Format("Y-m-d H:i:s") + list := make(g.List, 0, len(clauses)) + for i := range clauses { + list = append(list, g.Map{ + "task_id": taskId, + "seq": clauses[i].Seq, + "title": clauses[i].Title, + "content": clauses[i].Content, + "status": consts.TaskStatusPending, + "error_msg": "", + "created_at": now, + "updated_at": now, + }) + } tx, err := g.DB(consts.DbGroupDefault).Begin(ctx) if err != nil { return err } - now := gtime.Now().Format("Y-m-d H:i:s") - for i := range clauses { - clauses[i].TaskId = taskId - clauses[i].Status = consts.TaskStatusPending - clauses[i].CreatedAt = nil - clauses[i].UpdatedAt = nil - if _, err := tx.Model(consts.TableNameContractClause).Ctx(ctx).Data(g.Map{ - "task_id": clauses[i].TaskId, - "seq": clauses[i].Seq, - "title": clauses[i].Title, - "content": clauses[i].Content, - "status": clauses[i].Status, - "error_msg": "", - "created_at": now, - "updated_at": now, - }).Insert(); err != nil { + // Commit 成功后 IsClosed 为 true,跳过 Rollback,避免对已提交事务回滚产生报错日志 + defer func() { + if !tx.IsClosed() { _ = tx.Rollback() + } + }() + if _, err := tx.Model(consts.TableNameContractClause).Ctx(ctx).Data(list).Batch(100).Insert(); err != nil { + return err + } + return tx.Commit() +} + +// UpdateStatuses 批量更新条款状态(单表约束:IN ≤100 分批) +func (d *contractClauseDao) UpdateStatuses(ctx context.Context, ids []int64, status int, errorMsg string) error { + now := gtime.Now().Format("Y-m-d H:i:s") + for start := 0; start < len(ids); start += 100 { + end := min(start+100, len(ids)) + if _, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractClause).Ctx(ctx). + Data(g.Map{"status": status, "error_msg": errorMsg, "updated_at": now}). + WhereIn("id", ids[start:end]).Update(); err != nil { return err } } - return tx.Commit() + return nil } func (d *contractClauseDao) ListByTask(ctx context.Context, taskId int64) ([]*entity.ContractClause, error) { diff --git a/kb/dao/contract_mark_dao.go b/kb/dao/contract_mark_dao.go index b827c6b..6a61e67 100644 --- a/kb/dao/contract_mark_dao.go +++ b/kb/dao/contract_mark_dao.go @@ -2,12 +2,12 @@ package dao import ( "context" + "sort" "rag-local/kb/consts" "rag-local/kb/model/entity" "github.com/gogf/gf/v2/frame/g" - "github.com/gogf/gf/v2/os/gtime" ) var ContractMark = &contractMarkDao{} @@ -36,34 +36,6 @@ func init() { } } -func (d *contractMarkDao) InsertAll(ctx context.Context, marks []*entity.ContractMark) error { - if len(marks) == 0 { - return nil - } - tx, err := g.DB(consts.DbGroupDefault).Begin(ctx) - if err != nil { - return err - } - now := gtime.Now().Format("Y-m-d H:i:s") - for _, m := range marks { - if _, err := tx.Model(consts.TableNameContractMark).Ctx(ctx).Data(g.Map{ - "clause_id": m.ClauseId, - "chunk_id": m.ChunkId, - "dataset_id": m.DatasetId, - "law_title": m.LawTitle, - "law_item": m.LawItem, - "content": m.Content, - "reason": m.Reason, - "score": m.Score, - "created_at": now, - }).Insert(); err != nil { - _ = tx.Rollback() - return err - } - } - return tx.Commit() -} - func (d *contractMarkDao) ListByClause(ctx context.Context, clauseId int64) ([]*entity.ContractMark, error) { var list []*entity.ContractMark err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractMark).Ctx(ctx). @@ -74,15 +46,32 @@ func (d *contractMarkDao) ListByClause(ctx context.Context, clauseId int64) ([]* return list, err } +// ListByTask 某任务的全部标注(单表约束:先取条款 id,再按 IN 分批查询,内存按分排序) func (d *contractMarkDao) ListByTask(ctx context.Context, taskId int64) ([]*entity.ContractMark, error) { + clauseIds, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractClause).Ctx(ctx). + Fields("id").Where("task_id", taskId).Array() + if err != nil { + return nil, err + } var list []*entity.ContractMark - err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractMark).Ctx(ctx). - Where("clause_id IN (SELECT id FROM "+consts.TableNameContractClause+" WHERE task_id = ?)", taskId). - OrderDesc("score").Scan(&list) + for start := 0; start < len(clauseIds); start += 100 { + end := min(start+100, len(clauseIds)) + ids := make([]int64, 0, end-start) + for _, v := range clauseIds[start:end] { + ids = append(ids, v.Int64()) + } + var part []*entity.ContractMark + if err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractMark).Ctx(ctx). + WhereIn("clause_id", ids).Scan(&part); err != nil { + return nil, err + } + list = append(list, part...) + } if list == nil { list = make([]*entity.ContractMark, 0) } - return list, err + sort.Slice(list, func(i, j int) bool { return list[i].Score > list[j].Score }) + return list, nil } func (d *contractMarkDao) DeleteByClause(ctx context.Context, clauseId int64) error { diff --git a/kb/dao/contract_risk_dao.go b/kb/dao/contract_risk_dao.go index 2d37f62..e88c1d0 100644 --- a/kb/dao/contract_risk_dao.go +++ b/kb/dao/contract_risk_dao.go @@ -40,23 +40,30 @@ func (d *contractRiskDao) InsertAll(ctx context.Context, risks []*entity.Contrac if len(risks) == 0 { return nil } - tx, err := g.DB(consts.DbGroupDefault).Begin(ctx) - if err != nil { - return err - } now := gtime.Now().Format("Y-m-d H:i:s") + list := make(g.List, 0, len(risks)) for _, r := range risks { - if _, err := tx.Model(consts.TableNameContractRisk).Ctx(ctx).Data(g.Map{ + list = append(list, g.Map{ "task_id": r.TaskId, "clause_id": r.ClauseId, "level": r.Level, "desc": r.Desc, "laws": r.Laws, "created_at": now, - }).Insert(); err != nil { + }) + } + tx, err := g.DB(consts.DbGroupDefault).Begin(ctx) + if err != nil { + return err + } + // Commit 成功后 IsClosed 为 true,跳过 Rollback,避免对已提交事务回滚产生报错日志 + defer func() { + if !tx.IsClosed() { _ = tx.Rollback() - return err } + }() + if _, err := tx.Model(consts.TableNameContractRisk).Ctx(ctx).Data(list).Batch(100).Insert(); err != nil { + return err } return tx.Commit() } diff --git a/kb/dao/contract_task_dao.go b/kb/dao/contract_task_dao.go index 4607cd0..5862b9d 100644 --- a/kb/dao/contract_task_dao.go +++ b/kb/dao/contract_task_dao.go @@ -136,10 +136,23 @@ func (d *contractTaskDao) DeleteWithRelated(ctx context.Context, id int64) error _ = tx.Rollback() } }() + // 单表约束:先取条款 id,再按 IN 分批删除(条款数可能超 SQLite 变量上限) + clauseIds, err := tx.Model(consts.TableNameContractClause).Ctx(ctx).Fields("id").Where("task_id", id).Array() + if err != nil { + return err + } for _, table := range []string{consts.TableNameContractMark, consts.TableNameContractRisk} { - if _, err := tx.Exec( - "DELETE FROM "+table+" WHERE clause_id IN (SELECT id FROM "+consts.TableNameContractClause+" WHERE task_id = ?)", id); err != nil { - return err + for start := 0; start < len(clauseIds); start += 100 { + end := min(start+100, len(clauseIds)) + ids := make([]int64, 0, end-start) + for _, v := range clauseIds[start:end] { + ids = append(ids, v.Int64()) + } + if len(ids) > 0 { + if _, err := tx.Model(table).Ctx(ctx).WhereIn("clause_id", ids).Delete(); err != nil { + return err + } + } } } if _, err := tx.Model(consts.TableNameContractClause).Ctx(ctx).Where("task_id", id).Delete(); err != nil { diff --git a/kb/dao/document_dao.go b/kb/dao/document_dao.go index cd60d0c..c6fc44e 100644 --- a/kb/dao/document_dao.go +++ b/kb/dao/document_dao.go @@ -82,6 +82,24 @@ func (d *documentDao) List(ctx context.Context, datasetId int64, page, pageSize return list, total, err } +// ListByIds 批量按主键查文档(单表约束:IN ≤100 分批;调用方保证 ids 非空) +func (d *documentDao) ListByIds(ctx context.Context, ids []int64) ([]*entity.Document, error) { + var list []*entity.Document + for start := 0; start < len(ids); start += 100 { + end := min(start+100, len(ids)) + var part []*entity.Document + if err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDocument).Ctx(ctx). + FieldsEx("content").WhereIn("id", ids[start:end]).Scan(&part); err != nil { + return nil, err + } + list = append(list, part...) + } + if list == nil { + list = make([]*entity.Document, 0) + } + return list, nil +} + func (d *documentDao) Insert(ctx context.Context, data *entity.Document) (int64, error) { now := gtime.Now().Format("Y-m-d H:i:s") r, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDocument).Ctx(ctx).Data(g.Map{ diff --git a/kb/dao/kg_entity_dao.go b/kb/dao/kg_entity_dao.go index b68e692..c0909ca 100644 --- a/kb/dao/kg_entity_dao.go +++ b/kb/dao/kg_entity_dao.go @@ -124,6 +124,34 @@ func (d *kgEntityDao) UpsertBatch(ctx context.Context, datasetId, chunkId int64, return nil } +// HasGraphByDocument 判定某文档是否产出过图谱(单表约束:先取分块 id,再数关系)。 +// 不用实体行计数:实体按 (dataset_id,name) 去重、chunk_id 会被跨文件 upsert 覆盖,按实体计数会漏计。 +// 关系每 chunk 至少一条、不受覆盖影响。 +func (d *kgEntityDao) HasGraphByDocument(ctx context.Context, documentId int64) (bool, error) { + chunkIds, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameChunk).Ctx(ctx). + Fields("id").Where("document_id", documentId).Array() + if err != nil { + return false, err + } + // 分块数可能超 SQLite 变量上限,按 100 分批 + for start := 0; start < len(chunkIds); start += 100 { + end := min(start+100, len(chunkIds)) + ids := make([]int64, 0, end-start) + for _, v := range chunkIds[start:end] { + ids = append(ids, v.Int64()) + } + n, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameKgRelation).Ctx(ctx). + WhereIn("chunk_id", ids).Count() + if err != nil { + return false, err + } + if n > 0 { + return true, nil + } + } + return false, nil +} + func (d *kgEntityDao) DeleteByChunkIds(ctx context.Context, chunkIds []int64) error { if len(chunkIds) == 0 { return nil @@ -143,3 +171,79 @@ func (d *kgEntityDao) DeleteByDataset(ctx context.Context, datasetId int64) erro Where("dataset_id", datasetId).Delete() return err } + +// KgEntitySource 实体 → 出现文件清单(dao 内已按名聚合) +type KgEntitySource struct { + Name string + Files []string +} + +// SourcesByDataset 实体名 → 出现文件清单(单表约束:拆 4 条单表查询 + 内存组装)。 +// 有关系的实体经关系表溯源(关系不去重、每条带来源 chunk,覆盖实体全部出现文件); +// 孤立实体(无任何关系)由实体行 chunk_id 弱引用兜底(该弱引用是 upsert 的最近来源,不保证全部出现文件)。 +// 不再建实体-分块关联表:该组合已覆盖"图谱页标注来源文件"的展示需求。 +func (d *kgEntityDao) SourcesByDataset(ctx context.Context, datasetId int64) ([]KgEntitySource, error) { + // 1. 关系表(单表):数据集全部三元组的 head/tail 及其来源分块 + relRows, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameKgRelation).Ctx(ctx). + Fields("head", "tail", "chunk_id").Where("dataset_id", datasetId).All() + if err != nil { + return nil, err + } + // 2. 实体表(单表):孤立实体弱引用兜底 + entRows, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameKgEntity).Ctx(ctx). + Fields("name", "chunk_id").Where("dataset_id", datasetId).All() + if err != nil { + return nil, err + } + // 3. 分块表(单表):数据集全部分块 chunk_id → document_id + chunkRows, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameChunk).Ctx(ctx). + Fields("id", "document_id").Where("dataset_id", datasetId).All() + if err != nil { + return nil, err + } + // 4. 文档表(单表):数据集全部文档 id → 文件名 + docRows, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDocument).Ctx(ctx). + Fields("id", "filename").Where("dataset_id", datasetId).All() + if err != nil { + return nil, err + } + docByID := make(map[int64]string, len(docRows)) + for _, row := range docRows { + docByID[row["id"].Int64()] = row["filename"].String() + } + docByChunk := make(map[int64]int64, len(chunkRows)) + for _, row := range chunkRows { + docByChunk[row["id"].Int64()] = row["document_id"].Int64() + } + filesByName := make(map[string]map[string]bool) + add := func(name string, chunkId int64) { + docId, ok := docByChunk[chunkId] + if !ok { + return + } + fname := docByID[docId] + if name == "" || fname == "" { + return + } + if filesByName[name] == nil { + filesByName[name] = make(map[string]bool) + } + filesByName[name][fname] = true + } + for _, row := range relRows { + add(row["head"].String(), row["chunk_id"].Int64()) + add(row["tail"].String(), row["chunk_id"].Int64()) + } + for _, row := range entRows { + add(row["name"].String(), row["chunk_id"].Int64()) + } + out := make([]KgEntitySource, 0, len(filesByName)) + for name, set := range filesByName { + files := make([]string, 0, len(set)) + for f := range set { + files = append(files, f) + } + out = append(out, KgEntitySource{Name: name, Files: files}) + } + return out, nil +} diff --git a/kb/dao/parse_task_dao.go b/kb/dao/parse_task_dao.go index 3f2dbee..67cdbdd 100644 --- a/kb/dao/parse_task_dao.go +++ b/kb/dao/parse_task_dao.go @@ -99,6 +99,30 @@ func (d *parseTaskDao) DeleteByDocument(ctx context.Context, documentId int64) e return err } +// ResetRunning 启动恢复:进程异常退出会孤儿化 running 任务(轮询器只消费 pending),重置回 pending 让轮询器重新捡起 +func (d *parseTaskDao) ResetRunning(ctx context.Context) (int64, error) { + r, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameParseTask).Ctx(ctx). + Data(g.Map{ + "status": consts.TaskStatusPending, + "updated_at": gtime.Now().Format("Y-m-d H:i:s"), + }).Where("status", consts.TaskStatusRunning).Update() + if err != nil { + return 0, err + } + return r.RowsAffected() +} + +// HasPending 该文档是否存在待处理/运行中任务(重复入队防护) +func (d *parseTaskDao) HasPending(ctx context.Context, documentId int64) (bool, error) { + n, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameParseTask).Ctx(ctx). + Where("document_id", documentId). + WhereIn("status", []int{consts.TaskStatusPending, consts.TaskStatusRunning}).Count() + if err != nil { + return false, err + } + return n > 0, nil +} + func (d *parseTaskDao) NextPending(ctx context.Context) (*entity.ParseTask, error) { var m entity.ParseTask err := g.DB(consts.DbGroupDefault).Model(consts.TableNameParseTask).Ctx(ctx). diff --git a/kb/model/dto/document_dto.go b/kb/model/dto/document_dto.go index 732246c..a5671d4 100644 --- a/kb/model/dto/document_dto.go +++ b/kb/model/dto/document_dto.go @@ -52,6 +52,4 @@ type ReembedDocumentReq struct { Id int64 `v:"required" json:"id"` } -type ReembedDocumentRes struct { - Count int `json:"count"` -} +type ReembedDocumentRes struct{} diff --git a/kb/model/dto/kg_entity_dto.go b/kb/model/dto/kg_entity_dto.go index d75319d..7c93ce9 100644 --- a/kb/model/dto/kg_entity_dto.go +++ b/kb/model/dto/kg_entity_dto.go @@ -19,3 +19,17 @@ type ListKgEntityRes struct { Page int `json:"page"` PageSize int `json:"page_size"` } + +type SourcesKgEntityReq struct { + g.Meta `path:"/sources" method:"get" tags:"知识图谱" summary:"实体来源文件清单"` + DatasetId int64 `json:"dataset_id"` +} + +type KgEntitySourceItem struct { + Name string `json:"name"` + Files []string `json:"files"` +} + +type SourcesKgEntityRes struct { + List []*KgEntitySourceItem `json:"list"` +} diff --git a/kb/model/entity/contract_risk.go b/kb/model/entity/contract_risk.go index 3ff4c27..9ef8c8b 100644 --- a/kb/model/entity/contract_risk.go +++ b/kb/model/entity/contract_risk.go @@ -19,9 +19,10 @@ type ContractRisk struct { // LawRef 支撑法条引用(laws JSON 元素) type LawRef struct { - LawTitle string `json:"law_title"` - LawItem string `json:"law_item"` - Content string `json:"content"` + LawTitle string `json:"law_title"` + LawItem string `json:"law_item"` + Content string `json:"content"` + SourceFile string `json:"source_file"` // 法条所在语料文件名(判定时经 chunk 溯源填充) } // LawsRefs 解析 laws JSON 文本为法条引用数组(非法 JSON 返回空数组) diff --git a/kb/service/annotation_service.go b/kb/service/annotation_service.go index 839c8b9..3e31f6c 100644 --- a/kb/service/annotation_service.go +++ b/kb/service/annotation_service.go @@ -21,6 +21,7 @@ import ( "rag-local/kb/model/domain" "rag-local/kb/model/entity" + emodel "github.com/cloudwego/eino/components/model" "github.com/cloudwego/eino/schema" "github.com/gogf/gf/v2/errors/gerror" "github.com/gogf/gf/v2/frame/g" @@ -176,6 +177,8 @@ func (s *annotationService) processOne(ctx context.Context) { s.fail(ctx, task, "构建对话模型失败: "+err.Error()) return } + // 推理模型默认输出长思考链,烧光上下文(与 kg_extract 同策略);须在提交池前预置(共享实例可变字段) + chatModel.DisableThinking() embedders := make(map[int64]*OpenAIEmbedder) dsNames := make(map[int64]string) @@ -203,6 +206,7 @@ func (s *annotationService) processOne(ctx context.Context) { err error } ch := make(chan clauseJobOut, len(clauses)) + runIds := make([]int64, 0, len(clauses)) var wg sync.WaitGroup for _, cl := range clauses { // 断点续跑:已完成且已有风险记录 → 跳过;已完成但无风险记录 → 仅当存在旧格式法条标注时重跑迁移 @@ -216,10 +220,7 @@ func (s *annotationService) processOne(ctx context.Context) { continue } } - if err := dao.ContractClause.UpdateStatus(ctx, cl.Id, consts.TaskStatusRunning, ""); err != nil { - g.Log().Errorf(ctx, "mark clause running failed: %v", err) - continue - } + runIds = append(runIds, cl.Id) wg.Add(1) if err := common.AnnotationClausePool.AddWithRecover(ctx, func(ctx context.Context) { defer wg.Done() @@ -245,9 +246,31 @@ func (s *annotationService) processOne(ctx context.Context) { g.Log().Errorf(ctx, "submit clause %d failed: %v", cl.Id, err) } } + // 批量标记进行中(避免逐条款 UPDATE) + if len(runIds) > 0 { + if err := dao.ContractClause.UpdateStatuses(ctx, runIds, consts.TaskStatusRunning, ""); err != nil { + g.Log().Errorf(ctx, "mark clauses running failed: %v", err) + } + } go func() { wg.Wait(); close(ch) }() + // 进度按本地计数推进(断点续跑时已完成条款计入基数),每完成一条刷一次,避免逐条款读库统计 + progressDone := 0 + for _, cl := range clauses { + if cl.Status == consts.TaskStatusDone { + progressDone++ + } + } + totalClauses := len(clauses) + tickProgress := func() { + progressDone++ + if err := dao.ContractTask.UpdateProgress(ctx, task.Id, totalClauses, progressDone); err != nil { + g.Log().Warningf(ctx, "update annotation progress failed: %v", err) + } + } + failed := 0 + doneIds := make([]int64, 0, len(runIds)) for out := range ch { if out.err != nil { failed++ @@ -256,8 +279,8 @@ func (s *annotationService) processOne(ctx context.Context) { } if out.noCands { // 无候选视为完成(无标注),避免卡住进度 - _ = dao.ContractClause.UpdateStatus(ctx, out.clauseId, consts.TaskStatusDone, "") - s.updateProgress(ctx, task.Id) + doneIds = append(doneIds, out.clauseId) + tickProgress() continue } // 幂等:重跑前清旧风险与旧格式法条标注,避免重复记录 @@ -278,8 +301,14 @@ func (s *annotationService) processOne(ctx context.Context) { continue } } - _ = dao.ContractClause.UpdateStatus(ctx, out.clauseId, consts.TaskStatusDone, "") - s.updateProgress(ctx, task.Id) + doneIds = append(doneIds, out.clauseId) + tickProgress() + } + // 完成状态批量落库(成功与无候选统一刷一次;失败走上面逐条,error_msg 各异) + if len(doneIds) > 0 { + if err := dao.ContractClause.UpdateStatuses(ctx, doneIds, consts.TaskStatusDone, ""); err != nil { + g.Log().Errorf(ctx, "mark clauses done failed: %v", err) + } } msg := "" @@ -291,23 +320,6 @@ func (s *annotationService) processOne(ctx context.Context) { } } -// updateProgress 以库内实际完成数更新任务进度(断点续跑时跳过已 done 条款也能算对) -func (s *annotationService) updateProgress(ctx context.Context, taskId int64) { - doneList, err := dao.ContractClause.ListByTask(ctx, taskId) - if err != nil { - return - } - done := 0 - for _, c := range doneList { - if c.Status == consts.TaskStatusDone { - done++ - } - } - if err := dao.ContractTask.UpdateProgress(ctx, taskId, len(doneList), done); err != nil { - g.Log().Warningf(ctx, "update annotation progress failed: %v", err) - } -} - // annoRecallHit 单数据集召回结果(排名用于 RRF 融合) type annoRecallHit struct { ChunkId int64 @@ -359,13 +371,31 @@ func (s *annotationService) recallCandidates(ctx context.Context, clause *entity if len(cands) > consts.AnnoMaxCandidates { cands = cands[:consts.AnnoMaxCandidates] } - for i := range cands { - if chunk, err := dao.Chunk.GetOne(ctx, cands[i].ChunkId); err == nil && chunk != nil { - cands[i].ContentFull = chunk.Content - cands[i].Content = truncateRunes(chunk.Content, consts.AnnoCandidateMaxChars) - } + // 候选内容批量加载(单次 IN 查回内存映射,禁止逐条 GetOne 的 N+1);chunk 已删除的候选丢弃 + chunkIds := make([]int64, 0, len(cands)) + for _, c := range cands { + chunkIds = append(chunkIds, c.ChunkId) } - return cands, nil + chunks, err := dao.Chunk.ListByIds(ctx, chunkIds) + if err != nil { + g.Log().Warningf(ctx, "load candidate chunks failed: %v", err) + chunks = nil + } + contentByChunk := make(map[int64]string, len(chunks)) + for _, ch := range chunks { + contentByChunk[ch.Id] = ch.Content + } + kept := cands[:0] + for _, c := range cands { + content, ok := contentByChunk[c.ChunkId] + if !ok { + continue + } + c.ContentFull = content + c.Content = truncateRunes(content, consts.AnnoCandidateMaxChars) + kept = append(kept, c) + } + return kept, nil } // recallOneDataset 单数据集召回:向量检索 + FTS 检索(纯读,供池内并发调用) @@ -412,8 +442,19 @@ func (s *annotationService) judgeRisks(ctx context.Context, model *OpenAIChatMod sb.WriteString("你是资深法律顾问,负责审查合同条款的法律风险。请结合候选法律条文,识别该合同条款存在的法律风险点(条款与法律强制性规定冲突、遗漏法定必备内容、赔偿/补偿标准低于法定标准、期限或程序违法、表述模糊导致争议等)。\n\n【合同条款】\n") sb.WriteString(clause.Title + " " + clause.Content) sb.WriteString("\n\n【候选法律条文】\n") + // 候选按相关度降序,按上下文预算贪心填充(首条强制入队保底);被截断的候选不展示, + // 编号连续映射回 cands,LLM 引用编号受展示条数约束 + shown := 0 + budget := consts.AnnoJudgePromptBudget for i, c := range cands { - sb.WriteString(fmt.Sprintf("[%d]《%s》%s\n", i+1, c.LawTitle, c.Content)) + item := fmt.Sprintf("[%d]《%s》%s\n", i+1, c.LawTitle, truncateRunes(c.Content, consts.AnnoJudgeCandidateChars)) + itemLen := len([]rune(item)) + if shown > 0 && itemLen > budget { + break + } + sb.WriteString(item) + shown++ + budget -= itemLen } sb.WriteString(fmt.Sprintf("\n请输出该条款的风险点(0~%d 个,没有风险输出空数组)。每条风险点:\n", consts.RiskMaxPerClause)) sb.WriteString("- level:风险等级,high=违反强制性规定/可能导致合同无效或赔偿,mid=约定与法律不符但可补救,low=表述瑕疵或建议性提示\n") @@ -422,7 +463,8 @@ func (s *annotationService) judgeRisks(ctx context.Context, model *OpenAIChatMod sb.WriteString("只输出 JSON,不要其他内容:") sb.WriteString(`{"risks":[{"level":"high|mid|low","desc":"...","laws":[{"cand":1,"law_item":"第九十二条"}]}]}`) - msg, err := model.Generate(ctx, []*schema.Message{{Role: schema.User, Content: sb.String()}}) + msg, err := model.Generate(ctx, []*schema.Message{{Role: schema.User, Content: sb.String()}}, + emodel.WithMaxTokens(consts.AnnoJudgeMaxTokens)) if err != nil { return nil, err } @@ -443,6 +485,43 @@ func (s *annotationService) judgeRisks(ctx context.Context, model *OpenAIChatMod if err := json.Unmarshal([]byte(content), &resp); err != nil { return nil, gerror.Wrap(err, "解析风险判定结果失败: "+msg.Content) } + // 法条来源溯源(展示增强,失败仅告警不阻断判定):候选 chunk → 文档名,单表两条 SQL + 内存组装 + srcByChunk := make(map[int64]string, len(cands)) + chunkIds := make([]int64, 0, len(cands)) + seenChunk := make(map[int64]bool, len(cands)) + for _, c := range cands { + if c.ChunkId > 0 && !seenChunk[c.ChunkId] { + seenChunk[c.ChunkId] = true + chunkIds = append(chunkIds, c.ChunkId) + } + } + if len(chunkIds) > 0 { + docByChunk, err := dao.Chunk.ListDocumentIdsByChunkIds(ctx, chunkIds) + if err != nil { + g.Log().Warningf(ctx, "resolve chunk document failed: %v", err) + } else { + docIds := make([]int64, 0, len(docByChunk)) + for _, docId := range docByChunk { + docIds = append(docIds, docId) + } + if len(docIds) > 0 { + docs, err := dao.Document.ListByIds(ctx, docIds) + if err != nil { + g.Log().Warningf(ctx, "resolve document name failed: %v", err) + } else { + nameByDoc := make(map[int64]string, len(docs)) + for _, doc := range docs { + nameByDoc[doc.Id] = doc.Filename + } + for chunkId, docId := range docByChunk { + if f := nameByDoc[docId]; f != "" { + srcByChunk[chunkId] = f + } + } + } + } + } + } risks := make([]*entity.ContractRisk, 0, len(resp.Risks)) for _, r := range resp.Risks { desc := strings.TrimSpace(r.Desc) @@ -457,7 +536,7 @@ func (s *annotationService) judgeRisks(ctx context.Context, model *OpenAIChatMod } refs := make([]entity.LawRef, 0, len(r.Laws)) for _, lr := range r.Laws { - if lr.Cand < 1 || lr.Cand > len(cands) { + if lr.Cand < 1 || lr.Cand > shown { continue } c := cands[lr.Cand-1] @@ -466,9 +545,10 @@ func (s *annotationService) judgeRisks(ctx context.Context, model *OpenAIChatMod lawContent = c.Content } refs = append(refs, entity.LawRef{ - LawTitle: c.LawTitle, - LawItem: lawItem, - Content: lawContent, + LawTitle: c.LawTitle, + LawItem: lawItem, + Content: lawContent, + SourceFile: srcByChunk[c.ChunkId], }) } lawsJson, _ := json.Marshal(refs) @@ -680,8 +760,12 @@ h1{font-size:20px;text-align:center;margin-bottom:4px} `` + levelText[r.Level] + `` + `