v1.2.2
Co-Authored-By: Quentin Torroba <quentin.torroba@mistral.ai> Co-Authored-By: Michel Thomazo <michel.thomazo@mistral.ai>
This commit is contained in:
parent
402e898f39
commit
2e1e15120d
32 changed files with 391 additions and 549 deletions
|
|
@ -3,4 +3,4 @@ from __future__ import annotations
|
|||
from pathlib import Path
|
||||
|
||||
VIBE_ROOT = Path(__file__).parent
|
||||
__version__ = "1.2.1"
|
||||
__version__ = "1.2.2"
|
||||
|
|
|
|||
|
|
@ -38,7 +38,6 @@ class EventHandler:
|
|||
self.get_todos_collapsed = get_todos_collapsed
|
||||
self.current_tool_call: ToolCallMessage | None = None
|
||||
self.current_compact: CompactMessage | None = None
|
||||
self.tool_results: list[ToolResultMessage] = []
|
||||
|
||||
async def handle_event(
|
||||
self,
|
||||
|
|
@ -121,7 +120,6 @@ class EventHandler:
|
|||
)
|
||||
await self.mount_callback(tool_result)
|
||||
|
||||
self.tool_results.append(tool_result)
|
||||
self.current_tool_call = None
|
||||
|
||||
async def _handle_assistant_message(self, event: AssistantEvent) -> None:
|
||||
|
|
@ -153,6 +151,3 @@ class EventHandler:
|
|||
if self.current_compact:
|
||||
self.current_compact.stop_spinning(success=False)
|
||||
self.current_compact = None
|
||||
|
||||
def get_last_tool_result(self) -> ToolResultMessage | None:
|
||||
return self.tool_results[-1] if self.tool_results else None
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections import OrderedDict
|
||||
from collections.abc import AsyncGenerator, Callable
|
||||
from enum import StrEnum, auto
|
||||
import time
|
||||
|
|
@ -48,9 +47,9 @@ from vibe.core.types import (
|
|||
CompactStartEvent,
|
||||
LLMChunk,
|
||||
LLMMessage,
|
||||
LLMUsage,
|
||||
Role,
|
||||
SyncApprovalCallback,
|
||||
ToolCall,
|
||||
ToolCallEvent,
|
||||
ToolResultEvent,
|
||||
)
|
||||
|
|
@ -141,8 +140,6 @@ class Agent:
|
|||
config.effective_workdir,
|
||||
)
|
||||
|
||||
self._last_chunk: LLMChunk | None = None
|
||||
|
||||
@property
|
||||
def mode(self) -> AgentMode:
|
||||
return self._mode
|
||||
|
|
@ -267,11 +264,7 @@ class Agent:
|
|||
yield event
|
||||
|
||||
last_message = self.messages[-1]
|
||||
should_break_loop = (
|
||||
last_message.role != Role.tool
|
||||
and self._last_chunk is not None
|
||||
and self._last_chunk.finish_reason is not None
|
||||
)
|
||||
should_break_loop = last_message.role != Role.tool
|
||||
|
||||
self._flush_new_messages()
|
||||
|
||||
|
|
@ -293,9 +286,7 @@ class Agent:
|
|||
self.messages, self.stats, self.config, self.tool_manager
|
||||
)
|
||||
|
||||
async def _perform_llm_turn(
|
||||
self,
|
||||
) -> AsyncGenerator[AssistantEvent | ToolCallEvent | ToolResultEvent]:
|
||||
async def _perform_llm_turn(self) -> AsyncGenerator[BaseEvent, None]:
|
||||
if self.enable_streaming:
|
||||
async for event in self._stream_assistant_events():
|
||||
yield event
|
||||
|
|
@ -305,105 +296,40 @@ class Agent:
|
|||
yield assistant_event
|
||||
|
||||
last_message = self.messages[-1]
|
||||
last_chunk = self._last_chunk
|
||||
if last_chunk is None or last_chunk.usage is None:
|
||||
raise LLMResponseError("LLM response missing chunk or usage data")
|
||||
|
||||
parsed = self.format_handler.parse_message(last_message)
|
||||
resolved = self.format_handler.resolve_tool_calls(
|
||||
parsed, self.tool_manager, self.config
|
||||
)
|
||||
|
||||
if last_chunk.usage.completion_tokens > 0 and self.stats.last_turn_duration > 0:
|
||||
self.stats.tokens_per_second = (
|
||||
last_chunk.usage.completion_tokens / self.stats.last_turn_duration
|
||||
)
|
||||
|
||||
if not resolved.tool_calls and not resolved.failed_calls:
|
||||
return
|
||||
|
||||
async for event in self._handle_tool_calls(resolved):
|
||||
yield event
|
||||
|
||||
def _create_assistant_event(
|
||||
self, content: str, chunk: LLMChunk | None
|
||||
) -> AssistantEvent:
|
||||
return AssistantEvent(content=content)
|
||||
|
||||
async def _stream_assistant_events(self) -> AsyncGenerator[AssistantEvent]:
|
||||
chunks: list[LLMChunk] = []
|
||||
content_buffer = ""
|
||||
batched_chunk = LLMChunk(message=LLMMessage(role=Role.assistant))
|
||||
chunks_with_content = 0
|
||||
BATCH_SIZE = 5
|
||||
|
||||
async for chunk in self._chat_streaming():
|
||||
chunks.append(chunk)
|
||||
|
||||
if chunk.message.tool_calls and chunk.finish_reason is None:
|
||||
if chunk.message.content:
|
||||
content_buffer += chunk.message.content
|
||||
chunks_with_content += 1
|
||||
|
||||
if content_buffer:
|
||||
yield self._create_assistant_event(content_buffer, chunk)
|
||||
content_buffer = ""
|
||||
chunks_with_content = 0
|
||||
continue
|
||||
batched_chunk += chunk
|
||||
|
||||
if chunk.message.content:
|
||||
content_buffer += chunk.message.content
|
||||
chunks_with_content += 1
|
||||
|
||||
if chunks_with_content >= BATCH_SIZE:
|
||||
yield self._create_assistant_event(content_buffer, chunk)
|
||||
content_buffer = ""
|
||||
chunks_with_content = 0
|
||||
if chunks_with_content >= BATCH_SIZE:
|
||||
yield AssistantEvent(content=cast(str, batched_chunk.message.content))
|
||||
batched_chunk = LLMChunk(message=LLMMessage(role=Role.assistant))
|
||||
chunks_with_content = 0
|
||||
|
||||
if content_buffer:
|
||||
last_chunk = chunks[-1] if chunks else None
|
||||
yield self._create_assistant_event(content_buffer, last_chunk)
|
||||
|
||||
full_content = ""
|
||||
full_tool_calls_map = OrderedDict[int, ToolCall]()
|
||||
for chunk in chunks:
|
||||
full_content += chunk.message.content or ""
|
||||
if not chunk.message.tool_calls:
|
||||
continue
|
||||
|
||||
for tc in chunk.message.tool_calls:
|
||||
if tc.index is None:
|
||||
raise LLMResponseError("Tool call chunk missing index")
|
||||
if tc.index not in full_tool_calls_map:
|
||||
full_tool_calls_map[tc.index] = tc
|
||||
else:
|
||||
new_args_str = (
|
||||
full_tool_calls_map[tc.index].function.arguments or ""
|
||||
) + (tc.function.arguments or "")
|
||||
full_tool_calls_map[tc.index].function.arguments = new_args_str
|
||||
|
||||
full_tool_calls = list(full_tool_calls_map.values()) or None
|
||||
last_message = LLMMessage(
|
||||
role=Role.assistant, content=full_content, tool_calls=full_tool_calls
|
||||
)
|
||||
self.messages.append(last_message)
|
||||
finish_reason = next(
|
||||
(c.finish_reason for c in chunks if c.finish_reason is not None), None
|
||||
)
|
||||
self._last_chunk = LLMChunk(
|
||||
message=last_message, usage=chunks[-1].usage, finish_reason=finish_reason
|
||||
)
|
||||
if batched_chunk.message.content:
|
||||
yield AssistantEvent(content=batched_chunk.message.content)
|
||||
|
||||
async def _get_assistant_event(self) -> AssistantEvent:
|
||||
llm_result = await self._chat()
|
||||
if llm_result.usage is None:
|
||||
raise LLMResponseError(
|
||||
"Usage data missing in non-streaming completion response"
|
||||
)
|
||||
self._last_chunk = llm_result
|
||||
assistant_msg = llm_result.message
|
||||
self.messages.append(assistant_msg)
|
||||
|
||||
return AssistantEvent(content=assistant_msg.content or "")
|
||||
return AssistantEvent(content=llm_result.message.content or "")
|
||||
|
||||
async def _handle_tool_calls(
|
||||
self, resolved: ResolvedMessage
|
||||
|
|
@ -585,7 +511,6 @@ class Agent:
|
|||
|
||||
try:
|
||||
start_time = time.perf_counter()
|
||||
|
||||
async with self.backend as backend:
|
||||
result = await backend.complete(
|
||||
model=active_model,
|
||||
|
|
@ -599,31 +524,19 @@ class Agent:
|
|||
},
|
||||
max_tokens=max_tokens,
|
||||
)
|
||||
|
||||
end_time = time.perf_counter()
|
||||
|
||||
if result.usage is None:
|
||||
raise LLMResponseError(
|
||||
"Usage data missing in non-streaming completion response"
|
||||
)
|
||||
|
||||
self.stats.last_turn_duration = end_time - start_time
|
||||
self.stats.last_turn_prompt_tokens = result.usage.prompt_tokens
|
||||
self.stats.last_turn_completion_tokens = result.usage.completion_tokens
|
||||
self.stats.session_prompt_tokens += result.usage.prompt_tokens
|
||||
self.stats.session_completion_tokens += result.usage.completion_tokens
|
||||
self.stats.context_tokens = (
|
||||
result.usage.prompt_tokens + result.usage.completion_tokens
|
||||
)
|
||||
self._update_stats(usage=result.usage, time_seconds=end_time - start_time)
|
||||
|
||||
processed_message = self.format_handler.process_api_response_message(
|
||||
result.message
|
||||
)
|
||||
|
||||
return LLMChunk(
|
||||
message=processed_message,
|
||||
usage=result.usage,
|
||||
finish_reason=result.finish_reason,
|
||||
)
|
||||
self.messages.append(processed_message)
|
||||
return LLMChunk(message=processed_message, usage=result.usage)
|
||||
|
||||
except Exception as e:
|
||||
raise RuntimeError(
|
||||
|
|
@ -642,7 +555,8 @@ class Agent:
|
|||
tool_choice = self.format_handler.get_tool_choice()
|
||||
try:
|
||||
start_time = time.perf_counter()
|
||||
last_chunk = None
|
||||
usage = LLMUsage()
|
||||
chunk_agg = LLMChunk(message=LLMMessage(role=Role.assistant))
|
||||
async with self.backend as backend:
|
||||
async for chunk in backend.complete_streaming(
|
||||
model=active_model,
|
||||
|
|
@ -656,38 +570,40 @@ class Agent:
|
|||
},
|
||||
max_tokens=max_tokens,
|
||||
):
|
||||
last_chunk = chunk
|
||||
processed_message = (
|
||||
self.format_handler.process_api_response_message(chunk.message)
|
||||
)
|
||||
yield LLMChunk(
|
||||
message=processed_message,
|
||||
usage=chunk.usage,
|
||||
finish_reason=chunk.finish_reason,
|
||||
processed_chunk = LLMChunk(
|
||||
message=processed_message, usage=chunk.usage
|
||||
)
|
||||
|
||||
chunk_agg += processed_chunk
|
||||
usage += chunk.usage or LLMUsage()
|
||||
yield processed_chunk
|
||||
end_time = time.perf_counter()
|
||||
if last_chunk is None:
|
||||
raise LLMResponseError("Streamed completion returned no chunks")
|
||||
if last_chunk.usage is None:
|
||||
|
||||
if chunk_agg.usage is None:
|
||||
raise LLMResponseError(
|
||||
"Usage data missing in final chunk of streamed completion"
|
||||
)
|
||||
self._update_stats(usage=usage, time_seconds=end_time - start_time)
|
||||
|
||||
self.stats.last_turn_duration = end_time - start_time
|
||||
self.stats.last_turn_prompt_tokens = last_chunk.usage.prompt_tokens
|
||||
self.stats.last_turn_completion_tokens = last_chunk.usage.completion_tokens
|
||||
self.stats.session_prompt_tokens += last_chunk.usage.prompt_tokens
|
||||
self.stats.session_completion_tokens += last_chunk.usage.completion_tokens
|
||||
self.stats.context_tokens = (
|
||||
last_chunk.usage.prompt_tokens + last_chunk.usage.completion_tokens
|
||||
)
|
||||
self.messages.append(chunk_agg.message)
|
||||
|
||||
except Exception as e:
|
||||
raise RuntimeError(
|
||||
f"API error from {provider.name} (model: {active_model.name}): {e}"
|
||||
) from e
|
||||
|
||||
def _update_stats(self, usage: LLMUsage, time_seconds: float) -> None:
|
||||
self.stats.last_turn_duration = time_seconds
|
||||
self.stats.last_turn_prompt_tokens = usage.prompt_tokens
|
||||
self.stats.last_turn_completion_tokens = usage.completion_tokens
|
||||
self.stats.session_prompt_tokens += usage.prompt_tokens
|
||||
self.stats.session_completion_tokens += usage.completion_tokens
|
||||
self.stats.context_tokens = usage.prompt_tokens + usage.completion_tokens
|
||||
if time_seconds > 0 and usage.completion_tokens > 0:
|
||||
self.stats.tokens_per_second = usage.completion_tokens / time_seconds
|
||||
|
||||
async def _should_execute_tool(
|
||||
self, tool: BaseTool, args: dict[str, Any], tool_call_id: str
|
||||
) -> ToolDecision:
|
||||
|
|
|
|||
|
|
@ -143,17 +143,13 @@ class OpenAIAdapter(APIAdapter):
|
|||
message = LLMMessage.model_validate(data["choices"][0]["delta"])
|
||||
else:
|
||||
raise ValueError("Invalid response data")
|
||||
finish_reason = data["choices"][0].get("finish_reason", None)
|
||||
|
||||
elif "message" in data:
|
||||
message = LLMMessage.model_validate(data["message"])
|
||||
finish_reason = data["choices"][0].get("finish_reason", None)
|
||||
elif "delta" in data:
|
||||
message = LLMMessage.model_validate(data["delta"])
|
||||
finish_reason = None
|
||||
else:
|
||||
message = LLMMessage(role=Role.assistant, content="")
|
||||
finish_reason = None
|
||||
|
||||
usage_data = data.get("usage") or {}
|
||||
usage = LLMUsage(
|
||||
|
|
@ -161,7 +157,7 @@ class OpenAIAdapter(APIAdapter):
|
|||
completion_tokens=usage_data.get("completion_tokens", 0),
|
||||
)
|
||||
|
||||
return LLMChunk(message=message, usage=usage, finish_reason=finish_reason)
|
||||
return LLMChunk(message=message, usage=usage)
|
||||
|
||||
|
||||
class GenericBackend:
|
||||
|
|
@ -255,7 +251,7 @@ class GenericBackend:
|
|||
provider=self._provider.name,
|
||||
endpoint=url,
|
||||
response=e.response,
|
||||
headers=dict(e.response.headers.items()),
|
||||
headers=e.response.headers,
|
||||
model=model.name,
|
||||
messages=messages,
|
||||
temperature=temperature,
|
||||
|
|
@ -320,7 +316,7 @@ class GenericBackend:
|
|||
provider=self._provider.name,
|
||||
endpoint=url,
|
||||
response=e.response,
|
||||
headers=dict(e.response.headers.items()),
|
||||
headers=e.response.headers,
|
||||
model=model.name,
|
||||
messages=messages,
|
||||
temperature=temperature,
|
||||
|
|
@ -369,7 +365,11 @@ class GenericBackend:
|
|||
continue
|
||||
|
||||
DELIM_CHAR = ":"
|
||||
assert f"{DELIM_CHAR} " in line, "line should look like `key: value`"
|
||||
if f"{DELIM_CHAR} " not in line:
|
||||
raise ValueError(
|
||||
f"Stream chunk improperly formatted. "
|
||||
f"Expected `key{DELIM_CHAR} value`, received `{line}`"
|
||||
)
|
||||
delim_index = line.find(DELIM_CHAR)
|
||||
key = line[0:delim_index]
|
||||
value = line[delim_index + 2 :]
|
||||
|
|
@ -404,9 +404,8 @@ class GenericBackend:
|
|||
tool_choice=tool_choice,
|
||||
extra_headers=extra_headers,
|
||||
)
|
||||
assert result.usage is not None, (
|
||||
"Usage should be present in non-streaming completions"
|
||||
)
|
||||
if result.usage is None:
|
||||
raise ValueError("Missing usage in non streaming completion")
|
||||
|
||||
return result.usage.prompt_tokens
|
||||
|
||||
|
|
|
|||
|
|
@ -205,7 +205,6 @@ class MistralBackend:
|
|||
prompt_tokens=response.usage.prompt_tokens or 0,
|
||||
completion_tokens=response.usage.completion_tokens or 0,
|
||||
),
|
||||
finish_reason=response.choices[0].finish_reason,
|
||||
)
|
||||
|
||||
except mistralai.SDKError as e:
|
||||
|
|
@ -213,7 +212,7 @@ class MistralBackend:
|
|||
provider=self._provider.name,
|
||||
endpoint=self._server_url,
|
||||
response=e.raw_response,
|
||||
headers=dict(e.raw_response.headers.items()),
|
||||
headers=e.raw_response.headers,
|
||||
model=model.name,
|
||||
messages=messages,
|
||||
temperature=temperature,
|
||||
|
|
@ -279,7 +278,6 @@ class MistralBackend:
|
|||
if chunk.data.usage
|
||||
else 0,
|
||||
),
|
||||
finish_reason=chunk.data.choices[0].finish_reason,
|
||||
)
|
||||
|
||||
except mistralai.SDKError as e:
|
||||
|
|
@ -287,7 +285,7 @@ class MistralBackend:
|
|||
provider=self._provider.name,
|
||||
endpoint=self._server_url,
|
||||
response=e.raw_response,
|
||||
headers=dict(e.raw_response.headers.items()),
|
||||
headers=e.raw_response.headers,
|
||||
model=model.name,
|
||||
messages=messages,
|
||||
temperature=temperature,
|
||||
|
|
@ -325,8 +323,7 @@ class MistralBackend:
|
|||
tool_choice=tool_choice,
|
||||
extra_headers=extra_headers,
|
||||
)
|
||||
assert result.usage is not None, (
|
||||
"Usage should be present in non-streaming completions"
|
||||
)
|
||||
if result.usage is None:
|
||||
raise ValueError("Missing usage in non streaming completion")
|
||||
|
||||
return result.usage.prompt_tokens
|
||||
|
|
|
|||
|
|
@ -1,7 +1,9 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from abc import ABC
|
||||
from collections import OrderedDict
|
||||
from collections.abc import Awaitable, Callable
|
||||
import copy
|
||||
from enum import StrEnum, auto
|
||||
from typing import Annotated, Any, Literal
|
||||
|
||||
|
|
@ -185,19 +187,75 @@ class LLMMessage(BaseModel):
|
|||
"tool_call_id": getattr(v, "tool_call_id", None),
|
||||
}
|
||||
|
||||
def __add__(self, other: LLMMessage) -> LLMMessage:
|
||||
"""Careful: this is not commutative!"""
|
||||
if self.role != other.role:
|
||||
raise ValueError("Can't accumulate messages with different roles")
|
||||
|
||||
if self.name != other.name:
|
||||
raise ValueError("Can't accumulate messages with different names")
|
||||
|
||||
if self.tool_call_id != other.tool_call_id:
|
||||
raise ValueError("Can't accumulate messages with different tool_call_ids")
|
||||
|
||||
content = (self.content or "") + (other.content or "")
|
||||
if not content:
|
||||
content = None
|
||||
|
||||
tool_calls_map = OrderedDict[int, ToolCall]()
|
||||
for tool_calls in [self.tool_calls or [], other.tool_calls or []]:
|
||||
for tc in tool_calls:
|
||||
if tc.index is None:
|
||||
raise ValueError("Tool call chunk missing index")
|
||||
if tc.index not in tool_calls_map:
|
||||
tool_calls_map[tc.index] = copy.deepcopy(tc)
|
||||
else:
|
||||
existing_name = tool_calls_map[tc.index].function.name
|
||||
new_name = tc.function.name
|
||||
if existing_name and new_name and existing_name != new_name:
|
||||
raise ValueError(
|
||||
"Can't accumulate messages with different tool call names"
|
||||
)
|
||||
if new_name and not existing_name:
|
||||
tool_calls_map[tc.index].function.name = new_name
|
||||
new_args = (tool_calls_map[tc.index].function.arguments or "") + (
|
||||
tc.function.arguments or ""
|
||||
)
|
||||
tool_calls_map[tc.index].function.arguments = new_args
|
||||
|
||||
return LLMMessage(
|
||||
role=self.role,
|
||||
content=content,
|
||||
tool_calls=list(tool_calls_map.values()) or None,
|
||||
name=self.name,
|
||||
tool_call_id=self.tool_call_id,
|
||||
)
|
||||
|
||||
|
||||
class LLMUsage(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
prompt_tokens: int = 0
|
||||
completion_tokens: int = 0
|
||||
|
||||
def __add__(self, other: LLMUsage) -> LLMUsage:
|
||||
return LLMUsage(
|
||||
prompt_tokens=self.prompt_tokens + other.prompt_tokens,
|
||||
completion_tokens=self.completion_tokens + other.completion_tokens,
|
||||
)
|
||||
|
||||
|
||||
class LLMChunk(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
message: LLMMessage
|
||||
finish_reason: str | None = None
|
||||
usage: LLMUsage | None = None
|
||||
|
||||
def __add__(self, other: LLMChunk) -> LLMChunk:
|
||||
if self.usage is None and other.usage is None:
|
||||
new_usage = None
|
||||
else:
|
||||
new_usage = (self.usage or LLMUsage()) + (other.usage or LLMUsage())
|
||||
return LLMChunk(message=self.message + other.message, usage=new_usage)
|
||||
|
||||
|
||||
class BaseEvent(BaseModel, ABC):
|
||||
"""Abstract base class for all agent events."""
|
||||
|
|
@ -209,6 +267,13 @@ class AssistantEvent(BaseEvent):
|
|||
content: str
|
||||
stopped_by_middleware: bool = False
|
||||
|
||||
def __add__(self, other: AssistantEvent) -> AssistantEvent:
|
||||
return AssistantEvent(
|
||||
content=self.content + other.content,
|
||||
stopped_by_middleware=self.stopped_by_middleware
|
||||
or other.stopped_by_middleware,
|
||||
)
|
||||
|
||||
|
||||
class ToolCallEvent(BaseEvent):
|
||||
tool_name: str
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue