649 lines
28 KiB
Python
649 lines
28 KiB
Python
# -*- 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'
|
||
|
||
# —— 注入操作手册 ——
|
||
manual_content = self.env['ai.manual'].get_all_manual_content()
|
||
if manual_content:
|
||
system_msg += '\n--- 操作手册(用户询问操作方法时参考)---\n'
|
||
system_msg += manual_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)
|