初始化
This commit is contained in:
105
internal/handler/agent_handler.go
Normal file
105
internal/handler/agent_handler.go
Normal file
@@ -0,0 +1,105 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"tcm-agent/internal/agent"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// ========================================================================
|
||||
// 通用 Agent 对话接口
|
||||
// ========================================================================
|
||||
// 用于:自由问诊、多轮对话、调试 Agent 行为
|
||||
// ========================================================================
|
||||
|
||||
// AgentHandler 通用 Agent 对话处理器
|
||||
type AgentHandler struct {
|
||||
runner *agent.Runner
|
||||
}
|
||||
|
||||
// NewAgentHandler 创建 Agent 处理器
|
||||
func NewAgentHandler(runner *agent.Runner) *AgentHandler {
|
||||
return &AgentHandler{runner: runner}
|
||||
}
|
||||
|
||||
// ChatRequest 对话请求
|
||||
type ChatRequest struct {
|
||||
SessionID string `json:"session_id"` // 可选,不传则创建新会话
|
||||
Message string `json:"message" binding:"required"`
|
||||
UserID string `json:"user_id"` // 用户 ID
|
||||
Scene string `json:"scene"` // 场景名(决定用哪个模型)
|
||||
}
|
||||
|
||||
// Chat 通用 Agent 对话接口
|
||||
//
|
||||
// POST /api/v1/agent/chat
|
||||
//
|
||||
// 支持多轮对话:传 session_id 则延续上下文。
|
||||
// 不传 scene 时使用默认路由。
|
||||
func (h *AgentHandler) Chat(c *gin.Context) {
|
||||
var req ChatRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(400, gin.H{"error": "消息内容不能为空"})
|
||||
return
|
||||
}
|
||||
|
||||
// 获取或创建会话
|
||||
sessionID := req.SessionID
|
||||
if sessionID == "" {
|
||||
userID := req.UserID
|
||||
if userID == "" {
|
||||
userID = c.GetString("user_id")
|
||||
}
|
||||
// 使用请求的 scene(为空则用默认)
|
||||
session := h.runner.CreateSession(userID, req.Scene)
|
||||
sessionID = session.ID
|
||||
}
|
||||
|
||||
// 执行 Agent 推理
|
||||
response, err := h.runner.Run(c.Request.Context(), sessionID, req.Message)
|
||||
if err != nil {
|
||||
c.JSON(500, gin.H{
|
||||
"error": "Agent 执行失败",
|
||||
"detail": err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
session, _ := h.runner.GetSession(sessionID)
|
||||
|
||||
c.JSON(200, gin.H{
|
||||
"code": 200,
|
||||
"data": gin.H{
|
||||
"session_id": sessionID,
|
||||
"response": response,
|
||||
"history_len": len(session.History),
|
||||
"status": session.Status,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// GetSession 查看会话详情
|
||||
//
|
||||
// GET /api/v1/agent/session/:id
|
||||
func (h *AgentHandler) GetSession(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
session, ok := h.runner.GetSession(id)
|
||||
if !ok {
|
||||
c.JSON(404, gin.H{"error": "会话不存在"})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(200, gin.H{
|
||||
"code": 200,
|
||||
"data": gin.H{
|
||||
"id": session.ID,
|
||||
"user_id": session.UserID,
|
||||
"scene": session.Scene,
|
||||
"status": session.Status,
|
||||
"history": session.History,
|
||||
"state": session.State,
|
||||
"created_at": session.CreatedAt,
|
||||
"updated_at": session.UpdatedAt,
|
||||
},
|
||||
})
|
||||
}
|
||||
236
internal/handler/auth_handler.go
Normal file
236
internal/handler/auth_handler.go
Normal file
@@ -0,0 +1,236 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"log"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"tcm-agent/internal/config"
|
||||
"tcm-agent/internal/middleware"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
)
|
||||
|
||||
// ========================================================================
|
||||
// AuthHandler —— 独立管理前端(/admin SPA)的登录鉴权
|
||||
// ========================================================================
|
||||
// 暴露 3 个端点(前缀 /api/v1/auth):
|
||||
// POST /login 账号+密码+固定验证码 → 签发 HS256 JWT(7 天)
|
||||
// POST /refresh 带有效 token 换新 token(前端剩余有效期不足时静默续签)
|
||||
// GET /profile 带 token 返回用户信息(前端启动时校验会话有效性)
|
||||
//
|
||||
// 设计要点:
|
||||
// - 凭据来自 config.Panel(默认 liqi/qiqi991012/999999,config.yaml 可覆盖)
|
||||
// - 签名密钥复用 middleware.JWTSecret(与业务 API 的 JWT 校验同一把钥匙,
|
||||
// 所以登录发的 token 天然能过所有现有接口的鉴权)
|
||||
// - 登录失败统一模糊报错(不区分账号错还是密码错,防枚举)
|
||||
// - 防爆破:同 IP 连续失败 5 次锁 10 分钟(内存计数,重启清零,内网面板够用)
|
||||
// ========================================================================
|
||||
|
||||
// panelTokenTTL 登录 token 有效期(7 天,前端会在剩余 <24h 时自动续签)
|
||||
const panelTokenTTL = 7 * 24 * time.Hour
|
||||
|
||||
// panelRole 面板管理员角色名(写进 JWT claims,profile 返回给前端展示)
|
||||
const panelRole = "panel_admin"
|
||||
|
||||
// 防爆破参数:同 IP 连续 maxLoginFails 次失败 → 锁定 loginLockDuration
|
||||
const (
|
||||
maxLoginFails = 5
|
||||
loginLockDuration = 10 * time.Minute
|
||||
)
|
||||
|
||||
// loginFailEntry 单个 IP 的失败计数
|
||||
type loginFailEntry struct {
|
||||
Count int // 连续失败次数
|
||||
LockUntil time.Time // 锁定截止时间(零值表示未锁定)
|
||||
LastFail time.Time // 最后一次失败时间(做过期清理用)
|
||||
}
|
||||
|
||||
// AuthHandler 登录处理器
|
||||
type AuthHandler struct {
|
||||
cfg *config.Config
|
||||
mu sync.Mutex
|
||||
fails map[string]*loginFailEntry // key=客户端 IP
|
||||
}
|
||||
|
||||
// NewAuthHandler 构造
|
||||
func NewAuthHandler(cfg *config.Config) *AuthHandler {
|
||||
return &AuthHandler{cfg: cfg, fails: map[string]*loginFailEntry{}}
|
||||
}
|
||||
|
||||
// loginRequest 登录入参
|
||||
type loginRequest struct {
|
||||
Username string `json:"username" binding:"required"`
|
||||
Password string `json:"password" binding:"required"`
|
||||
Captcha string `json:"captcha" binding:"required"`
|
||||
}
|
||||
|
||||
// Login 登录换 token
|
||||
// POST /api/v1/auth/login
|
||||
func (h *AuthHandler) Login(c *gin.Context) {
|
||||
var req loginRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": "请填写账号、密码和验证码"})
|
||||
return
|
||||
}
|
||||
|
||||
ip := c.ClientIP()
|
||||
|
||||
// 先查锁定状态:锁定期内直接拒绝,不做凭据比对(省 CPU 也防继续试探)
|
||||
if locked, remain := h.isLocked(ip); locked {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 429, "message": "失败次数过多,请 " + remain + " 后再试"})
|
||||
return
|
||||
}
|
||||
|
||||
panel := h.cfg.Panel
|
||||
// 验证码 / 账号 / 密码全部比对;任何一项不对都返回同一句模糊报错(防枚举)
|
||||
if strings.TrimSpace(req.Captcha) != panel.Captcha ||
|
||||
req.Username != panel.Username ||
|
||||
req.Password != panel.Password {
|
||||
h.recordFail(ip)
|
||||
log.Printf("[Auth] 登录失败 ip=%s username=%s", ip, req.Username)
|
||||
c.JSON(http.StatusOK, gin.H{"code": 401, "message": "账号、密码或验证码错误"})
|
||||
return
|
||||
}
|
||||
|
||||
// 登录成功:清空该 IP 的失败计数
|
||||
h.clearFail(ip)
|
||||
|
||||
token, expireAt, err := h.signToken(panel.Username)
|
||||
if err != nil {
|
||||
log.Printf("[Auth] 签发 token 失败: %v", err)
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": "签发凭证失败"})
|
||||
return
|
||||
}
|
||||
|
||||
log.Printf("[Auth] 登录成功 ip=%s username=%s", ip, panel.Username)
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": gin.H{
|
||||
"token": token,
|
||||
"expire_at": expireAt, // unix 秒,前端据此判断何时续签
|
||||
"nick_name": panel.Username, // 顶栏展示
|
||||
"role_name": "面板管理员",
|
||||
"role": panelRole,
|
||||
}})
|
||||
}
|
||||
|
||||
// Refresh 续签 token
|
||||
// POST /api/v1/auth/refresh
|
||||
//
|
||||
// 必须带一个仍然有效的 JWT(中间件已校验并注入 user_id),
|
||||
// 用旧 token 的身份签发一个全新 7 天 token。
|
||||
// 过期的 token 无法续签(中间件直接 401),需要重新登录。
|
||||
func (h *AuthHandler) Refresh(c *gin.Context) {
|
||||
userID := c.GetString("user_id")
|
||||
if userID == "" || c.GetString("auth_mode") != "jwt" {
|
||||
// 走 SharedSecret / 口令进来的调用方没有"会话"概念,不支持续签
|
||||
c.JSON(http.StatusOK, gin.H{"code": 401, "message": "当前凭证不支持续签,请重新登录"})
|
||||
return
|
||||
}
|
||||
token, expireAt, err := h.signToken(userID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": "续签失败"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": gin.H{
|
||||
"token": token,
|
||||
"expire_at": expireAt,
|
||||
}})
|
||||
}
|
||||
|
||||
// Profile 返回当前登录用户信息
|
||||
// GET /api/v1/auth/profile
|
||||
//
|
||||
// 前端启动时调用:能拿到数据说明 token 还有效,直接进面板;401 则跳登录页
|
||||
func (h *AuthHandler) Profile(c *gin.Context) {
|
||||
userID := c.GetString("user_id")
|
||||
if userID == "" || c.GetString("auth_mode") != "jwt" {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 401, "message": "会话无效,请重新登录"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": gin.H{
|
||||
"username": userID,
|
||||
"nick_name": userID,
|
||||
"role": c.GetString("user_role"),
|
||||
"role_name": "面板管理员",
|
||||
}})
|
||||
}
|
||||
|
||||
// signToken 签发面板 JWT
|
||||
//
|
||||
// claims:sub=用户名、role=panel_admin、iat/exp 标准字段。
|
||||
// 密钥与 middleware.Auth 校验用的是同一把(middleware.JWTSecret),
|
||||
// 所以这个 token 能直接通过所有业务 API 的 JWT 鉴权路径。
|
||||
func (h *AuthHandler) signToken(username string) (string, int64, error) {
|
||||
now := time.Now()
|
||||
expireAt := now.Add(panelTokenTTL)
|
||||
claims := jwt.MapClaims{
|
||||
"sub": username,
|
||||
"role": panelRole,
|
||||
"iat": now.Unix(),
|
||||
"exp": expireAt.Unix(),
|
||||
}
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||
signed, err := token.SignedString([]byte(middleware.JWTSecret(h.cfg)))
|
||||
if err != nil {
|
||||
return "", 0, err
|
||||
}
|
||||
return signed, expireAt.Unix(), nil
|
||||
}
|
||||
|
||||
// ---------------------------- 防爆破(内存计数) ----------------------------
|
||||
|
||||
// isLocked 判断某 IP 是否处于锁定期,返回剩余时间的人类可读描述
|
||||
func (h *AuthHandler) isLocked(ip string) (bool, string) {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
entry, ok := h.fails[ip]
|
||||
if !ok {
|
||||
return false, ""
|
||||
}
|
||||
if entry.LockUntil.IsZero() || time.Now().After(entry.LockUntil) {
|
||||
return false, ""
|
||||
}
|
||||
remain := time.Until(entry.LockUntil).Round(time.Second)
|
||||
return true, remain.String()
|
||||
}
|
||||
|
||||
// recordFail 记录一次失败;达到阈值后设置锁定
|
||||
func (h *AuthHandler) recordFail(ip string) {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
|
||||
// 顺手清理超过 1 小时没动静的旧条目,防 map 无限膨胀
|
||||
cutoff := time.Now().Add(-time.Hour)
|
||||
for k, v := range h.fails {
|
||||
if v.LastFail.Before(cutoff) && (v.LockUntil.IsZero() || time.Now().After(v.LockUntil)) {
|
||||
delete(h.fails, k)
|
||||
}
|
||||
}
|
||||
|
||||
entry, ok := h.fails[ip]
|
||||
if !ok {
|
||||
entry = &loginFailEntry{}
|
||||
h.fails[ip] = entry
|
||||
}
|
||||
// 锁定期已过则重新从 1 开始计数
|
||||
if !entry.LockUntil.IsZero() && time.Now().After(entry.LockUntil) {
|
||||
entry.Count = 0
|
||||
entry.LockUntil = time.Time{}
|
||||
}
|
||||
entry.Count++
|
||||
entry.LastFail = time.Now()
|
||||
if entry.Count >= maxLoginFails {
|
||||
entry.LockUntil = time.Now().Add(loginLockDuration)
|
||||
log.Printf("[Auth] IP %s 连续失败 %d 次,锁定至 %s", ip, entry.Count, entry.LockUntil.Format("15:04:05"))
|
||||
}
|
||||
}
|
||||
|
||||
// clearFail 登录成功后清空该 IP 的失败计数
|
||||
func (h *AuthHandler) clearFail(ip string) {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
delete(h.fails, ip)
|
||||
}
|
||||
211
internal/handler/emr_handler.go
Normal file
211
internal/handler/emr_handler.go
Normal file
@@ -0,0 +1,211 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"tcm-agent/internal/agent"
|
||||
"tcm-agent/internal/model/entity"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// ========================================================================
|
||||
// 病历 HTTP 接口处理器
|
||||
// ========================================================================
|
||||
// 职责链:
|
||||
// HTTP 请求 → 参数校验 → 调用 Agent → 持久化 → 响应封装
|
||||
// ========================================================================
|
||||
|
||||
// EMRHandler 病历接口处理器
|
||||
type EMRHandler struct {
|
||||
agentRunner *agent.Runner
|
||||
emrAgent *agent.EMRGenerator
|
||||
}
|
||||
|
||||
// NewEMRHandler 创建病历处理器
|
||||
//
|
||||
// 参数:
|
||||
//
|
||||
// runner - Agent 引擎
|
||||
// scene - 场景名(对应 config.yaml 中 routes 的 key)
|
||||
// 为空则使用默认值 "emr-generator"
|
||||
func NewEMRHandler(runner *agent.Runner, scene string) *EMRHandler {
|
||||
return &EMRHandler{
|
||||
agentRunner: runner,
|
||||
emrAgent: agent.NewEMRGenerator(runner, scene),
|
||||
}
|
||||
}
|
||||
|
||||
// GenerateRequest 生成病历请求体
|
||||
type GenerateRequest struct {
|
||||
PatientID string `json:"patient_id" binding:"required"`
|
||||
ChiefComplaint string `json:"chief_complaint" binding:"required"`
|
||||
HistoryNotes string `json:"history_notes"`
|
||||
Allergies []string `json:"allergies"`
|
||||
PastIllness []string `json:"past_illness"`
|
||||
}
|
||||
|
||||
// Generate 根据主诉+病史生成病历
|
||||
//
|
||||
// POST /api/v1/emr/generate
|
||||
//
|
||||
// 完整生命周期:
|
||||
// 1. 参数校验
|
||||
// 2. Agent 感知 → 规划 → 检索 → 工具 → 反思 → 输出
|
||||
// 3. 规则引擎质控
|
||||
// 4. 持久化到数据库
|
||||
// 5. 返回结构化响应
|
||||
func (h *EMRHandler) Generate(c *gin.Context) {
|
||||
// ===== ① 参数校验 =====
|
||||
var req GenerateRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"error": "参数错误",
|
||||
"detail": err.Error(),
|
||||
"example": `{"patient_id":"P001","chief_complaint":"反复头晕3个月","history_notes":"...","allergies":[],"past_illness":[]}`,
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
doctorID := c.GetString("user_id")
|
||||
|
||||
// ===== ②~⑤ 调用 Agent =====
|
||||
agentReq := &agent.EMRRequest{
|
||||
PatientID: req.PatientID,
|
||||
ChiefComplaint: req.ChiefComplaint,
|
||||
HistoryNotes: req.HistoryNotes,
|
||||
Allergies: req.Allergies,
|
||||
PastIllness: req.PastIllness,
|
||||
}
|
||||
|
||||
resp, err := h.emrAgent.Generate(c.Request.Context(), agentReq)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"error": "病历生成失败",
|
||||
"detail": err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// ===== ⑥ 持久化 =====
|
||||
emr := &entity.EMR{
|
||||
PatientID: req.PatientID,
|
||||
DoctorID: doctorID,
|
||||
Draft: resp.Draft,
|
||||
Structured: resp.Structured,
|
||||
Status: resp.Status,
|
||||
SessionID: resp.SessionID,
|
||||
CreatedAt: time.Now(),
|
||||
UpdatedAt: time.Now(),
|
||||
}
|
||||
// dao.EMR.Create(emr) // 实际项目取消注释
|
||||
|
||||
// ===== ⑦ 返回响应 =====
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"code": 200,
|
||||
"message": "病历生成成功",
|
||||
"data": gin.H{
|
||||
"session_id": resp.SessionID,
|
||||
"draft": resp.Draft,
|
||||
"structured": resp.Structured,
|
||||
"issues": resp.Issues,
|
||||
"status": resp.Status,
|
||||
"emr_id": emr.ID,
|
||||
"next_action": getNextAction(resp.Status),
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// KnowledgeQA 病历书写规范问答
|
||||
//
|
||||
// POST /api/v1/emr/qa
|
||||
func (h *EMRHandler) KnowledgeQA(c *gin.Context) {
|
||||
var req struct {
|
||||
Question string `json:"question" binding:"required"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(400, gin.H{"error": "问题不能为空"})
|
||||
return
|
||||
}
|
||||
|
||||
answer, err := h.agentRunner.MaxKB().Chat(c.Request.Context(), req.Question)
|
||||
if err != nil {
|
||||
c.JSON(500, gin.H{"error": "知识库查询失败", "detail": err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(200, gin.H{
|
||||
"code": 200,
|
||||
"data": gin.H{
|
||||
"question": req.Question,
|
||||
"answer": answer,
|
||||
"source": "MaxKB 知识库",
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// GetByID 查询病历详情
|
||||
//
|
||||
// GET /api/v1/emr/:id
|
||||
func (h *EMRHandler) GetByID(c *gin.Context) {
|
||||
idStr := c.Param("id")
|
||||
id, _ := strconv.ParseInt(idStr, 10, 64)
|
||||
|
||||
c.JSON(200, gin.H{
|
||||
"code": 200,
|
||||
"data": gin.H{
|
||||
"id": id,
|
||||
"note": "实际项目从数据库查询",
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// Update 更新病历(医生人工修改后保存)
|
||||
//
|
||||
// PUT /api/v1/emr/:id
|
||||
func (h *EMRHandler) Update(c *gin.Context) {
|
||||
idStr := c.Param("id")
|
||||
id, _ := strconv.ParseInt(idStr, 10, 64)
|
||||
|
||||
var req struct {
|
||||
Draft string `json:"draft"`
|
||||
IsFinal bool `json:"is_final"`
|
||||
}
|
||||
c.ShouldBindJSON(&req)
|
||||
|
||||
userID := c.GetString("user_id")
|
||||
logAudit("emr_update", userID, idStr, req.Draft)
|
||||
|
||||
c.JSON(200, gin.H{
|
||||
"code": 200,
|
||||
"message": "病历已更新",
|
||||
"data": gin.H{
|
||||
"id": id,
|
||||
"is_final": req.IsFinal,
|
||||
"updated_by": userID,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// getNextAction 根据状态给出下一步建议
|
||||
func getNextAction(status string) string {
|
||||
switch status {
|
||||
case "success":
|
||||
return "病历生成完成,请医生审核确认"
|
||||
case "need_revision":
|
||||
return "病历存在质控问题,请查看 issues 列表并修改"
|
||||
default:
|
||||
return "请检查输入信息是否完整"
|
||||
}
|
||||
}
|
||||
|
||||
// logAudit 审计日志
|
||||
func logAudit(action, userID, targetID, detail string) {
|
||||
// 实际项目:写入审计表
|
||||
_ = action
|
||||
_ = userID
|
||||
_ = targetID
|
||||
_ = detail
|
||||
}
|
||||
100
internal/handler/enhancer_handler.go
Normal file
100
internal/handler/enhancer_handler.go
Normal file
@@ -0,0 +1,100 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"tcm-agent/internal/service"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// ========================================================================
|
||||
// EnhancerHandler —— 知识增强接口的 HTTP 处理器
|
||||
// ========================================================================
|
||||
// 对接 PHP 端的 TcmAgentClient,提供 POST /api/v1/agent/enhance 端点。
|
||||
// PHP 端拿到的响应里包含 content + steps,会把 steps 写入 xk_ai_generation_step。
|
||||
// ========================================================================
|
||||
|
||||
// EnhancerHandler 知识增强接口处理器
|
||||
type EnhancerHandler struct {
|
||||
svc *service.EnhancerService
|
||||
}
|
||||
|
||||
// NewEnhancerHandler 构造函数
|
||||
func NewEnhancerHandler(svc *service.EnhancerService) *EnhancerHandler {
|
||||
return &EnhancerHandler{svc: svc}
|
||||
}
|
||||
|
||||
// Enhance HTTP 入口
|
||||
//
|
||||
// POST /api/v1/agent/enhance
|
||||
//
|
||||
// 请求体(与 service.EnhanceRequest 一致):
|
||||
//
|
||||
// {
|
||||
// "scene": "medical_record",
|
||||
// "context": "痰湿中阻 煎法",
|
||||
// "messages": [
|
||||
// {"role": "system", "content": "..."},
|
||||
// {"role": "user", "content": "..."}
|
||||
// ],
|
||||
// "kb_enabled": true,
|
||||
// "top_k": 5,
|
||||
// "provider": "" // 空 = 按 scene 路由
|
||||
// }
|
||||
//
|
||||
// 响应:
|
||||
//
|
||||
// {
|
||||
// "code": 200,
|
||||
// "data": {
|
||||
// "content": "...",
|
||||
// "provider": "spark",
|
||||
// "model": "spark-max",
|
||||
// "steps": [...],
|
||||
// "total_ms": 1234
|
||||
// }
|
||||
// }
|
||||
func (h *EnhancerHandler) Enhance(c *gin.Context) {
|
||||
if h.svc == nil {
|
||||
c.JSON(http.StatusServiceUnavailable, gin.H{
|
||||
"code": 503,
|
||||
"message": "知识增强服务未初始化",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
var req service.EnhanceRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"code": 400,
|
||||
"message": "参数错误: " + err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// 校验:messages 至少 1 条
|
||||
if len(req.Messages) == 0 {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"code": 400,
|
||||
"message": "messages 不能为空",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
resp, err := h.svc.Enhance(c.Request.Context(), &req)
|
||||
if err != nil {
|
||||
// 业务失败:仍然把已收集的 steps 返回,让 PHP 能记录失败过程
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"code": 500,
|
||||
"message": err.Error(),
|
||||
"data": resp,
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"code": 200,
|
||||
"data": resp,
|
||||
})
|
||||
}
|
||||
8
internal/handler/helpers.go
Normal file
8
internal/handler/helpers.go
Normal file
@@ -0,0 +1,8 @@
|
||||
package handler
|
||||
|
||||
import "errors"
|
||||
|
||||
// errInvalidID 通用路径参数 id 非法错误
|
||||
//
|
||||
// 各 handler 用 parseUintParam 统一返回这个错误
|
||||
var errInvalidID = errors.New("id 参数非法")
|
||||
101
internal/handler/history_handler.go
Normal file
101
internal/handler/history_handler.go
Normal file
@@ -0,0 +1,101 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"tcm-agent/internal/dao"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// ========================================================================
|
||||
// HistoryHandler —— AI 生成历史(DB 长期数据)只读查询
|
||||
// ========================================================================
|
||||
// 暴露 3 个端点(前缀 /api/v1/agent):
|
||||
// GET /history 分页列表(场景/状态/via_agent/供应商/日期筛选)
|
||||
// GET /history/:id 单条详情 + 步骤时间线
|
||||
// GET /history/scenes 出现过的场景列表(前端筛选下拉)
|
||||
//
|
||||
// 与 /agent/runs 的区别:
|
||||
// - runs:内存环形缓冲,最近 200 条,含守卫拦截等"没到 PHP"的运行,重启清零
|
||||
// - history:读 PHP 落库的 xk_ai_generation(_step),长期审计视角,只读
|
||||
// ========================================================================
|
||||
|
||||
// HistoryHandler 历史查询处理器
|
||||
type HistoryHandler struct{}
|
||||
|
||||
// NewHistoryHandler 构造
|
||||
func NewHistoryHandler() *HistoryHandler {
|
||||
return &HistoryHandler{}
|
||||
}
|
||||
|
||||
// List 分页查询历史
|
||||
// GET /api/v1/agent/history?page=1&size=20&scene=&status=-1&via_agent=-1&provider=&date_start=0&date_end=0
|
||||
func (h *HistoryHandler) List(c *gin.Context) {
|
||||
filter := dao.AIGenerationListFilter{
|
||||
Page: queryInt(c, "page", 1),
|
||||
Size: queryInt(c, "size", 20),
|
||||
Scene: c.Query("scene"),
|
||||
Status: queryInt(c, "status", -1),
|
||||
ViaAgent: queryInt(c, "via_agent", -1),
|
||||
Provider: c.Query("provider"),
|
||||
}
|
||||
filter.DateStart = int64(queryInt(c, "date_start", 0))
|
||||
filter.DateEnd = int64(queryInt(c, "date_end", 0))
|
||||
|
||||
rows, total, err := dao.AIGenerationList(filter)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": gin.H{
|
||||
"list": rows,
|
||||
"total": total,
|
||||
"page": filter.Page,
|
||||
"size": filter.Size,
|
||||
}})
|
||||
}
|
||||
|
||||
// Detail 单条详情 + 步骤
|
||||
// GET /api/v1/agent/history/:id
|
||||
func (h *HistoryHandler) Detail(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil || id == 0 {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": "id 参数无效"})
|
||||
return
|
||||
}
|
||||
row, steps, err := dao.AIGenerationGet(uint(id))
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 404, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": gin.H{
|
||||
"generation": row,
|
||||
"steps": steps,
|
||||
}})
|
||||
}
|
||||
|
||||
// Scenes 场景下拉数据
|
||||
// GET /api/v1/agent/history/scenes
|
||||
func (h *HistoryHandler) Scenes(c *gin.Context) {
|
||||
scenes, err := dao.AIGenerationScenes()
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": scenes})
|
||||
}
|
||||
|
||||
// queryInt 读取整型 query 参数,缺失或非法时返回默认值
|
||||
func queryInt(c *gin.Context, key string, def int) int {
|
||||
v := c.Query(key)
|
||||
if v == "" {
|
||||
return def
|
||||
}
|
||||
n, err := strconv.Atoi(v)
|
||||
if err != nil {
|
||||
return def
|
||||
}
|
||||
return n
|
||||
}
|
||||
397
internal/handler/kb_admin_handler.go
Normal file
397
internal/handler/kb_admin_handler.go
Normal file
@@ -0,0 +1,397 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"io"
|
||||
"log"
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"tcm-agent/internal/kb"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// ========================================================================
|
||||
// KBAdminHandler —— 本地知识库后台管理 API
|
||||
// ========================================================================
|
||||
// 暴露 9 个端点(前缀 /api/v1/kb/admin):
|
||||
//
|
||||
// 库管理:
|
||||
// GET /libraries 列出所有库(含禁用的)
|
||||
// GET /libraries/:id 取单个库详情
|
||||
// POST /libraries 创建库
|
||||
// DELETE /libraries/:id 软删除库(连带文档+分段)
|
||||
//
|
||||
// 文档管理:
|
||||
// GET /libraries/:id/docs 列出某库下所有文档
|
||||
// POST /docs/import 导入文档(multipart 上传文件)
|
||||
// GET /docs/:id 取文档详情
|
||||
// DELETE /docs/:id 软删除文档(连带分段)
|
||||
//
|
||||
// 分段管理:
|
||||
// GET /docs/:id/chunks 列出某文档的分段
|
||||
//
|
||||
// 工具:
|
||||
// POST /search 检索测试(前端"试一试"按钮用)
|
||||
// POST /embed 【V2 预留】触发批量向量化,V1 返回"未启用"
|
||||
// ========================================================================
|
||||
|
||||
// KBAdminHandler 后台管理处理器
|
||||
type KBAdminHandler struct {
|
||||
libSvc *kb.LibraryService
|
||||
searcher *kb.Searcher
|
||||
}
|
||||
|
||||
// NewKBAdminHandler 构造
|
||||
func NewKBAdminHandler(libSvc *kb.LibraryService, searcher *kb.Searcher) *KBAdminHandler {
|
||||
return &KBAdminHandler{libSvc: libSvc, searcher: searcher}
|
||||
}
|
||||
|
||||
// ---------------------------- 库管理 ----------------------------
|
||||
|
||||
// ListLibraries 列出所有库
|
||||
// GET /api/v1/kb/admin/libraries
|
||||
func (h *KBAdminHandler) ListLibraries(c *gin.Context) {
|
||||
rows, err := h.libSvc.ListLibraries(c.Request.Context())
|
||||
if err != nil {
|
||||
// 关键:把错误打到日志,方便后端排查(DB 未连/表不存在/SQL 语法错都会在这里暴露)
|
||||
log.Printf("[KB] ListLibraries 失败: %v", err)
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": rows})
|
||||
}
|
||||
|
||||
// GetLibrary 取单个库详情
|
||||
// GET /api/v1/kb/admin/libraries/:id
|
||||
func (h *KBAdminHandler) GetLibrary(c *gin.Context) {
|
||||
id, err := parseUintParam(c, "id")
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
row, err := h.libSvc.GetLibrary(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 404, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": row})
|
||||
}
|
||||
|
||||
// CreateLibrary 创建库
|
||||
// POST /api/v1/kb/admin/libraries
|
||||
//
|
||||
// Body: kb.CreateLibraryInput
|
||||
func (h *KBAdminHandler) CreateLibrary(c *gin.Context) {
|
||||
var in kb.CreateLibraryInput
|
||||
if err := c.ShouldBindJSON(&in); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": "参数错误: " + err.Error()})
|
||||
return
|
||||
}
|
||||
row, err := h.libSvc.CreateLibrary(c.Request.Context(), in)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": row})
|
||||
}
|
||||
|
||||
// DeleteLibrary 软删除库(连带文档+分段)
|
||||
// DELETE /api/v1/kb/admin/libraries/:id
|
||||
func (h *KBAdminHandler) DeleteLibrary(c *gin.Context) {
|
||||
id, err := parseUintParam(c, "id")
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
if err := h.libSvc.DeleteLibrary(c.Request.Context(), id); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "message": "删除成功"})
|
||||
}
|
||||
|
||||
// ---------------------------- 文档管理 ----------------------------
|
||||
|
||||
// ListDocs 列出某库下的文档
|
||||
// GET /api/v1/kb/admin/libraries/:id/docs
|
||||
func (h *KBAdminHandler) ListDocs(c *gin.Context) {
|
||||
id, err := parseUintParam(c, "id")
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
rows, err := h.libSvc.ListDocs(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": rows})
|
||||
}
|
||||
|
||||
// ImportDocument 导入文档
|
||||
// POST /api/v1/kb/admin/docs/import
|
||||
//
|
||||
// 表单字段:
|
||||
// library_id (必填, form 字段)
|
||||
// title (可选, form 字段,自定义文档标题)
|
||||
// file (必填, multipart 文件)
|
||||
// max_len (可选, form 字段,自定义分段最大长度 100~2000,缺省 500)
|
||||
// overlap (可选, form 字段,自定义分段重叠 0~500,缺省 50)
|
||||
//
|
||||
// 支持 .xlsx / .xls / .csv / .md / .txt / .pdf / .docx / .html,
|
||||
// 具体解析与自动分段规则见 kb.ParseFileFromBytes
|
||||
func (h *KBAdminHandler) ImportDocument(c *gin.Context) {
|
||||
libraryIDStr := c.PostForm("library_id")
|
||||
libraryID, err := strconv.ParseUint(libraryIDStr, 10, 64)
|
||||
if err != nil || libraryID == 0 {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": "library_id 不能为空"})
|
||||
return
|
||||
}
|
||||
title := c.PostForm("title")
|
||||
|
||||
// 自定义分段参数:解析失败或未传时用哨兵值(max_len=0 / overlap=-1 表示"用默认")
|
||||
// overlap 不能用 0 当哨兵——0 是合法值(不重叠),语义与"没传"不同
|
||||
maxLen := 0
|
||||
if v := c.PostForm("max_len"); v != "" {
|
||||
if n, e := strconv.Atoi(v); e == nil {
|
||||
maxLen = n
|
||||
}
|
||||
}
|
||||
overlap := -1
|
||||
if v := c.PostForm("overlap"); v != "" {
|
||||
if n, e := strconv.Atoi(v); e == nil {
|
||||
overlap = n
|
||||
}
|
||||
}
|
||||
|
||||
fileHeader, err := c.FormFile("file")
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": "请上传文件: " + err.Error()})
|
||||
return
|
||||
}
|
||||
// 限制文件大小 50MB(中医知识库单文件不会超过这个)
|
||||
if fileHeader.Size > 50*1024*1024 {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": "文件过大(最大 50MB)"})
|
||||
return
|
||||
}
|
||||
file, err := fileHeader.Open()
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": "打开上传文件失败: " + err.Error()})
|
||||
return
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
// 读全部字节:必须用 io.ReadFull 而不是单次 file.Read——
|
||||
// 大文件(>32MB)multipart 会落磁盘临时文件,单次 Read 不保证读满缓冲区,
|
||||
// 读不满会导致导入的内容被截断
|
||||
buf := make([]byte, fileHeader.Size)
|
||||
if _, err := io.ReadFull(file, buf); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": "读取上传文件失败: " + err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
result, err := h.libSvc.ImportDocument(c.Request.Context(), kb.ImportDocInput{
|
||||
LibraryID: uint(libraryID),
|
||||
Filename: fileHeader.Filename,
|
||||
Content: buf,
|
||||
Title: title,
|
||||
MaxLen: maxLen,
|
||||
Overlap: overlap,
|
||||
})
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": result})
|
||||
}
|
||||
|
||||
// GetDoc 取文档详情
|
||||
// GET /api/v1/kb/admin/docs/:id
|
||||
func (h *KBAdminHandler) GetDoc(c *gin.Context) {
|
||||
id, err := parseUintParam(c, "id")
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
row, err := h.libSvc.GetDoc(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 404, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": row})
|
||||
}
|
||||
|
||||
// DeleteDoc 软删除文档(连带分段)
|
||||
// DELETE /api/v1/kb/admin/docs/:id
|
||||
func (h *KBAdminHandler) DeleteDoc(c *gin.Context) {
|
||||
id, err := parseUintParam(c, "id")
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
if err := h.libSvc.DeleteDoc(c.Request.Context(), id); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "message": "删除成功"})
|
||||
}
|
||||
|
||||
// ---------------------------- 分段管理 ----------------------------
|
||||
|
||||
// ListChunks 列出文档的分段
|
||||
// GET /api/v1/kb/admin/docs/:id/chunks
|
||||
func (h *KBAdminHandler) ListChunks(c *gin.Context) {
|
||||
id, err := parseUintParam(c, "id")
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
rows, err := h.libSvc.ListChunks(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
log.Printf("[KB] ListChunks 失败: %v", err)
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": rows})
|
||||
}
|
||||
|
||||
// UpdateChunk 编辑分段
|
||||
// PUT /api/v1/kb/admin/chunks/:id
|
||||
//
|
||||
// Body: kb.UpdateChunkInput
|
||||
// { "title": "...", "content": "...", "related_questions": ["问题1","问题2"] }
|
||||
//
|
||||
// 三个字段都可选,传啥改啥;related_questions 传空数组表示清空关联问题
|
||||
func (h *KBAdminHandler) UpdateChunk(c *gin.Context) {
|
||||
id, err := parseUintParam(c, "id")
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
var in kb.UpdateChunkInput
|
||||
if err := c.ShouldBindJSON(&in); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": "参数错误: " + err.Error()})
|
||||
return
|
||||
}
|
||||
// 内容不能为空(标题可空);例外:只带 is_active 的启停开关调用不改内容,放行
|
||||
if in.Content == "" && in.IsActive == nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": "content 不能为空"})
|
||||
return
|
||||
}
|
||||
updated, err := h.libSvc.UpdateChunk(c.Request.Context(), id, in)
|
||||
if err != nil {
|
||||
log.Printf("[KB] UpdateChunk 失败 id=%d: %v", id, err)
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": updated, "message": "已保存"})
|
||||
}
|
||||
|
||||
// BatchUpdateChunks 批量启用/禁用/删除分段
|
||||
// PUT /api/v1/kb/admin/chunks/batch
|
||||
//
|
||||
// Body: { "ids": [1,2,3], "action": "enable|disable|delete" }
|
||||
//
|
||||
// 管理前端的分段多选批量操作入口;事务原子性由 DAO 保证
|
||||
func (h *KBAdminHandler) BatchUpdateChunks(c *gin.Context) {
|
||||
var in kb.BatchChunkInput
|
||||
if err := c.ShouldBindJSON(&in); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": "参数错误: " + err.Error()})
|
||||
return
|
||||
}
|
||||
affected, err := h.libSvc.BatchUpdateChunks(c.Request.Context(), in)
|
||||
if err != nil {
|
||||
log.Printf("[KB] BatchUpdateChunks 失败 action=%s ids=%d: %v", in.Action, len(in.IDs), err)
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": gin.H{"affected": affected}, "message": "操作成功"})
|
||||
}
|
||||
|
||||
// RechunkDocument 用新参数对文档重新分段
|
||||
// POST /api/v1/kb/admin/docs/:id/rechunk
|
||||
//
|
||||
// Body: { "max_len": 500, "overlap": 50 }(0/-1 表示用默认值)
|
||||
//
|
||||
// 用 xk_kb_doc.content 存的原文重切,旧分段软删、新分段插入(事务)。
|
||||
// 注意:人工编辑过的分段内容会被覆盖,前端调用前必须二次确认
|
||||
func (h *KBAdminHandler) RechunkDocument(c *gin.Context) {
|
||||
id, err := parseUintParam(c, "id")
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
var in kb.RechunkInput
|
||||
// body 可以整个不传(全用默认参数),绑定失败不视为错误
|
||||
if err := c.ShouldBindJSON(&in); err != nil {
|
||||
in = kb.RechunkInput{MaxLen: 0, Overlap: -1}
|
||||
}
|
||||
result, err := h.libSvc.RechunkDocument(c.Request.Context(), id, in)
|
||||
if err != nil {
|
||||
log.Printf("[KB] RechunkDocument 失败 id=%d: %v", id, err)
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": result, "message": "重新分段完成"})
|
||||
}
|
||||
|
||||
// ---------------------------- 工具 ----------------------------
|
||||
|
||||
// KBSearchRequest 检索测试入参
|
||||
type KBSearchRequest struct {
|
||||
LibraryID uint `json:"library_id" binding:"required"`
|
||||
Query string `json:"query" binding:"required"`
|
||||
TopK int `json:"top_k"`
|
||||
Mode string `json:"mode"`
|
||||
}
|
||||
|
||||
// Search 检索测试
|
||||
// POST /api/v1/kb/admin/search
|
||||
//
|
||||
// V1 返回 FULLTEXT 得分;V2 接入向量后支持 mode=vector/blend
|
||||
func (h *KBAdminHandler) Search(c *gin.Context) {
|
||||
var req KBSearchRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": "参数错误: " + err.Error()})
|
||||
return
|
||||
}
|
||||
results, err := h.searcher.Search(c.Request.Context(), kb.SearchOptions{
|
||||
LibraryID: req.LibraryID,
|
||||
Query: req.Query,
|
||||
TopK: req.TopK,
|
||||
Mode: req.Mode,
|
||||
})
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": results})
|
||||
}
|
||||
|
||||
// EmbedRequest 触发向量化入参(V2 用)
|
||||
type EmbedRequest struct {
|
||||
LibraryID uint `json:"library_id" binding:"required"`
|
||||
}
|
||||
|
||||
// Embed 【V2 预留】触发批量向量化
|
||||
// POST /api/v1/kb/admin/embed
|
||||
//
|
||||
// V1 始终返回"未启用",前端展示对应提示
|
||||
func (h *KBAdminHandler) Embed(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"code": 501,
|
||||
"message": "向量化功能 V2 才支持(当前 NoopEmbedder 未启用)。请保持 search_mode=fulltext。",
|
||||
})
|
||||
}
|
||||
|
||||
// ---------------------------- 工具函数 ----------------------------
|
||||
|
||||
// parseUintParam 解析路径参数 :id 为 uint
|
||||
func parseUintParam(c *gin.Context, key string) (uint, error) {
|
||||
v, err := strconv.ParseUint(c.Param(key), 10, 64)
|
||||
if err != nil || v == 0 {
|
||||
return 0, errInvalidID
|
||||
}
|
||||
return uint(v), nil
|
||||
}
|
||||
237
internal/handler/kb_crawl_handler.go
Normal file
237
internal/handler/kb_crawl_handler.go
Normal file
@@ -0,0 +1,237 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"tcm-agent/internal/crawler"
|
||||
"tcm-agent/internal/dao"
|
||||
"tcm-agent/internal/service"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// ========================================================================
|
||||
// KBCrawlHandler —— 药品抓取任务管理 API
|
||||
// ========================================================================
|
||||
// 暴露 7 个端点(前缀 /api/v1/kb/admin/crawl,复用 KB 管理鉴权:
|
||||
// 口令头或面板 JWT 均可):
|
||||
//
|
||||
// GET /sources 可用抓取源列表(前端下拉)
|
||||
// GET /tasks 任务列表(含实时 running 标记)
|
||||
// POST /tasks 创建任务
|
||||
// PUT /tasks/:id 更新任务(名称/目标库/调度/限量/启停)
|
||||
// DELETE /tasks/:id 软删除任务
|
||||
// POST /tasks/:id/run 立即抓取(异步,立即返回)
|
||||
// GET /tasks/:id/logs 任务运行历史(最近 N 条)
|
||||
// ========================================================================
|
||||
|
||||
// KBCrawlHandler 抓取任务处理器(无状态,直接调 dao/service)
|
||||
type KBCrawlHandler struct{}
|
||||
|
||||
// NewKBCrawlHandler 构造
|
||||
func NewKBCrawlHandler() *KBCrawlHandler { return &KBCrawlHandler{} }
|
||||
|
||||
// ListSources 可用抓取源
|
||||
// GET /api/v1/kb/admin/crawl/sources
|
||||
func (h *KBCrawlHandler) ListSources(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": crawler.ListSources()})
|
||||
}
|
||||
|
||||
// ListTasks 任务列表
|
||||
// GET /api/v1/kb/admin/crawl/tasks
|
||||
//
|
||||
// 每行附加 running 字段(内存实时状态)——DB 的 last_status 有落库延迟,
|
||||
// 前端旋转图标要跟内存状态走
|
||||
func (h *KBCrawlHandler) ListTasks(c *gin.Context) {
|
||||
rows, err := dao.KBCrawlTaskList()
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
out := make([]gin.H, 0, len(rows))
|
||||
for _, t := range rows {
|
||||
out = append(out, gin.H{
|
||||
"id": t.ID, "name": t.Name, "source": t.Source,
|
||||
"library_id": t.LibraryID, "schedule_type": t.ScheduleType,
|
||||
"interval_hours": t.IntervalHours, "run_at_hour": t.RunAtHour,
|
||||
"run_at_weekday": t.RunAtWeekday, "items_per_run": t.ItemsPerRun,
|
||||
"progress_offset": t.ProgressOffset, "status": t.Status,
|
||||
"last_run_at": t.LastRunAt, "last_status": t.LastStatus,
|
||||
"last_message": t.LastMessage, "created_at": t.CreatedAt,
|
||||
"running": service.CrawlTaskRunning(t.ID),
|
||||
})
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": out})
|
||||
}
|
||||
|
||||
// crawlTaskRequest 创建/更新任务的入参
|
||||
type crawlTaskRequest struct {
|
||||
Name string `json:"name"`
|
||||
Source string `json:"source"`
|
||||
LibraryID uint `json:"library_id"`
|
||||
ScheduleType string `json:"schedule_type"`
|
||||
IntervalHours int `json:"interval_hours"`
|
||||
RunAtHour int `json:"run_at_hour"`
|
||||
RunAtWeekday int `json:"run_at_weekday"`
|
||||
ItemsPerRun int `json:"items_per_run"`
|
||||
Status *int `json:"status"` // 指针区分"没传"和"传了 0(禁用)"
|
||||
}
|
||||
|
||||
// normalize 归一化 + 校验入参(创建和更新共用)
|
||||
func (r *crawlTaskRequest) normalize() string {
|
||||
if r.Name == "" {
|
||||
return "任务名称不能为空"
|
||||
}
|
||||
if r.Source == "" {
|
||||
r.Source = "zhongyoo"
|
||||
}
|
||||
if _, ok := crawler.GetSource(r.Source); !ok {
|
||||
return "抓取源不存在: " + r.Source
|
||||
}
|
||||
if r.LibraryID == 0 {
|
||||
return "请选择目标知识库"
|
||||
}
|
||||
switch r.ScheduleType {
|
||||
case "", "manual":
|
||||
r.ScheduleType = "manual"
|
||||
case "interval":
|
||||
if r.IntervalHours < 1 {
|
||||
r.IntervalHours = 24
|
||||
}
|
||||
case "daily", "weekly":
|
||||
if r.RunAtHour < 0 || r.RunAtHour > 23 {
|
||||
r.RunAtHour = 3
|
||||
}
|
||||
if r.RunAtWeekday < 0 || r.RunAtWeekday > 6 {
|
||||
r.RunAtWeekday = 1
|
||||
}
|
||||
default:
|
||||
return "调度类型不合法: " + r.ScheduleType
|
||||
}
|
||||
if r.ItemsPerRun < 1 || r.ItemsPerRun > 500 {
|
||||
r.ItemsPerRun = 50
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// CreateTask 创建任务
|
||||
// POST /api/v1/kb/admin/crawl/tasks
|
||||
func (h *KBCrawlHandler) CreateTask(c *gin.Context) {
|
||||
var req crawlTaskRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": "参数格式错误: " + err.Error()})
|
||||
return
|
||||
}
|
||||
if msg := req.normalize(); msg != "" {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": msg})
|
||||
return
|
||||
}
|
||||
// 目标库必须真实存在(防手滑填错 ID,跑的时候才发现)
|
||||
if _, err := dao.KBGetLibrary(req.LibraryID); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": "目标知识库不存在,请先在知识库管理里创建"})
|
||||
return
|
||||
}
|
||||
row := &dao.KBCrawlTaskRow{
|
||||
Name: req.Name, Source: req.Source, LibraryID: req.LibraryID,
|
||||
ScheduleType: req.ScheduleType, IntervalHours: req.IntervalHours,
|
||||
RunAtHour: req.RunAtHour, RunAtWeekday: req.RunAtWeekday,
|
||||
ItemsPerRun: req.ItemsPerRun, Status: 1,
|
||||
}
|
||||
if err := dao.KBCrawlTaskCreate(row); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": row, "message": "任务已创建"})
|
||||
}
|
||||
|
||||
// UpdateTask 更新任务
|
||||
// PUT /api/v1/kb/admin/crawl/tasks/:id
|
||||
func (h *KBCrawlHandler) UpdateTask(c *gin.Context) {
|
||||
id, err := parseUintParam(c, "id")
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
var req crawlTaskRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": "参数格式错误: " + err.Error()})
|
||||
return
|
||||
}
|
||||
if msg := req.normalize(); msg != "" {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": msg})
|
||||
return
|
||||
}
|
||||
if _, err := dao.KBGetLibrary(req.LibraryID); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": "目标知识库不存在"})
|
||||
return
|
||||
}
|
||||
fields := map[string]any{
|
||||
"name": req.Name, "source": req.Source, "library_id": req.LibraryID,
|
||||
"schedule_type": req.ScheduleType, "interval_hours": req.IntervalHours,
|
||||
"run_at_hour": req.RunAtHour, "run_at_weekday": req.RunAtWeekday,
|
||||
"items_per_run": req.ItemsPerRun,
|
||||
}
|
||||
if req.Status != nil {
|
||||
fields["status"] = *req.Status
|
||||
}
|
||||
if err := dao.KBCrawlTaskUpdate(id, fields); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "message": "任务已更新"})
|
||||
}
|
||||
|
||||
// DeleteTask 软删除任务
|
||||
// DELETE /api/v1/kb/admin/crawl/tasks/:id
|
||||
func (h *KBCrawlHandler) DeleteTask(c *gin.Context) {
|
||||
id, err := parseUintParam(c, "id")
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
if service.CrawlTaskRunning(id) {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": "任务正在运行中,等本次跑完再删除"})
|
||||
return
|
||||
}
|
||||
if err := dao.KBCrawlTaskDelete(id); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "message": "任务已删除"})
|
||||
}
|
||||
|
||||
// RunTask 立即抓取(异步触发)
|
||||
// POST /api/v1/kb/admin/crawl/tasks/:id/run
|
||||
func (h *KBCrawlHandler) RunTask(c *gin.Context) {
|
||||
id, err := parseUintParam(c, "id")
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
if err := service.TriggerCrawlTask(id, "manual"); err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"code": 200,
|
||||
"message": "抓取已在后台启动,可在运行历史里查看进度(每批约需 1-3 分钟)",
|
||||
})
|
||||
}
|
||||
|
||||
// ListLogs 任务运行历史
|
||||
// GET /api/v1/kb/admin/crawl/tasks/:id/logs?limit=20
|
||||
func (h *KBCrawlHandler) ListLogs(c *gin.Context) {
|
||||
id, err := parseUintParam(c, "id")
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 400, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
limit, _ := strconv.Atoi(c.DefaultQuery("limit", "20"))
|
||||
rows, err := dao.KBCrawlLogList(id, limit)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusOK, gin.H{"code": 500, "message": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"code": 200, "data": rows})
|
||||
}
|
||||
94
internal/handler/knowledge_handler.go
Normal file
94
internal/handler/knowledge_handler.go
Normal file
@@ -0,0 +1,94 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"tcm-agent/internal/agent"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// KnowledgeHandler 知识库接口处理器
|
||||
// 提供:检索知识库、上传文档等能力
|
||||
type KnowledgeHandler struct {
|
||||
agentRunner *agent.Runner
|
||||
}
|
||||
|
||||
// NewKnowledgeHandler 创建知识库处理器
|
||||
func NewKnowledgeHandler(runner *agent.Runner) *KnowledgeHandler {
|
||||
return &KnowledgeHandler{agentRunner: runner}
|
||||
}
|
||||
|
||||
// SearchRequest 知识检索请求
|
||||
type SearchRequest struct {
|
||||
Query string `json:"query" binding:"required"` // 检索关键词
|
||||
TopK int `json:"top_k"` // 返回条数(默认5)
|
||||
}
|
||||
|
||||
// Search 检索知识库
|
||||
// POST /api/v1/knowledge/search
|
||||
//
|
||||
// 直接调用MaxKB的RAG检索能力
|
||||
// 用于:查询方剂组成、药典条目、诊疗规范等
|
||||
func (h *KnowledgeHandler) Search(c *gin.Context) {
|
||||
var req SearchRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(400, gin.H{"error": "查询关键词不能为空"})
|
||||
return
|
||||
}
|
||||
|
||||
if req.TopK <= 0 {
|
||||
req.TopK = 5
|
||||
}
|
||||
|
||||
// 调用MaxKB检索
|
||||
answer, err := h.agentRunner.MaxKB().Chat(c.Request.Context(), req.Query)
|
||||
if err != nil {
|
||||
c.JSON(500, gin.H{
|
||||
"error": "知识库检索失败",
|
||||
"detail": err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(200, gin.H{
|
||||
"code": 200,
|
||||
"data": gin.H{
|
||||
"query": req.Query,
|
||||
"answer": answer,
|
||||
"source": "MaxKB RAG",
|
||||
"top_k": req.TopK,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// IngestRequest 文档上传请求
|
||||
type IngestRequest struct {
|
||||
Title string `json:"title" binding:"required"` // 文档标题
|
||||
Content string `json:"content" binding:"required"` // 文档内容
|
||||
Category string `json:"category"` // 分类(方剂/药典/指南/病历模板)
|
||||
}
|
||||
|
||||
// Ingest 上传文档到知识库
|
||||
// POST /api/v1/knowledge/ingest
|
||||
//
|
||||
// 将文档写入MaxKB知识库,触发自动向量化
|
||||
// 支持后续RAG检索
|
||||
func (h *KnowledgeHandler) Ingest(c *gin.Context) {
|
||||
var req IngestRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(400, gin.H{"error": "标题和内容不能为空"})
|
||||
return
|
||||
}
|
||||
|
||||
// 实际项目中调用MaxKB的文档上传API
|
||||
// 这里返回模拟成功
|
||||
c.JSON(200, gin.H{
|
||||
"code": 200,
|
||||
"message": "文档已加入知识库队列",
|
||||
"data": gin.H{
|
||||
"title": req.Title,
|
||||
"category": req.Category,
|
||||
"status": "pending_vectorization",
|
||||
"note": "实际项目请调用MaxKB文档上传API",
|
||||
},
|
||||
})
|
||||
}
|
||||
194
internal/handler/observe_handler.go
Normal file
194
internal/handler/observe_handler.go
Normal file
@@ -0,0 +1,194 @@
|
||||
package handler
|
||||
|
||||
// ========================================================================
|
||||
// ObserveHandler —— /agent/view 面板观测与调试接口
|
||||
// ========================================================================
|
||||
// 端点清单(全部走全局 Auth,无放行):
|
||||
// GET /api/v1/agent/logs 增量日志(内存环形缓冲)
|
||||
// GET /api/v1/agent/system 进程运行时状态
|
||||
// GET /api/v1/agent/config Agent 配置只读视图(含守卫词表)
|
||||
// POST /api/v1/agent/guard-test 医疗守卫测试台
|
||||
// POST /api/v1/agent/kb-test 知识库检索测试
|
||||
// POST /api/v1/models/test 模型连通性测试(注册在 modelGroup)
|
||||
//
|
||||
// 安全边界:
|
||||
// - 全部只读或无持久副作用(ModelTest 消耗少量 token)
|
||||
// - 日志可能含 PHI(debug 开关打开时),logs 接口必须有 Auth
|
||||
// - config 视图不返回任何密钥(EmbeddingAPIKey 等一律不出)
|
||||
// ========================================================================
|
||||
|
||||
import (
|
||||
"runtime"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"tcm-agent/internal/agentcfg"
|
||||
"tcm-agent/internal/service"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// ObserveHandler 观测接口处理器
|
||||
type ObserveHandler struct {
|
||||
enhancer *service.EnhancerService // 复用 enhance 同款检索/模型解析逻辑
|
||||
startedAt time.Time // 进程启动时间(算 uptime)
|
||||
}
|
||||
|
||||
// NewObserveHandler 构造函数(router.Setup 启动时创建一次)
|
||||
func NewObserveHandler(enhancer *service.EnhancerService) *ObserveHandler {
|
||||
return &ObserveHandler{
|
||||
enhancer: enhancer,
|
||||
startedAt: time.Now(),
|
||||
}
|
||||
}
|
||||
|
||||
// Logs 增量拉取内存日志
|
||||
//
|
||||
// GET /api/v1/agent/logs?since_id=0&limit=200
|
||||
// 面板首次加载 since_id=0 全量拉,之后带上一次返回的最大 ID 增量拉
|
||||
func (h *ObserveHandler) Logs(c *gin.Context) {
|
||||
sinceID, _ := strconv.ParseInt(c.DefaultQuery("since_id", "0"), 10, 64)
|
||||
limit, _ := strconv.Atoi(c.DefaultQuery("limit", "200"))
|
||||
entries := service.MemLog.List(sinceID, limit)
|
||||
// 返回最大 ID 作为下次轮询游标(无新日志时沿用请求值)
|
||||
lastID := sinceID
|
||||
if len(entries) > 0 {
|
||||
lastID = entries[len(entries)-1].ID
|
||||
}
|
||||
c.JSON(200, gin.H{
|
||||
"code": 200,
|
||||
"data": gin.H{
|
||||
"entries": entries,
|
||||
"last_id": lastID,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// System 进程运行时状态
|
||||
//
|
||||
// GET /api/v1/agent/system
|
||||
// 面板「系统状态」Tab 5s 轮询,全部从 runtime 取,零外部依赖
|
||||
func (h *ObserveHandler) System(c *gin.Context) {
|
||||
var ms runtime.MemStats
|
||||
runtime.ReadMemStats(&ms)
|
||||
|
||||
concUsed, concCap := service.GetConcurrency()
|
||||
runUsed, runCap := service.RunLogUsage()
|
||||
|
||||
c.JSON(200, gin.H{
|
||||
"code": 200,
|
||||
"data": gin.H{
|
||||
"started_at": h.startedAt.Unix(),
|
||||
"uptime_seconds": int64(time.Since(h.startedAt).Seconds()),
|
||||
"go_version": runtime.Version(),
|
||||
"goroutines": runtime.NumGoroutine(),
|
||||
// 内存三个关键指标:堆占用 / 累计分配 / 从 OS 拿到的总量
|
||||
"heap_alloc_mb": float64(ms.HeapAlloc) / 1024 / 1024,
|
||||
"sys_mb": float64(ms.Sys) / 1024 / 1024,
|
||||
"num_gc": ms.NumGC,
|
||||
// enhance 并发状态(信号量探针)
|
||||
"enhance_concurrency": gin.H{"used": concUsed, "capacity": concCap},
|
||||
// RunLog 缓冲占用
|
||||
"runlog_usage": gin.H{"used": runUsed, "capacity": runCap},
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// Config Agent 配置只读视图
|
||||
//
|
||||
// GET /api/v1/agent/config
|
||||
// 数据来自 agentcfg.Get()(xk_system_config 实时加载,30s 缓存)+ 守卫词表。
|
||||
// 面板只展示不修改——开关的写入口统一在 PHP 后台,避免双写
|
||||
func (h *ObserveHandler) Config(c *gin.Context) {
|
||||
cfg := agentcfg.Get()
|
||||
white, black := service.MedicalGuardKeywords()
|
||||
|
||||
c.JSON(200, gin.H{
|
||||
"code": 200,
|
||||
"data": gin.H{
|
||||
"react": gin.H{
|
||||
"enabled": cfg.ReAct.Enabled,
|
||||
"max_iterations": cfg.ReAct.MaxIterations,
|
||||
"planning_enabled": cfg.ReAct.PlanningEnabled,
|
||||
"reflection_enabled": cfg.ReAct.ReflectionEnabled,
|
||||
"json_repair_enabled": cfg.ReAct.JSONRepairEnabled,
|
||||
},
|
||||
"token_budget": gin.H{
|
||||
"enabled": cfg.TokenBudget.Enabled,
|
||||
"per_request": cfg.TokenBudget.PerRequest,
|
||||
"max_tokens_per_call": cfg.TokenBudget.MaxTokensPerCall,
|
||||
},
|
||||
"kb": gin.H{
|
||||
"source": cfg.KB.Source,
|
||||
"embedding_provider": cfg.KB.EmbeddingProvider,
|
||||
"top_k": cfg.KB.TopK,
|
||||
"search_mode": cfg.KB.SearchMode,
|
||||
},
|
||||
"medical_guard": gin.H{
|
||||
"enabled": cfg.MedicalGuard.Enabled,
|
||||
"whitelist": white,
|
||||
"blacklist": black,
|
||||
},
|
||||
"debug": gin.H{
|
||||
"log_request_body": cfg.Debug.LogRequestBody,
|
||||
},
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// GuardTest 医疗守卫测试台
|
||||
//
|
||||
// POST /api/v1/agent/guard-test body: {"text": "..."}
|
||||
// 用与 Enhance 入口完全相同的守卫逻辑跑一遍给定文本
|
||||
func (h *ObserveHandler) GuardTest(c *gin.Context) {
|
||||
var req struct {
|
||||
Text string `json:"text" binding:"required"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(400, gin.H{"code": 400, "message": "text 必填"})
|
||||
return
|
||||
}
|
||||
result := service.GuardTest(req.Text)
|
||||
c.JSON(200, gin.H{
|
||||
"code": 200,
|
||||
"data": gin.H{
|
||||
"passed": result.Passed,
|
||||
"reason": result.Reason,
|
||||
"hit_whitelist": result.HitWhitelist,
|
||||
"hit_blacklist": result.HitBlacklist,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// KBTest 知识库检索测试
|
||||
//
|
||||
// POST /api/v1/agent/kb-test body: {"query": "...", "top_k": 5}
|
||||
// 走 enhance 主流程同款检索路径(含 ai_kb_source 分流),
|
||||
// 回答「Agent 实际会检索到什么」
|
||||
func (h *ObserveHandler) KBTest(c *gin.Context) {
|
||||
var req struct {
|
||||
Query string `json:"query" binding:"required"`
|
||||
TopK int `json:"top_k"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(400, gin.H{"code": 400, "message": "query 必填"})
|
||||
return
|
||||
}
|
||||
result := h.enhancer.KBTest(c.Request.Context(), req.Query, req.TopK)
|
||||
c.JSON(200, gin.H{"code": 200, "data": result})
|
||||
}
|
||||
|
||||
// ModelTest 模型连通性测试
|
||||
//
|
||||
// POST /api/v1/models/test body: {"provider": "", "message": ""}
|
||||
// provider 留空 = 测当前生效配置;会真实调一次 LLM(max_tokens 64)
|
||||
func (h *ObserveHandler) ModelTest(c *gin.Context) {
|
||||
var req struct {
|
||||
Provider string `json:"provider"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
// body 可以整个为空(全默认),解析失败也不阻断
|
||||
_ = c.ShouldBindJSON(&req)
|
||||
result := h.enhancer.ModelTest(c.Request.Context(), req.Provider, req.Message)
|
||||
c.JSON(200, gin.H{"code": 200, "data": result})
|
||||
}
|
||||
252
internal/handler/prescription_handler.go
Normal file
252
internal/handler/prescription_handler.go
Normal file
@@ -0,0 +1,252 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"tcm-agent/internal/agent"
|
||||
"tcm-agent/internal/model/entity"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// ========================================================================
|
||||
// 处方 HTTP 接口处理器
|
||||
// ========================================================================
|
||||
// 职责链:
|
||||
// HTTP 请求 → 参数校验 → 调用 Agent → 规则校验 → 持久化 → 响应
|
||||
// ========================================================================
|
||||
|
||||
// PrescriptionHandler 处方接口处理器
|
||||
type PrescriptionHandler struct {
|
||||
agentRunner *agent.Runner
|
||||
rxAgent *agent.PrescriptionGenerator
|
||||
}
|
||||
|
||||
// NewPrescriptionHandler 创建处方处理器
|
||||
//
|
||||
// 参数:
|
||||
//
|
||||
// runner - Agent 引擎
|
||||
// scene - 场景名(为空则使用默认值 "prescription")
|
||||
func NewPrescriptionHandler(runner *agent.Runner, scene string) *PrescriptionHandler {
|
||||
return &PrescriptionHandler{
|
||||
agentRunner: runner,
|
||||
rxAgent: agent.NewPrescriptionGenerator(runner, scene),
|
||||
}
|
||||
}
|
||||
|
||||
// PrescriptionGenerateRequest 处方生成请求体(避免与 emr_handler.GenerateRequest 重名)
|
||||
type PrescriptionGenerateRequest struct {
|
||||
PatientID string `json:"patient_id" binding:"required"`
|
||||
EMRText string `json:"emr_text" binding:"required"`
|
||||
Diagnosis string `json:"diagnosis" binding:"required"`
|
||||
Age int `json:"age"`
|
||||
IsPregnant bool `json:"is_pregnant"`
|
||||
Allergies []string `json:"allergies"`
|
||||
}
|
||||
|
||||
// Generate 根据病历生成处方
|
||||
//
|
||||
// POST /api/v1/prescription/generate
|
||||
func (h *PrescriptionHandler) Generate(c *gin.Context) {
|
||||
// ===== ① 参数校验 =====
|
||||
var req PrescriptionGenerateRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"error": "参数错误",
|
||||
"detail": err.Error(),
|
||||
"example": `{"patient_id":"P001","emr_text":"...","diagnosis":"痰湿中阻证","age":45,"is_pregnant":false,"allergies":[]}`,
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
doctorID := c.GetString("user_id")
|
||||
|
||||
// ===== ②~⑤ 调用 Agent =====
|
||||
agentReq := &agent.PrescriptionRequest{
|
||||
PatientID: req.PatientID,
|
||||
EMRText: req.EMRText,
|
||||
Diagnosis: req.Diagnosis,
|
||||
Age: req.Age,
|
||||
IsPregnant: req.IsPregnant,
|
||||
Allergies: req.Allergies,
|
||||
}
|
||||
|
||||
resp, err := h.rxAgent.Generate(c.Request.Context(), agentReq)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"error": "处方生成失败",
|
||||
"detail": err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// ===== ⑥ 持久化 =====
|
||||
rx := &entity.Prescription{
|
||||
PatientID: req.PatientID,
|
||||
DoctorID: doctorID,
|
||||
SessionID: resp.SessionID,
|
||||
Draft: resp.Draft,
|
||||
Status: resp.Status,
|
||||
Blocked: resp.Blocked,
|
||||
Warnings: sliceToJSON(resp.Warnings),
|
||||
CreatedAt: time.Now(),
|
||||
UpdatedAt: time.Now(),
|
||||
}
|
||||
// dao.Prescription.Create(rx)
|
||||
|
||||
// ===== ⑦ 返回响应 =====
|
||||
httpStatus := http.StatusOK
|
||||
if resp.Blocked {
|
||||
httpStatus = http.StatusConflict // 409
|
||||
}
|
||||
|
||||
c.JSON(httpStatus, gin.H{
|
||||
"code": httpStatus,
|
||||
"message": getPrescriptionMessage(resp.Status),
|
||||
"data": gin.H{
|
||||
"session_id": resp.SessionID,
|
||||
"prescription": resp.Prescription,
|
||||
"draft": resp.Draft,
|
||||
"warnings": resp.Warnings,
|
||||
"blocked": resp.Blocked,
|
||||
"status": resp.Status,
|
||||
"rx_id": rx.ID,
|
||||
"next_action": getPrescriptionNextAction(resp.Status),
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// Validate 仅校验处方(不生成)
|
||||
//
|
||||
// POST /api/v1/prescription/validate
|
||||
func (h *PrescriptionHandler) Validate(c *gin.Context) {
|
||||
var req struct {
|
||||
PrescriptionText string `json:"prescription_text" binding:"required"`
|
||||
PatientID string `json:"patient_id"`
|
||||
IsPregnant bool `json:"is_pregnant"`
|
||||
Allergies []string `json:"allergies"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(400, gin.H{"error": "处方文本不能为空"})
|
||||
return
|
||||
}
|
||||
|
||||
agentReq := &agent.PrescriptionRequest{
|
||||
PatientID: req.PatientID,
|
||||
IsPregnant: req.IsPregnant,
|
||||
Allergies: req.Allergies,
|
||||
}
|
||||
warnings, blocked := h.rxAgent.Validate(c.Request.Context(), req.PrescriptionText, agentReq)
|
||||
|
||||
c.JSON(200, gin.H{
|
||||
"code": 200,
|
||||
"data": gin.H{
|
||||
"warnings": warnings,
|
||||
"blocked": blocked,
|
||||
"safe": !blocked && len(warnings) == 0,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// GetByID 查询处方详情
|
||||
//
|
||||
// GET /api/v1/prescription/:id
|
||||
func (h *PrescriptionHandler) GetByID(c *gin.Context) {
|
||||
idStr := c.Param("id")
|
||||
id, _ := strconv.ParseInt(idStr, 10, 64)
|
||||
|
||||
c.JSON(200, gin.H{
|
||||
"code": 200,
|
||||
"data": gin.H{
|
||||
"id": id,
|
||||
"note": "实际项目从数据库查询处方详情",
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// Approve 医生审核确认处方
|
||||
//
|
||||
// POST /api/v1/prescription/:id/approve
|
||||
//
|
||||
// 这是 Human-in-the-Loop 的关键环节:
|
||||
// AI 生成 → 规则校验 → 医生最终审核 → 生效
|
||||
func (h *PrescriptionHandler) Approve(c *gin.Context) {
|
||||
idStr := c.Param("id")
|
||||
id, _ := strconv.ParseInt(idStr, 10, 64)
|
||||
|
||||
var req struct {
|
||||
Approved bool `json:"approved"`
|
||||
DoctorNote string `json:"doctor_note"`
|
||||
ModifiedRx string `json:"modified_rx"`
|
||||
}
|
||||
c.ShouldBindJSON(&req)
|
||||
|
||||
doctorID := c.GetString("user_id")
|
||||
|
||||
logAudit("prescription_approve", doctorID, idStr, req.DoctorNote)
|
||||
|
||||
status := "approved"
|
||||
if !req.Approved {
|
||||
status = "rejected"
|
||||
}
|
||||
|
||||
c.JSON(200, gin.H{
|
||||
"code": 200,
|
||||
"message": "审核完成",
|
||||
"data": gin.H{
|
||||
"id": id,
|
||||
"status": status,
|
||||
"approved_by": doctorID,
|
||||
"doctor_note": req.DoctorNote,
|
||||
"modified_rx": req.ModifiedRx,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// getPrescriptionMessage 根据状态返回提示信息
|
||||
func getPrescriptionMessage(status string) string {
|
||||
switch status {
|
||||
case "success":
|
||||
return "处方生成成功,请医生审核"
|
||||
case "blocked":
|
||||
return "处方被安全规则拦截,已自动修正,请查看"
|
||||
case "need_review":
|
||||
return "处方存在警告,需医生重点关注"
|
||||
default:
|
||||
return "处方生成完成"
|
||||
}
|
||||
}
|
||||
|
||||
// getPrescriptionNextAction 下一步操作建议
|
||||
func getPrescriptionNextAction(status string) string {
|
||||
switch status {
|
||||
case "success":
|
||||
return "请医生审核处方并确认"
|
||||
case "blocked":
|
||||
return "处方已被拦截修正,请医生重新审阅"
|
||||
case "need_review":
|
||||
return "存在安全警告,请医生评估后决定"
|
||||
default:
|
||||
return "请检查输入信息"
|
||||
}
|
||||
}
|
||||
|
||||
// sliceToJSON 将字符串切片转为 JSON 字符串(用于数据库存储)
|
||||
func sliceToJSON(s []string) string {
|
||||
if len(s) == 0 {
|
||||
return "[]"
|
||||
}
|
||||
// 简单拼接,生产环境用 json.Marshal
|
||||
result := "["
|
||||
for i, v := range s {
|
||||
if i > 0 {
|
||||
result += ","
|
||||
}
|
||||
result += `"` + v + `"`
|
||||
}
|
||||
result += "]"
|
||||
return result
|
||||
}
|
||||
Reference in New Issue
Block a user