forked from lsdefine/GenericAgent
- Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathagent_loop.py
More file actions
Latest commit
96 lines (86 loc) · 4.04 KB
/
Copy pathagent_loop.py
File metadata and controls
96 lines (86 loc) · 4.04 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
importjson, re
fromdataclassesimportdataclass
fromtypingimportAny, Optional
@dataclass
classStepOutcome:
data: Any
next_prompt: Optional[str] =None
should_exit: bool=False
deftry_call_generator(func, *args, **kwargs):
ret=func(*args, **kwargs)
ifhasattr(ret, '__iter__') andnotisinstance(ret, (str, bytes, dict, list)):
ret=yieldfromret
returnret
classBaseHandler:
deftool_before_callback(self, tool_name, args, response): pass
deftool_after_callback(self, tool_name, args, response, ret): pass
defnext_prompt_patcher(self, next_prompt, outcome, turn): returnnext_prompt
defdispatch(self, tool_name, args, response):
method_name=f"do_{tool_name}"
ifhasattr(self, method_name):
_=yieldfromtry_call_generator(self.tool_before_callback, tool_name, args, response)
ret=yieldfromtry_call_generator(getattr(self, method_name), args, response)
_=yieldfromtry_call_generator(self.tool_after_callback, tool_name, args, response, ret)
returnret
eliftool_name=='bad_json':
returnStepOutcome(None, next_prompt=args.get('msg', 'bad_json'), should_exit=False)
else:
yieldf"未知工具: {tool_name}\n"
returnStepOutcome(None, next_prompt=f"未知工具 {tool_name}", should_exit=False)
defjson_default(o):
ifisinstance(o, set): returnlist(o)
returnstr(o)
defexhaust(g):
try:
whileTrue: next(g)
exceptStopIterationase: returne.value
defget_pretty_json(data):
ifisinstance(data, dict) and"script"indata:
data=data.copy()
data["script"] =data["script"].replace("; ", ";\n ")
returnjson.dumps(data, indent=2, ensure_ascii=False).replace('\\n', '\n')
defagent_runner_loop(client, system_prompt, user_input, handler, tools_schema, max_turns=15, verbose=True):
messages= [
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_input}
]
forturninrange(max_turns):
yieldf"**LLM Running (Turn {turn+1}) ...**\n\n"
if (turn+1) %10==0: client.last_tools=''# 每10轮重置一次工具描述,避免上下文过大导致的模型性能下降
response_gen=client.chat(messages=messages, tools=tools_schema)
ifverbose:
response=yieldfromresponse_gen
yield'\n\n'
else:
response=exhaust(response_gen)
yieldresponse.content
ifnotresponse.tool_calls:
tool_name, args='no_tool', {}
else:
tool_call=response.tool_calls[0]
tool_name=tool_call.function.name
args=json.loads(tool_call.function.arguments)
iftool_name=='no_tool': pass
else:
showarg=get_pretty_json(args)
ifnotverboseandlen(showarg) >200: showarg=showarg[:200] +' ...'
yieldf"🛠️ **正在调用工具:** `{tool_name}` 📥**参数:**\n````text\n{showarg}\n````\n"
handler.current_turn=turn+1
gen=handler.dispatch(tool_name, args, response)
ifverbose:
yield'`````\n'
outcome=yieldfromgen
yield'`````\n'
else:
outcome=exhaust(gen)
ifoutcome.next_promptisNone: return {'result': 'CURRENT_TASK_DONE', 'data': outcome.data}
ifoutcome.should_exit: return {'result': 'EXITED', 'data': outcome.data}
ifoutcome.next_prompt.startswith('未知工具'): client.last_tools=''
next_prompt=""
ifoutcome.dataisnotNone:
datastr=json.dumps(outcome.data, ensure_ascii=False, default=json_default) iftype(outcome.data) in [dict, list] elsestr(outcome.data)
next_prompt+=f"<tool_result>\n{datastr}\n</tool_result>\n\n"
next_prompt+=outcome.next_prompt
next_prompt=handler.next_prompt_patcher(next_prompt, outcome, turn+1)
messages= [{"role": "user", "content": next_prompt}]
return {'result': 'MAX_TURNS_EXCEEDED'}