Files
gzth/yuthon_ai_agent/models/ai_agent.py
T
2026-05-28 15:00:06 +08:00

642 lines
28 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- coding: utf-8 -*-
import json
import logging
from datetime import datetime
import pytz
import requests
from odoo import models, fields, api
from odoo.exceptions import AccessError
_logger = logging.getLogger(__name__)
# 单次对话最多允许的工具调用轮次(防止死循环)
MAX_TOOL_ITERATIONS = 6
# search_read 默认/最大返回条数
DEFAULT_LIMIT = 20
MAX_LIMIT = 200
class AiProvider(models.Model):
"""AI 服务配置"""
_name = 'ai.provider'
_description = 'AI 服务商'
create_date = fields.Datetime(string='创建时间', readonly=True)
write_date = fields.Datetime(string='最后更新', readonly=True)
name = fields.Char(string='名称', required=True)
provider_type = fields.Selection([
('openai', 'OpenAI'),
('azure', 'Azure OpenAI'),
('local', '本地模型 (Ollama/LocalAI)'),
('custom', '自定义'),
], string='类型', default='openai', required=True)
api_base = fields.Char(
string='API 地址',
default='https://api.openai.com/v1',
help='API Base URL,包含 /v1'
)
api_key = fields.Char(string='API Key')
model_name = fields.Char(
string='模型名称',
default='gpt-4o-mini',
help='如 gpt-4o-mini, qwen-plus, llama3 等'
)
max_tokens = fields.Integer(string='最大 Token', default=4096)
temperature = fields.Float(string='温度', default=0.7)
timeout = fields.Integer(
string='超时(秒)',
default=20,
help='AI 单次请求的最大等待秒数,超过后返回超时错误。默认 20 秒。'
)
active = fields.Boolean(string='启用', default=True)
enable_tools = fields.Boolean(
string='启用工具调用',
default=True,
help='开启后 AI 可以查询/创建/修改 Odoo 数据(Function Calling)。要求模型支持 tools 参数(如 gpt-4o, qwen-plus, deepseek-chat 等)。'
)
allow_write = fields.Boolean(
string='允许写入',
default=False,
help='允许 AI 调用 create/write/unlink 工具。建议仅在受信任环境开启。'
)
# 规则配置
rule_ids = fields.Many2many(
'ai.rule', 'ai_rule_provider_rel', 'provider_id', 'rule_id',
string='关联规则',
help='为该服务商绑定的 AI 规则(回复指令、限制规则、技能接口)',
)
# Token 统计(关联到本服务商)
provider_total_tokens = fields.Integer(
string='总 Token', compute='_compute_provider_token_stats'
)
provider_call_count = fields.Integer(
string='调用次数', compute='_compute_provider_token_stats'
)
def _compute_provider_token_stats(self):
"""统计该服务商的 Token 总量和调用次数"""
for rec in self:
usage_data = self.env['ai.token.usage'].sudo().read_group(
[('provider_id', '=', rec.id)],
['total_tokens:sum'],
[],
)
rec.provider_total_tokens = usage_data[0]['total_tokens'] if usage_data else 0
rec.provider_call_count = self.env['ai.token.usage'].sudo().search_count(
[('provider_id', '=', rec.id)]
)
def action_view_provider_token_usage(self):
"""打开该服务商的 Token 用量明细"""
self.ensure_one()
return {
'type': 'ir.actions.act_window',
'name': f'{self.name} - Token 用量',
'res_model': 'ai.token.usage',
'view_mode': 'tree',
'domain': [('provider_id', '=', self.id)],
'context': {'search_default_group_user': 1},
}
def get_applicable_rules(self, user=None):
"""获取对当前用户生效的规则(按优先级排序)"""
self.ensure_one()
user = user or self.env.user
rules = self.rule_ids.filtered(lambda r: r.active)
# 过滤用户组:规则未配用户组 = 对所有人生效,配了则检查用户是否在组中
applicable = rules.filtered(
lambda r: not r.group_ids or (r.group_ids & user.groups_id)
)
return applicable.sorted('priority')
class AiConversation(models.Model):
"""对话会话"""
_name = 'ai.conversation'
_description = 'AI 对话'
_order = 'create_date desc'
create_date = fields.Datetime(string='创建时间', readonly=True)
write_date = fields.Datetime(string='最后更新', readonly=True)
name = fields.Char(string='标题', required=True, default='新对话')
provider_id = fields.Many2one(
'ai.provider', string='AI 服务',
required=True,
default=lambda self: self.env['ai.provider'].search([('active', '=', True)], limit=1)
)
message_ids = fields.One2many(
'ai.message', 'conversation_id', string='消息'
)
message_count = fields.Integer(
string='消息数', compute='_compute_message_count', store=True
)
create_uid = fields.Many2one('res.users', string='用户', default=lambda self: self.env.uid)
# 从 provider 映射过来的可读字段,方便在对话页直观查看当前助手配置
provider_model_name = fields.Char(related='provider_id.model_name', string='模型名称', readonly=True)
provider_timeout = fields.Integer(related='provider_id.timeout', string='超时(秒)', readonly=True)
provider_enable_tools = fields.Boolean(related='provider_id.enable_tools', string='工具调用', readonly=True)
provider_allow_write = fields.Boolean(related='provider_id.allow_write', string='允许写入', readonly=True)
@api.depends('message_ids')
def _compute_message_count(self):
for rec in self:
rec.message_count = len(rec.message_ids)
def action_clear(self):
"""清空对话消息"""
self.ensure_one()
self.message_ids.unlink()
def open_chat(self):
"""打开聊天界面"""
self.ensure_one()
return {
'type': 'ir.actions.client',
'tag': 'yuthon_ai_agent.chat',
'name': self.name,
'params': {'res_id': self.id},
}
# ---------------------------------------------------------
# Function Calling: 工具定义 & 执行
# ---------------------------------------------------------
def _get_tools_schema(self):
"""返回 OpenAI 兼容的工具列表(JSON Schema)。"""
tools = [
{
'type': 'function',
'function': {
'name': 'odoo_search_read',
'description': '搜索并读取 Odoo 数据库记录。常见模型:res.partner(联系人)、product.product(产品)、sale.order(销售订单)、account.move(发票)、stock.picking(出入库)、hr.employee(员工)、res.users(用户)。',
'parameters': {
'type': 'object',
'properties': {
'model': {'type': 'string', 'description': '模型技术名称,如 res.partner'},
'domain': {'type': 'array', 'description': "Odoo domain,如 [['name','ilike','张三']],默认 []", 'items': {}},
'fields': {'type': 'array', 'items': {'type': 'string'}, 'description': '要读取的字段列表,留空则只读 id/display_name'},
'limit': {'type': 'integer', 'description': f'返回条数,默认 {DEFAULT_LIMIT},最大 {MAX_LIMIT}'},
'order': {'type': 'string', 'description': '排序,如 "create_date desc"'},
},
'required': ['model'],
},
},
},
{
'type': 'function',
'function': {
'name': 'odoo_fields_get',
'description': '获取模型字段定义(类型、label、关系字段指向的模型)。当不清楚字段名时先调用此工具。',
'parameters': {
'type': 'object',
'properties': {
'model': {'type': 'string'},
'allfields': {'type': 'array', 'items': {'type': 'string'}, 'description': '可选:只返回指定字段'},
},
'required': ['model'],
},
},
},
{
'type': 'function',
'function': {
'name': 'odoo_search_count',
'description': '统计满足条件的记录数。',
'parameters': {
'type': 'object',
'properties': {
'model': {'type': 'string'},
'domain': {'type': 'array', 'items': {}},
},
'required': ['model'],
},
},
},
{
'type': 'function',
'function': {
'name': 'save_memory',
'description': (
'将用户的偏好、习惯、工作内容等信息保存为长期记忆。'
'当对话中发现值得记住的用户信息时主动调用。'
'分类:habit(习惯偏好)、work(工作内容)、knowledge(专业知识)、'
'style(沟通风格)、personal(个人信息)、other(其他)。'
),
'parameters': {
'type': 'object',
'properties': {
'category': {
'type': 'string',
'enum': ['habit', 'work', 'knowledge', 'style', 'personal', 'other'],
'description': '记忆类别',
},
'content': {
'type': 'string',
'description': '要记住的内容,简洁一句话概括',
},
},
'required': ['category', 'content'],
},
},
},
]
if self.provider_id.allow_write:
tools += [
{
'type': 'function',
'function': {
'name': 'odoo_create',
'description': '在 Odoo 中创建一条记录,返回新记录 id 与 display_name。',
'parameters': {
'type': 'object',
'properties': {
'model': {'type': 'string'},
'values': {'type': 'object', 'description': '字段名 -> 值。Many2one 字段传 id;Many2many/One2many 用 [(6,0,[ids])] 等命令格式'},
},
'required': ['model', 'values'],
},
},
},
{
'type': 'function',
'function': {
'name': 'odoo_write',
'description': '更新已有记录。',
'parameters': {
'type': 'object',
'properties': {
'model': {'type': 'string'},
'ids': {'type': 'array', 'items': {'type': 'integer'}},
'values': {'type': 'object'},
},
'required': ['model', 'ids', 'values'],
},
},
},
{
'type': 'function',
'function': {
'name': 'odoo_unlink',
'description': '删除记录(请谨慎,先与用户确认)。',
'parameters': {
'type': 'object',
'properties': {
'model': {'type': 'string'},
'ids': {'type': 'array', 'items': {'type': 'integer'}},
},
'required': ['model', 'ids'],
},
},
},
]
return tools
def _execute_tool(self, name, args):
"""执行单个工具调用,返回可 JSON 序列化的结果字典。
所有 ORM 操作都走当前用户权限(self.env),自动遵循 ir.rule / ir.model.access。
"""
try:
# save_memory 工具单独处理(不需要 model 参数)
if name == 'save_memory':
profile = self.env['ai.user.profile'].get_or_create()
category = args.get('category', 'other')
content = args.get('content', '')
if not content:
return {'error': '记忆内容不能为空'}
# 检查是否已有相似记忆,避免重复
existing = self.env['ai.user.memory'].search([
('profile_id', '=', profile.id),
('content', '=', content),
('active', '=', True),
], limit=1)
if existing:
return {'success': True, 'message': '该记忆已存在,跳过'}
self.env['ai.user.memory'].create({
'profile_id': profile.id,
'category': category,
'content': content,
'source': 'auto',
})
return {'success': True, 'message': f'已记住:{content}'}
model_name = args.get('model')
if not model_name or model_name not in self.env:
return {'error': f'未知模型:{model_name}'}
Model = self.env[model_name]
if name == 'odoo_search_read':
domain = args.get('domain') or []
fields_ = args.get('fields') or []
limit = min(int(args.get('limit') or DEFAULT_LIMIT), MAX_LIMIT)
order = args.get('order') or None
records = Model.search_read(domain, fields_, limit=limit, order=order)
# 将 UTC 时间转换为用户时区
records = self._convert_records_timezone(records, model_name, fields_)
return {'success': True, 'count': len(records), 'records': records}
if name == 'odoo_search_count':
domain = args.get('domain') or []
return {'success': True, 'count': Model.search_count(domain)}
if name == 'odoo_fields_get':
allfields = args.get('allfields') or []
fg = Model.fields_get(allfields, attributes=['type', 'string', 'relation', 'required', 'readonly', 'selection'])
return {'success': True, 'fields': fg}
if name == 'odoo_create':
if not self.provider_id.allow_write:
return {'error': '当前 Provider 未启用写入权限'}
rec = Model.create(args.get('values') or {})
return {'success': True, 'id': rec.id, 'display_name': rec.display_name}
if name == 'odoo_write':
if not self.provider_id.allow_write:
return {'error': '当前 Provider 未启用写入权限'}
ids = args.get('ids') or []
Model.browse(ids).write(args.get('values') or {})
return {'success': True, 'updated_ids': ids}
if name == 'odoo_unlink':
if not self.provider_id.allow_write:
return {'error': '当前 Provider 未启用写入权限'}
ids = args.get('ids') or []
Model.browse(ids).unlink()
return {'success': True, 'deleted_ids': ids}
return {'error': f'未知工具:{name}'}
except Exception as e:
_logger.exception('Tool %s failed: %s', name, e)
return {'error': f'{type(e).__name__}: {e}'}
def _execute_skill_tool(self, name, args):
"""尝试执行技能规则工具,找不到返回 None"""
provider = self.provider_id
if not provider:
return None
skill_rule = self.env['ai.rule'].sudo().search([
('rule_type', '=', 'skill'),
('skill_name', '=', name),
('active', '=', True),
('provider_ids', 'in', [provider.id]),
], limit=1)
if not skill_rule:
# 也搜索不限服务商的全局技能
skill_rule = self.env['ai.rule'].sudo().search([
('rule_type', '=', 'skill'),
('skill_name', '=', name),
('active', '=', True),
('provider_ids', '=', False),
], limit=1)
if skill_rule:
return skill_rule.execute_skill(args)
return None
def _call_chat_completion(self, provider, messages, tools=None):
"""封装一次 HTTP 调用,返回 (assistant_message_dict, error_str)。
同时记录 Token 使用量。
"""
payload = {
'model': provider.model_name,
'messages': messages,
'max_tokens': provider.max_tokens,
'temperature': provider.temperature,
}
if tools:
payload['tools'] = tools
payload['tool_choice'] = 'auto'
try:
resp = requests.post(
f"{provider.api_base}/chat/completions",
headers={
'Authorization': f'Bearer {provider.api_key}',
'Content-Type': 'application/json',
},
json=payload,
timeout=max(provider.timeout or 20, 1),
)
resp.raise_for_status()
data = resp.json()
# 记录 Token 使用量
usage = data.get('usage') or {}
if usage:
self.env['ai.token.usage'].sudo().create({
'user_id': self.env.uid,
'conversation_id': self.id,
'provider_id': provider.id,
'model_name': provider.model_name,
'prompt_tokens': usage.get('prompt_tokens', 0),
'completion_tokens': usage.get('completion_tokens', 0),
'total_tokens': usage.get('total_tokens', 0),
})
return data['choices'][0]['message'], None
except requests.exceptions.Timeout:
return None, f'错误:AI 响应超时(超过 {provider.timeout or 20} 秒),请重试或在服务商配置中调高超时。'
except requests.exceptions.RequestException as e:
return None, f'错误:API 调用失败 — {e}'
except (KeyError, IndexError, ValueError) as e:
return None, f'错误:解析 AI 响应失败 — {e}'
def _convert_records_timezone(self, records, model_name, fields_list):
"""将查询结果中的 datetime 字段从 UTC 转换为用户时区"""
if not records:
return records
# 获取用户时区
user_tz_name = self.env.user.tz or 'Asia/Shanghai'
try:
user_tz = pytz.timezone(user_tz_name)
except pytz.exceptions.UnknownTimeZoneError:
user_tz = pytz.timezone('Asia/Shanghai')
utc_tz = pytz.UTC
# 获取模型的 datetime 字段列表
try:
model_fields = self.env[model_name].fields_get(attributes=['type'])
except Exception:
return records
datetime_fields = set()
for fname, finfo in model_fields.items():
if finfo.get('type') == 'datetime':
datetime_fields.add(fname)
# 如果指定了fields_list,只处理请求的字段
if fields_list:
datetime_fields = datetime_fields & set(fields_list)
if not datetime_fields:
return records
# 转换每条记录中的 datetime 字段
for record in records:
for fname in datetime_fields:
val = record.get(fname)
if not val:
continue
try:
if isinstance(val, str):
dt_utc = datetime.strptime(val, '%Y-%m-%d %H:%M:%S')
elif isinstance(val, datetime):
dt_utc = val
else:
continue
dt_utc = utc_tz.localize(dt_utc)
dt_local = dt_utc.astimezone(user_tz)
record[fname] = dt_local.strftime('%Y-%m-%d %H:%M:%S')
except (ValueError, TypeError):
continue
return records
# ---------------------------------------------------------
# 主入口
# ---------------------------------------------------------
def send_message(self, content):
"""发送用户消息并获取 AI 回复(支持工具调用)。"""
self.ensure_one()
# 权限校验:只有 AI Agent 用户组才能使用
if not self.env.user.has_group('yuthon_ai_agent.group_ai_agent_user'):
return {'error': '您没有使用 AI 助手的权限,请联系管理员开通。'}
provider = self.provider_id
if not provider or not provider.active:
return {'error': 'AI 服务未配置或已停用。'}
# 1. 保存用户消息
self.env['ai.message'].create({
'conversation_id': self.id,
'role': 'user',
'content': content,
})
# 2. 构建发送给模型的 messages
# 获取用户画像
profile = self.env['ai.user.profile'].get_or_create()
user_context = profile.get_system_context()
system_msg = (
'你是一个 Odoo 17 ERP 系统的智能助手,可以调用工具直接查询/操作数据库。\n'
'常见模型:res.partner(联系人)、product.product/product.template(产品)、'
'sale.order(销售订单)、purchase.order(采购单)、account.move(发票/账单)、'
'stock.picking(出入库)、hr.employee(员工)、res.users(用户)、res.company(公司)。\n'
'当用户提到“查/搜/统计/创建/修改/删除”等数据相关需求时,必须主动调用相应工具,'
'不要凭空编造结果。不确定字段名时先调用 odoo_fields_get。\n'
'所有回答使用中文,简洁清晰。\n'
'【格式要求】回复不要使用 Markdown 语法(禁止使用 | 表格、** 加粗、# 标题等)。\n'
'列表数据请用编号+冒号格式,每条一行,字段用空格分隔,例如:\n'
'1. 张三 研发部 2025-03-01\n'
'2. 李四 财务部 2025-03-15\n'
'多字段时用“ / ”分隔,保持简洁可读。\n'
'查询结果中的时间已经转换为北京时间(UTC+8),直接展示即可,无需再做时区转换。\n\n'
'【重要】你具有记忆能力。当对话中发现用户的新偏好、工作内容、习惯等值得长期记住的信息时,'
'主动调用 save_memory 工具保存。下次对话你会自动获取这些记忆。\n'
'称呼用户时使用其设定的称呼,不要说“您”、“亲”等通用称谓。\n'
)
# —— 注入 AI 规则 ——
applicable_rules = provider.get_applicable_rules()
instruction_rules = applicable_rules.filtered(lambda r: r.rule_type == 'instruction')
restriction_rules = applicable_rules.filtered(lambda r: r.rule_type == 'restriction')
skill_rules = applicable_rules.filtered(lambda r: r.rule_type == 'skill' and r.skill_name)
if instruction_rules:
system_msg += '\n--- 回复指令(必须遵守)---\n'
for rule in instruction_rules:
system_msg += f'• {rule.name}:{rule.content}\n'
if restriction_rules:
system_msg += '\n--- 限制规则(禁止回答以下话题)---\n'
for rule in restriction_rules:
system_msg += f'• {rule.name}:{rule.content}\n'
system_msg += '如果用户询问以上禁止话题,礼貌拒绝并说明无法回答。\n'
if user_context:
system_msg += f'\n--- 当前用户信息 ---\n{user_context}\n'
messages = [{'role': 'system', 'content': system_msg}]
for msg in self.message_ids:
if msg.role in ('user', 'assistant') and msg.content:
messages.append({'role': msg.role, 'content': msg.content})
tools = self._get_tools_schema() if provider.enable_tools else None
# 把技能规则注册为额外的工具
if skill_rules:
if tools is None:
tools = []
for rule in skill_rules:
tools.append(rule.get_tool_schema())
# 3. 多轮 tool_calls 循环
final_content = ''
for _ in range(MAX_TOOL_ITERATIONS):
assistant_msg, err = self._call_chat_completion(provider, messages, tools)
if err:
final_content = err
break
tool_calls = assistant_msg.get('tool_calls')
if not tool_calls:
final_content = assistant_msg.get('content') or ''
break
# 把 assistant 的工具调用消息追加到上下文
messages.append({
'role': 'assistant',
'content': assistant_msg.get('content') or '',
'tool_calls': tool_calls,
})
# 依次执行工具,写回 role=tool 消息
for tc in tool_calls:
fn = tc.get('function', {})
fname = fn.get('name', '')
try:
fargs = json.loads(fn.get('arguments') or '{}')
except json.JSONDecodeError:
fargs = {}
# 先尝试技能接口工具,找不到再走内置工具
result = self._execute_skill_tool(fname, fargs)
if result is None:
result = self._execute_tool(fname, fargs)
messages.append({
'role': 'tool',
'tool_call_id': tc.get('id'),
'content': json.dumps(result, ensure_ascii=False, default=str),
})
else:
final_content = final_content or '错误:工具调用次数超限,请重新提问。'
# 4. 保存最终 AI 回复
self.env['ai.message'].create({
'conversation_id': self.id,
'role': 'assistant',
'content': final_content,
})
# 5. 首次发送自动更新标题
if self.name == '新对话' and len(self.message_ids) >= 2:
self.name = content[:30] + ('...' if len(content) > 30 else '')
return final_content
class AiMessage(models.Model):
"""对话消息"""
_name = 'ai.message'
_description = 'AI 消息'
_order = 'create_date'
conversation_id = fields.Many2one(
'ai.conversation', string='对话', required=True, ondelete='cascade'
)
role = fields.Selection([
('user', '用户'),
('assistant', '助手'),
('system', '系统'),
], string='角色', required=True)
content = fields.Text(string='内容', required=True)
create_uid = fields.Many2one('res.users', string='发送者', default=lambda self: self.env.uid)