Co-Authored-By: Quentin Torroba <quentin.torroba@mistral.ai>
Co-Authored-By: Michel Thomazo <michel.thomazo@mistral.ai>
This commit is contained in:
Mathias Gesbert 2025-12-22 13:33:20 +01:00 committed by Mathias Gesbert
parent 402e898f39
commit 2e1e15120d
32 changed files with 391 additions and 549 deletions

View file

@ -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"

View file

@ -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

View file

@ -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:

View file

@ -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

View file

@ -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

View file

@ -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