import { Router } from 'express' import { authMiddleware, AuthRequest } from '../middleware/auth' import { chat, chatStream, reviewContract, matchCase, predictRisks } from '../services/ai.service' import { seedKnowledgeBase, addKnowledge, searchKnowledge, ensureRAGTable } from '../services/rag.service' import prisma from '../lib/prisma' import { z } from 'zod' const router = Router() const PLAN_LIMITS: Record = { FREE: { chat: 10, review: 3, case: 3 }, PRO: { chat: 100, review: 20, case: 20 }, ENTERPRISE: { chat: 0, review: 0, case: 0 }, } async function checkUsageLimit(orgId: string, type: 'chat' | 'review' | 'case'): Promise { const org = await prisma.organization.findUnique({ where: { id: orgId } }) if (!org) return const limits = PLAN_LIMITS[org.plan] || PLAN_LIMITS.FREE const limit = limits[type] if (limit === 0) return const now = new Date() const monthStart = new Date(now.getFullYear(), now.getMonth(), 1) const count = await prisma.auditLog.count({ where: { orgId, action: `AI_${type.toUpperCase()}`, createdAt: { gte: monthStart }, }, }) if (count >= limit) { throw { code: 'USAGE_LIMIT', message: `本月 AI${type === 'chat' ? '问答' : type === 'review' ? '合同审查' : '案例匹配'}次数已达上限(${limit}次),请升级套餐` } } } async function recordUsage(orgId: string, userId: string, type: 'chat' | 'review' | 'case'): Promise { const month = new Date().toISOString().slice(0, 7) await prisma.auditLog.create({ data: { orgId, userId, action: `AI_${type.toUpperCase()}`, entity: 'AI', entityId: null, detail: { month, type } as any, ip: '', }, }) } async function buildOrgContext(orgId: string): Promise { const [employees, risks] = await Promise.all([ prisma.employee.findMany({ where: { orgId, status: 'ACTIVE' }, include: { contracts: { orderBy: { createdAt: 'desc' }, take: 1 } }, }), prisma.riskItem.findMany({ where: { orgId, status: 'PENDING' }, include: { employee: true }, }), ]) const now = new Date() const empSummary = employees.map((e) => { const contract = e.contracts[0] const daysToExpire = contract?.endDate ? Math.floor((new Date(contract.endDate).getTime() - now.getTime()) / (1000 * 60 * 60 * 24)) : null const specialStatus: string[] = [] if (e.isPregnant) specialStatus.push('孕期/哺乳期') if (e.isInMedicalPeriod) specialStatus.push('医疗期') if (e.isWorkInjured) specialStatus.push('工伤') return `- ${e.name}(${e.department}),入职${e.hireDate.toISOString().slice(0, 10)},${contract ? `合同:${contract.contractType},${contract.endDate ? `到期${contract.endDate.toISOString().slice(0, 10)}(剩余${daysToExpire}天)` : '无固定期限'}` : '未签合同'}${specialStatus.length > 0 ? `,特殊状态:${specialStatus.join('/')}` : ''}` }).join('\n') const riskSummary = risks.map((r) => `- ${r.title}(${r.level}):${r.description || '无详细描述'}`).join('\n') return `员工列表(${employees.length}人): ${empSummary} 当前风险项(${risks.length}项): ${riskSummary}` } router.post('/chat', authMiddleware, async (req: AuthRequest, res, next) => { try { const { messages } = req.body as { messages: { role: 'user' | 'assistant'; content: string }[] } if (!messages || !Array.isArray(messages)) { return res.status(400).json({ success: false, error: { code: 'BAD_REQUEST', message: '缺少 messages 参数' } }) } await checkUsageLimit(req.user!.orgId, 'chat') const orgContext = await buildOrgContext(req.user!.orgId) const reply = await chat(messages, orgContext) await recordUsage(req.user!.orgId, req.user!.id, 'chat') res.json({ success: true, data: { reply } }) } catch (err) { next(err) } }) router.post('/chat-stream', authMiddleware, async (req: AuthRequest, res, next) => { try { const { messages } = req.body as { messages: { role: 'user' | 'assistant'; content: string }[] } if (!messages || !Array.isArray(messages)) { return res.status(400).json({ success: false, error: { code: 'BAD_REQUEST', message: '缺少 messages 参数' } }) } await checkUsageLimit(req.user!.orgId, 'chat') const orgContext = await buildOrgContext(req.user!.orgId) res.setHeader('Content-Type', 'text/event-stream') res.setHeader('Cache-Control', 'no-cache') res.setHeader('Connection', 'keep-alive') let usageRecorded = false try { for await (const delta of chatStream(messages, orgContext)) { res.write(`data: ${JSON.stringify({ delta })}\n\n`) } res.write('data: [DONE]\n\n') } finally { if (!usageRecorded) { await recordUsage(req.user!.orgId, req.user!.id, 'chat') usageRecorded = true } } res.end() } catch (err) { if (!res.headersSent) next(err) else res.end() } }) router.post('/review', authMiddleware, async (req: AuthRequest, res, next) => { try { const { contractText } = req.body as { contractText: string } if (!contractText) { return res.status(400).json({ success: false, error: { code: 'BAD_REQUEST', message: '缺少合同文本' } }) } await checkUsageLimit(req.user!.orgId, 'review') const result = await reviewContract(contractText) await recordUsage(req.user!.orgId, req.user!.id, 'review') res.json({ success: true, data: { text: result.text, structured: result.structured } }) } catch (err) { next(err) } }) router.post('/match-case', authMiddleware, async (req: AuthRequest, res, next) => { try { const { scenario } = req.body as { scenario: string } if (!scenario) { return res.status(400).json({ success: false, error: { code: 'BAD_REQUEST', message: '缺少争议情形描述' } }) } await checkUsageLimit(req.user!.orgId, 'case') const result = await matchCase(scenario) await recordUsage(req.user!.orgId, req.user!.id, 'case') res.json({ success: true, data: { result } }) } catch (err) { next(err) } }) // 案例匹配结果转待办(RiskItem) router.post('/case-to-todo', authMiddleware, async (req: AuthRequest, res, next) => { try { const schema = z.object({ employeeId: z.string().min(1), title: z.string().min(1), description: z.string().min(1), level: z.enum(['HIGH', 'MEDIUM', 'LOW']).default('MEDIUM'), type: z.enum(['CONTRACT', 'SALARY', 'TERMINATION', 'MONTHLY', 'ONBOARDING']).default('TERMINATION'), }) const data = schema.parse(req.body) const risk = await prisma.riskItem.create({ data: { orgId: req.user!.orgId, employeeId: data.employeeId, title: data.title, description: data.description, level: data.level, type: data.type, status: 'PENDING', }, }) res.json({ success: true, data: risk }) } catch (err) { next(err) } }) router.get('/predict', authMiddleware, async (req: AuthRequest, res, next) => { try { const department = req.query.department as string const employeeId = req.query.employeeId as string const riskType = req.query.riskType as string let orgContext = await buildOrgContext(req.user!.orgId) if (employeeId) { const emp = await prisma.employee.findFirst({ where: { id: employeeId, orgId: req.user!.orgId }, include: { contracts: { orderBy: { createdAt: 'desc' }, take: 1 } } }) if (emp) { const contract = emp.contracts[0] orgContext = `员工详情: - 姓名:${emp.name} - 部门:${emp.department} - 入职日期:${emp.hireDate.toISOString().slice(0, 10)} - 状态:${emp.status} - 特殊状态:${emp.isPregnant ? '孕期/哺乳期 ' : ''}${emp.isInMedicalPeriod ? '医疗期 ' : ''}${emp.isWorkInjured ? '工伤' : '无'} - 合同:${contract ? `${contract.contractType},${contract.startDate.toISOString().slice(0, 10)}至${contract.endDate ? contract.endDate.toISOString().slice(0, 10) : '无固定期限'}` : '未签合同'}\n${orgContext}` } } else if (department) { const employees = await prisma.employee.findMany({ where: { orgId: req.user!.orgId, department, status: 'ACTIVE' }, include: { contracts: { orderBy: { createdAt: 'desc' }, take: 1 } } }) const empSummary = employees.map(e => `- ${e.name},入职${e.hireDate.toISOString().slice(0, 10)},${e.contracts[0] ? e.contracts[0].contractType : '未签合同'}`).join('\n') orgContext = `部门【${department}】员工列表(${employees.length}人):\n${empSummary}\n\n${orgContext}` } if (riskType && riskType !== 'all') { orgContext = `请重点关注【${riskType === 'contract' ? '合同' : riskType === 'salary' ? '薪酬' : riskType === 'termination' ? '解聘' : riskType}】类风险。\n\n${orgContext}` } const result = await predictRisks(orgContext) res.json({ success: true, data: { result } }) } catch (err) { next(err) } }) // ========== AI 会话历史 ========== router.get('/conversations', authMiddleware, async (req: AuthRequest, res, next) => { try { const conversations = await prisma.aIConversation.findMany({ where: { orgId: req.user!.orgId, userId: req.user!.id }, orderBy: { updatedAt: 'desc' }, take: 50, select: { id: true, title: true, createdAt: true, updatedAt: true }, }) res.json({ success: true, data: conversations }) } catch (err) { next(err) } }) router.get('/conversations/:id', authMiddleware, async (req: AuthRequest, res, next) => { try { const conv = await prisma.aIConversation.findFirst({ where: { id: req.params.id, orgId: req.user!.orgId, userId: req.user!.id }, }) if (!conv) return res.status(404).json({ success: false, error: { code: 'NOT_FOUND', message: '会话不存在' } }) res.json({ success: true, data: conv }) } catch (err) { next(err) } }) router.post('/conversations', authMiddleware, async (req: AuthRequest, res, next) => { try { const { title, messages } = req.body as { title?: string; messages: any[] } const conv = await prisma.aIConversation.create({ data: { orgId: req.user!.orgId, userId: req.user!.id, title: title || (messages.find(m => m.role === 'user')?.content.slice(0, 30) || '新对话'), messages: messages || [], }, }) res.json({ success: true, data: conv }) } catch (err) { next(err) } }) router.put('/conversations/:id', authMiddleware, async (req: AuthRequest, res, next) => { try { const { title, messages } = req.body as { title?: string; messages?: any[] } const conv = await prisma.aIConversation.updateMany({ where: { id: req.params.id, orgId: req.user!.orgId, userId: req.user!.id }, data: { ...(title ? { title } : {}), ...(messages ? { messages } : {}), }, }) if (conv.count === 0) return res.status(404).json({ success: false, error: { code: 'NOT_FOUND', message: '会话不存在' } }) res.json({ success: true }) } catch (err) { next(err) } }) router.delete('/conversations/:id', authMiddleware, async (req: AuthRequest, res, next) => { try { const conv = await prisma.aIConversation.deleteMany({ where: { id: req.params.id, orgId: req.user!.orgId, userId: req.user!.id }, }) if (conv.count === 0) return res.status(404).json({ success: false, error: { code: 'NOT_FOUND', message: '会话不存在' } }) res.json({ success: true }) } catch (err) { next(err) } }) // ========== AI 审查记录保存到员工档案 ========== router.post('/review/save', authMiddleware, async (req: AuthRequest, res, next) => { try { const schema = z.object({ employeeId: z.string(), type: z.enum(['REVIEW', 'CASE']), input: z.string(), result: z.string(), }) const data = schema.parse(req.body) const record = await prisma.aIReviewRecord.create({ data: { orgId: req.user!.orgId, employeeId: data.employeeId, type: data.type, input: data.input, result: data.result, createdBy: req.user!.id, }, }) res.json({ success: true, data: record }) } catch (err) { next(err) } }) router.get('/review/employee/:employeeId', authMiddleware, async (req: AuthRequest, res, next) => { try { const records = await prisma.aIReviewRecord.findMany({ where: { orgId: req.user!.orgId, employeeId: req.params.employeeId }, orderBy: { createdAt: 'desc' }, take: 20, }) res.json({ success: true, data: records }) } catch (err) { next(err) } }) // RAG 知识库管理 router.post('/rag/seed', authMiddleware, async (_req: AuthRequest, res, next) => { try { await seedKnowledgeBase() res.json({ success: true, data: { message: '知识库初始化完成' } }) } catch (err) { next(err) } }) router.post('/rag/add', authMiddleware, async (req: AuthRequest, res, next) => { try { const { title, content, source, category } = req.body if (!title || !content) { return res.status(400).json({ success: false, error: { code: 'BAD_REQUEST', message: '缺少 title 或 content' } }) } const result = await addKnowledge(title, content, source || '自定义', category || '其他') res.json({ success: true, data: result }) } catch (err) { next(err) } }) router.post('/rag/search', authMiddleware, async (req: AuthRequest, res, next) => { try { const { query, topK } = req.body if (!query) { return res.status(400).json({ success: false, error: { code: 'BAD_REQUEST', message: '缺少 query' } }) } const results = await searchKnowledge(query, topK || 5) res.json({ success: true, data: { results } }) } catch (err) { next(err) } }) // 知识库列表 router.get('/rag/list', authMiddleware, async (req: AuthRequest, res, next) => { try { await ensureRAGTable() const category = req.query.category as string | undefined const items = category ? await prisma.$queryRaw`SELECT id, title, content, source, category, created_at FROM rag_knowledge WHERE category = ${category} ORDER BY created_at DESC LIMIT 200` as any[] : await prisma.$queryRaw`SELECT id, title, content, source, category, created_at FROM rag_knowledge ORDER BY created_at DESC LIMIT 200` as any[] res.json({ success: true, data: items }) } catch (err) { next(err) } }) // 删除知识条目 router.delete('/rag/:id', authMiddleware, async (req: AuthRequest, res, next) => { try { await ensureRAGTable() await prisma.$executeRaw`DELETE FROM rag_knowledge WHERE id = ${req.params.id}` res.json({ success: true }) } catch (err) { next(err) } }) export default router