File size: 3,827 Bytes
430354e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
97
98
99
100
from helpers.tool import Tool, Response
from helpers import parallel_tools
from helpers.strings import sanitize_string


class ParallelTool(Tool):
    async def before_execution(self, **kwargs):
        self.log = None

    async def after_execution(self, response: Response, **kwargs):
        text = sanitize_string(response.message.strip())
        self.agent.hist_add_tool_result(
            self.name,
            text,
            **(response.additional or {}),
        )
        await parallel_tools.collect_parallel_jobs(
            self.agent,
            getattr(self, "_collect_job_ids", []),
            promote_parent_history=True,
        )

    async def execute(self, **kwargs) -> Response:
        self._collect_job_ids = []
        args = {**self.args, **kwargs}
        action = str(args.get("action") or "").strip().lower()

        try:
            timeout = parallel_tools.coerce_timeout(args.get("timeout"))
            job_ids = parallel_tools.normalize_job_ids(args.get("job_ids"))

            if action == "cancel":
                results = await parallel_tools.cancel_parallel_jobs(self.agent, job_ids)
                return Response(
                    message=parallel_tools.format_parallel_results(results),
                    break_loop=False,
                )

            raw_calls = parallel_tools.extract_tool_calls(args)
            started_jobs = []
            if raw_calls is not None:
                calls = parallel_tools.normalize_parallel_tool_calls(raw_calls)
                started_jobs = await parallel_tools.start_parallel_jobs(self.agent, calls)

            started_job_ids = [job.id for job in started_jobs]
            all_job_ids = [*job_ids, *started_job_ids]

            if not all_job_ids:
                return Response(
                    message=(
                        "Error: provide `tool_calls` to start parallel jobs, "
                        "or `job_ids` to await/cancel existing jobs."
                    ),
                    break_loop=False,
                )

            wait_default = action not in {"start", "background", "collect"}
            wait = parallel_tools.coerce_bool(args.get("wait"), wait_default)
            if action in {"await", "wait"}:
                wait = True

            if not wait:
                if not started_jobs and not job_ids:
                    return Response(
                        message="Error: `wait: false` requires `tool_calls` to start new jobs.",
                        break_loop=False,
                    )
                if not job_ids:
                    return Response(
                        message=parallel_tools.format_started_jobs(started_jobs),
                        break_loop=False,
                    )
                results = await parallel_tools.await_parallel_jobs(
                    self.agent,
                    all_job_ids,
                    timeout=timeout,
                    collect=False,
                    wait=False,
                )
                self._collect_job_ids = [result["job_id"] for result in results]
                return Response(
                    message=parallel_tools.format_parallel_results(results),
                    break_loop=False,
                )

            results = await parallel_tools.await_parallel_jobs(
                self.agent,
                all_job_ids,
                timeout=timeout,
                collect=False,
                wait=True,
            )
            self._collect_job_ids = [result["job_id"] for result in results]
            return Response(
                message=parallel_tools.format_parallel_results(results),
                break_loop=False,
            )
        except ValueError as exc:
            return Response(message=f"Error: {exc}", break_loop=False)