v1.3.2
This commit is contained in:
parent
c79e2cf487
commit
08d8e85447
18 changed files with 287 additions and 51 deletions
|
|
@ -3,4 +3,4 @@ from __future__ import annotations
|
|||
from pathlib import Path
|
||||
|
||||
VIBE_ROOT = Path(__file__).parent
|
||||
__version__ = "1.3.1"
|
||||
__version__ = "1.3.2"
|
||||
|
|
|
|||
|
|
@ -9,11 +9,12 @@ from textual.app import ComposeResult
|
|||
from textual.containers import Horizontal
|
||||
from textual.widgets import Static
|
||||
|
||||
from vibe.cli.textual_ui.widgets.spinner import BrailleSpinner
|
||||
from vibe.cli.textual_ui.widgets.spinner import SpinnerMixin, SpinnerType
|
||||
|
||||
|
||||
class LoadingWidget(Static):
|
||||
class LoadingWidget(SpinnerMixin, Static):
|
||||
TARGET_COLORS = ("#FFD800", "#FFAF00", "#FF8205", "#FA500F", "#E10500")
|
||||
SPINNER_TYPE = SpinnerType.BRAILLE
|
||||
|
||||
EASTER_EGGS: ClassVar[list[str]] = [
|
||||
"Eating a chocolatine",
|
||||
|
|
@ -50,12 +51,11 @@ class LoadingWidget(Static):
|
|||
|
||||
def __init__(self, status: str | None = None) -> None:
|
||||
super().__init__(classes="loading-widget")
|
||||
self.init_spinner()
|
||||
self.status = status or self._get_default_status()
|
||||
self.current_color_index = 0
|
||||
self.transition_progress = 0
|
||||
self._spinner = BrailleSpinner()
|
||||
self.char_widgets: list[Static] = []
|
||||
self.spinner_widget: Static | None = None
|
||||
self.ellipsis_widget: Static | None = None
|
||||
self.hint_widget: Static | None = None
|
||||
self.start_time: float | None = None
|
||||
|
|
@ -89,10 +89,10 @@ class LoadingWidget(Static):
|
|||
|
||||
def compose(self) -> ComposeResult:
|
||||
with Horizontal(classes="loading-container"):
|
||||
self.spinner_widget = Static(
|
||||
self._indicator_widget = Static(
|
||||
self._spinner.current_frame(), classes="loading-indicator"
|
||||
)
|
||||
yield self.spinner_widget
|
||||
yield self._indicator_widget
|
||||
|
||||
with Horizontal(classes="loading-status"):
|
||||
for char in self.status:
|
||||
|
|
@ -120,12 +120,20 @@ class LoadingWidget(Static):
|
|||
self.char_widgets.append(widget)
|
||||
status_container.mount(widget)
|
||||
|
||||
self.update_animation()
|
||||
self._update_animation()
|
||||
|
||||
def on_mount(self) -> None:
|
||||
self.start_time = time()
|
||||
self.update_animation()
|
||||
self.set_interval(0.1, self.update_animation)
|
||||
self._update_animation()
|
||||
self.start_spinner_timer()
|
||||
|
||||
def on_resize(self) -> None:
|
||||
self.refresh_spinner()
|
||||
|
||||
def _update_spinner_frame(self) -> None:
|
||||
if not self._is_spinning:
|
||||
return
|
||||
self._update_animation()
|
||||
|
||||
def _get_color_for_position(self, position: int) -> str:
|
||||
current_color = self.TARGET_COLORS[self.current_color_index]
|
||||
|
|
@ -136,13 +144,13 @@ class LoadingWidget(Static):
|
|||
return next_color
|
||||
return current_color
|
||||
|
||||
def update_animation(self) -> None:
|
||||
def _update_animation(self) -> None:
|
||||
total_elements = 1 + len(self.char_widgets) + 2
|
||||
|
||||
if self.spinner_widget:
|
||||
if self._indicator_widget:
|
||||
spinner_char = self._spinner.next_frame()
|
||||
color = self._get_color_for_position(0)
|
||||
self.spinner_widget.update(f"[{color}]{spinner_char}[/]")
|
||||
self._indicator_widget.update(f"[{color}]{spinner_char}[/]")
|
||||
|
||||
for i, widget in enumerate(self.char_widgets):
|
||||
position = 1 + i
|
||||
|
|
|
|||
|
|
@ -153,6 +153,9 @@ class ReasoningMessage(SpinnerMixin, StreamingMessageBase):
|
|||
def on_mount(self) -> None:
|
||||
self.start_spinner_timer()
|
||||
|
||||
def on_resize(self) -> None:
|
||||
self.refresh_spinner()
|
||||
|
||||
async def on_click(self) -> None:
|
||||
await self._toggle_collapsed()
|
||||
|
||||
|
|
|
|||
|
|
@ -129,6 +129,10 @@ class SpinnerMixin:
|
|||
return
|
||||
self._indicator_widget.update(self._spinner.next_frame())
|
||||
|
||||
def refresh_spinner(self) -> None:
|
||||
if self._indicator_widget:
|
||||
self._indicator_widget.refresh()
|
||||
|
||||
def stop_spinning(self, success: bool = True) -> None:
|
||||
self._is_spinning = False
|
||||
if self._spinner_timer:
|
||||
|
|
|
|||
|
|
@ -38,6 +38,9 @@ class StatusMessage(SpinnerMixin, Static):
|
|||
self.update_display()
|
||||
self.start_spinner_timer()
|
||||
|
||||
def on_resize(self) -> None:
|
||||
self.refresh_spinner()
|
||||
|
||||
def _update_spinner_frame(self) -> None:
|
||||
if not self._is_spinning:
|
||||
return
|
||||
|
|
|
|||
|
|
@ -131,6 +131,7 @@ class ProviderConfig(BaseModel):
|
|||
api_key_env_var: str = ""
|
||||
api_style: str = "openai"
|
||||
backend: Backend = Backend.GENERIC
|
||||
reasoning_field_name: str = "reasoning_content"
|
||||
|
||||
|
||||
class _MCPBase(BaseModel):
|
||||
|
|
|
|||
|
|
@ -46,7 +46,9 @@ class APIAdapter(Protocol):
|
|||
api_key: str | None = None,
|
||||
) -> PreparedRequest: ...
|
||||
|
||||
def parse_response(self, data: dict[str, Any]) -> LLMChunk: ...
|
||||
def parse_response(
|
||||
self, data: dict[str, Any], provider: ProviderConfig
|
||||
) -> LLMChunk: ...
|
||||
|
||||
|
||||
BACKEND_ADAPTERS: dict[str, APIAdapter] = {}
|
||||
|
|
@ -103,6 +105,20 @@ class OpenAIAdapter(APIAdapter):
|
|||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
return headers
|
||||
|
||||
def _reasoning_to_api(
|
||||
self, msg_dict: dict[str, Any], field_name: str
|
||||
) -> dict[str, Any]:
|
||||
if field_name != "reasoning_content" and "reasoning_content" in msg_dict:
|
||||
msg_dict[field_name] = msg_dict.pop("reasoning_content")
|
||||
return msg_dict
|
||||
|
||||
def _reasoning_from_api(
|
||||
self, msg_dict: dict[str, Any], field_name: str
|
||||
) -> dict[str, Any]:
|
||||
if field_name != "reasoning_content" and field_name in msg_dict:
|
||||
msg_dict["reasoning_content"] = msg_dict.pop(field_name)
|
||||
return msg_dict
|
||||
|
||||
def prepare_request(
|
||||
self,
|
||||
*,
|
||||
|
|
@ -116,7 +132,11 @@ class OpenAIAdapter(APIAdapter):
|
|||
provider: ProviderConfig,
|
||||
api_key: str | None = None,
|
||||
) -> PreparedRequest:
|
||||
converted_messages = [msg.model_dump(exclude_none=True) for msg in messages]
|
||||
field_name = provider.reasoning_field_name
|
||||
converted_messages = [
|
||||
self._reasoning_to_api(msg.model_dump(exclude_none=True), field_name)
|
||||
for msg in messages
|
||||
]
|
||||
|
||||
payload = self.build_payload(
|
||||
model_name, converted_messages, temperature, tools, max_tokens, tool_choice
|
||||
|
|
@ -130,25 +150,37 @@ class OpenAIAdapter(APIAdapter):
|
|||
payload["stream_options"] = stream_options
|
||||
|
||||
headers = self.build_headers(api_key)
|
||||
|
||||
body = json.dumps(payload).encode("utf-8")
|
||||
|
||||
return PreparedRequest(self.endpoint, headers, body)
|
||||
|
||||
def parse_response(self, data: dict[str, Any]) -> LLMChunk:
|
||||
def _parse_message(
|
||||
self, data: dict[str, Any], field_name: str
|
||||
) -> LLMMessage | None:
|
||||
if data.get("choices"):
|
||||
if "message" in data["choices"][0]:
|
||||
message = LLMMessage.model_validate(data["choices"][0]["message"])
|
||||
elif "delta" in data["choices"][0]:
|
||||
message = LLMMessage.model_validate(data["choices"][0]["delta"])
|
||||
else:
|
||||
raise ValueError("Invalid response data")
|
||||
choice = data["choices"][0]
|
||||
if "message" in choice:
|
||||
msg_dict = self._reasoning_from_api(choice["message"], field_name)
|
||||
return LLMMessage.model_validate(msg_dict)
|
||||
if "delta" in choice:
|
||||
msg_dict = self._reasoning_from_api(choice["delta"], field_name)
|
||||
return LLMMessage.model_validate(msg_dict)
|
||||
raise ValueError("Invalid response data: missing message or delta")
|
||||
|
||||
elif "message" in data:
|
||||
message = LLMMessage.model_validate(data["message"])
|
||||
elif "delta" in data:
|
||||
message = LLMMessage.model_validate(data["delta"])
|
||||
else:
|
||||
if "message" in data:
|
||||
msg_dict = self._reasoning_from_api(data["message"], field_name)
|
||||
return LLMMessage.model_validate(msg_dict)
|
||||
if "delta" in data:
|
||||
msg_dict = self._reasoning_from_api(data["delta"], field_name)
|
||||
return LLMMessage.model_validate(msg_dict)
|
||||
|
||||
return None
|
||||
|
||||
def parse_response(
|
||||
self, data: dict[str, Any], provider: ProviderConfig
|
||||
) -> LLMChunk:
|
||||
message = self._parse_message(data, provider.reasoning_field_name)
|
||||
if message is None:
|
||||
message = LLMMessage(role=Role.assistant, content="")
|
||||
|
||||
usage_data = data.get("usage") or {}
|
||||
|
|
@ -244,7 +276,7 @@ class GenericBackend:
|
|||
|
||||
try:
|
||||
res_data, _ = await self._make_request(url, body, headers)
|
||||
return adapter.parse_response(res_data)
|
||||
return adapter.parse_response(res_data, self._provider)
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
raise BackendErrorBuilder.build_http_error(
|
||||
|
|
@ -309,7 +341,7 @@ class GenericBackend:
|
|||
|
||||
try:
|
||||
async for res_data in self._make_streaming_request(url, body, headers):
|
||||
yield adapter.parse_response(res_data)
|
||||
yield adapter.parse_response(res_data, self._provider)
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
raise BackendErrorBuilder.build_http_error(
|
||||
|
|
|
|||
|
|
@ -160,6 +160,13 @@ class MistralBackend:
|
|||
else None
|
||||
)
|
||||
|
||||
reasoning_field = getattr(provider, "reasoning_field_name", "reasoning_content")
|
||||
if reasoning_field != "reasoning_content":
|
||||
raise ValueError(
|
||||
f"Mistral backend does not support custom reasoning_field_name "
|
||||
f"(got '{reasoning_field}'). Mistral uses ThinkChunk for reasoning."
|
||||
)
|
||||
|
||||
# Mistral SDK takes server URL without api version as input
|
||||
url_pattern = r"(https?://[^/]+)(/v.*)"
|
||||
match = re.match(url_pattern, self._provider.api_base)
|
||||
|
|
|
|||
|
|
@ -179,15 +179,11 @@ class LLMMessage(BaseModel):
|
|||
if isinstance(v, dict):
|
||||
v.setdefault("content", "")
|
||||
v.setdefault("role", "assistant")
|
||||
v.setdefault(
|
||||
"reasoning_content", v.get("reasoning_content") or v.get("reasoning")
|
||||
)
|
||||
return v
|
||||
return {
|
||||
"role": str(getattr(v, "role", "assistant")),
|
||||
"content": getattr(v, "content", ""),
|
||||
"reasoning_content": getattr(v, "reasoning_content", None)
|
||||
or getattr(v, "reasoning", None),
|
||||
"reasoning_content": getattr(v, "reasoning_content", None),
|
||||
"tool_calls": getattr(v, "tool_calls", None),
|
||||
"name": getattr(v, "name", None),
|
||||
"tool_call_id": getattr(v, "tool_call_id", None),
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue