From 1b94db779024a41c7caf4c97a20a7c814dd41854 Mon Sep 17 00:00:00 2001 From: presstab Date: Fri, 19 Jun 2026 14:53:59 -0600 Subject: [PATCH 1/4] feat: add OpenRouter quantization support for model configuration - Add `/model quantizations [int4, int8]` command - Pass quantizations as provider routing in OpenAI stream for OpenRouter models - Introduce `_parse_quantizations`, `set_model_quantizations`, `update_model_quantizations` - Update help text and tests accordingly - Fix error handling in fetch_context_phase by broadening exception catch --- .../agents/pipeline/fetch_context_phase.py | 4 +- src/jrdev/commands/help.py | 4 +- src/jrdev/commands/model.py | 77 +++++++++++++ src/jrdev/core/application.py | 9 ++ src/jrdev/models/model_list.py | 14 ++- src/jrdev/services/streaming/openai_stream.py | 8 +- tests/test_commands_model.py | 44 +++++++- tests/test_openai_stream.py | 102 ++++++++++++++++++ 8 files changed, 252 insertions(+), 10 deletions(-) create mode 100644 tests/test_openai_stream.py diff --git a/src/jrdev/agents/pipeline/fetch_context_phase.py b/src/jrdev/agents/pipeline/fetch_context_phase.py index a79b2db..4a81a84 100644 --- a/src/jrdev/agents/pipeline/fetch_context_phase.py +++ b/src/jrdev/agents/pipeline/fetch_context_phase.py @@ -142,8 +142,8 @@ async def ask_files_sufficient(self, files: List[str], user_task: str): self.app.ui.update_task_info(sub_task_str, update={"sub_task_finished": True}) json_content = cutoff_string(response, "```json", "```") - tool_calls = json.loads(json_content) try: + tool_calls = json.loads(json_content) if tool_calls: tool = tool_calls.get("tool") if tool and tool == "read": @@ -154,7 +154,7 @@ async def ask_files_sufficient(self, files: List[str], user_task: str): if file not in files: files.append(file) self.app.logger.info(f"Adding file {file}") - except AttributeError as e: + except (json.JSONDecodeError, TypeError, AttributeError) as e: self.app.logger.error("ask_files_sufficient: malformed additional files response %s", str(e)) self.app.ui.print_text("ask_files_sufficient response malformed", PrintType.ERROR) return files diff --git a/src/jrdev/commands/help.py b/src/jrdev/commands/help.py index bfffdb6..c9ea905 100644 --- a/src/jrdev/commands/help.py +++ b/src/jrdev/commands/help.py @@ -82,7 +82,7 @@ async def handle_help(app: Any, args: List[str], _worker_id: str): app.ui.print_text( f" {cmd_format}{format_command_with_args('/model', ' [args]')}{reset} " "- Manage and add models (/model add|edit " - ")", + "; /model quantizations [int4, int8] for OpenRouter)", print_type=None, ) app.ui.print_text(f" {cmd_format}/models{reset} - List all available models", print_type=None) @@ -215,7 +215,7 @@ async def handle_help_plain(app: Any, _args: List[str]): app.ui.print_text( " /model [args] - Manage and add models (/model add|edit " - " )", + " ; /model quantizations [int4, int8] for OpenRouter)", print_type=None, ) app.ui.print_text(" /models - List all available models", print_type=None) diff --git a/src/jrdev/commands/model.py b/src/jrdev/commands/model.py index d9c7397..0c88c45 100644 --- a/src/jrdev/commands/model.py +++ b/src/jrdev/commands/model.py @@ -9,6 +9,8 @@ from jrdev.ui.ui import PrintType from jrdev.utils.string_utils import is_valid_context_window, is_valid_cost, is_valid_name +OPENROUTER_PROVIDER = "open_router" + def _parse_bool(val: str) -> bool: """Parse a string to boolean, accepting common true/false values.""" @@ -21,6 +23,31 @@ def _parse_bool(val: str) -> bool: raise ValueError(f"Invalid boolean value: {val}") +def _parse_quantizations(args: List[str]) -> List[str] | None: + """Parse a quantization list from CLI args such as [int4, int8].""" + raw = " ".join(args).strip() + if raw.startswith("[") and raw.endswith("]"): + raw = raw[1:-1].strip() + + if not raw: + return None + + if "," in raw: + quantizations = [item.strip() for item in raw.split(",")] + else: + quantizations = raw.split() + + quantizations = [item for item in quantizations if item] + if not quantizations: + return None + + for quantization in quantizations: + if not is_valid_name(quantization, max_len=32): + return None + + return quantizations + + # pylint: disable=too-many-return-statements def _parse_model_arguments(app, args: List[str], start_idx: int) -> Union[dict, None]: """Parse and validate model arguments starting from given index.""" @@ -101,6 +128,8 @@ async def handle_model(app, args: List[str], _worker_id: str): list - Shows all models available in your configuration. set - Sets the active model for the chat. remove - Removes a model from your configuration. + quantizations [values] + - Sets OpenRouter provider quantization filters. add - Adds a new model to your configuration. edit @@ -125,6 +154,8 @@ async def handle_model(app, args: List[str], _worker_id: str): " /model list - Shows all models available in your user_models.json.\n" " /model set - Sets the active model for the chat.\n" " /model remove - Removes a model from your user_models.json.\n" + " /model quantizations [int4, int8]\n" + " - Set OpenRouter provider quantization filters for a model.\n" " /model add \n" " - Add a new model to your user_models.json.\n" " and are the cost per 1,000,000 tokens (as a float, in dollars).\n" @@ -145,6 +176,10 @@ async def handle_model(app, args: List[str], _worker_id: str): return subcommand = args[1].lower() + if len(args) >= 3 and args[2].lower() == "quantizations": + _handle_quantizations(app, args, available_model_names) + return + _handle_subcommand(app, subcommand, args, available_model_names) if subcommand not in ["list", "set", "remove", "add", "edit"]: @@ -174,6 +209,48 @@ def _handle_subcommand(app: Any, subcommand: str, args: List[str], available_mod _handle_edit(app, args, available_model_names) +def _handle_quantizations(app: Any, args: List[str], available_model_names: List[str]) -> None: + if len(args) < 4: + app.ui.print_text("Usage: /model quantizations [int4, int8]", PrintType.ERROR) + return + + model_name = args[1] + if model_name not in available_model_names: + app.ui.print_text(f"Error: Model '{model_name}' not found in your configuration.", PrintType.ERROR) + return + + model_info = next((model for model in app.get_models() if model["name"] == model_name), None) + if not model_info: + app.ui.print_text(f"Error: Model '{model_name}' not found in your configuration.", PrintType.ERROR) + return + + provider = model_info.get("provider") + if provider != OPENROUTER_PROVIDER: + app.ui.print_text( + f"Error: Quantizations can only be set for OpenRouter models. '{model_name}' uses provider '{provider}'.", + PrintType.ERROR, + ) + return + + quantizations = _parse_quantizations(args[3:]) + if quantizations is None: + app.ui.print_text( + "Invalid quantizations. Usage: /model quantizations [int4, int8]", PrintType.ERROR + ) + return + + if not app.set_model_quantizations(model_name, quantizations): + app.logger.info(f"Failed to set quantizations for model {model_name}") + app.ui.print_text(f"Failed to set quantizations for model '{model_name}'", PrintType.ERROR) + return + + app.logger.info(f"Set quantizations for model {model_name}: {quantizations}") + app.ui.print_text( + f"Set quantizations for model '{model_name}' to: {', '.join(quantizations)}", + PrintType.SUCCESS, + ) + + def _handle_set(app: Any, args: List[str], available_model_names: List[str]) -> None: if len(args) < 3: app.ui.print_text("Usage: /model set ", PrintType.ERROR) diff --git a/src/jrdev/core/application.py b/src/jrdev/core/application.py index 0dd36ed..b27d438 100644 --- a/src/jrdev/core/application.py +++ b/src/jrdev/core/application.py @@ -457,6 +457,15 @@ def edit_model(self, model_name: str, provider: str, is_think: bool, input_cost: self.ui.model_list_updated() return True + def set_model_quantizations(self, model_name: str, quantizations: List[str]) -> bool: + """Set OpenRouter provider quantization filters for a model and flush to disk.""" + if not self.state.model_list.update_model_quantizations(model_name, quantizations): + return False + + save_models(self.state.model_list.get_model_list()) + self.ui.model_list_updated() + return True + def refresh_model_list(self): # 1) grab every model from user's config (single source of truth) models = load_models() diff --git a/src/jrdev/models/model_list.py b/src/jrdev/models/model_list.py index 238cdef..2a130ba 100644 --- a/src/jrdev/models/model_list.py +++ b/src/jrdev/models/model_list.py @@ -75,6 +75,18 @@ def update_model(self, model_name: str, provider: str, is_think: bool, input_cos return True return False + def update_model_quantizations(self, model_name: str, quantizations: List[str]) -> bool: + """ + Update OpenRouter provider quantization filters for an existing model. + Returns True if the model was updated, False if the model was not found. + """ + with self._lock: + for m in self._model_list: + if m["name"] == model_name: + m["quantizations"] = quantizations + return True + return False + def add_model(self, model_name: str, provider: str, is_think: bool, input_cost: int, output_cost: int, context_window: int) -> bool: """ Add a new model to the model list if it does not already exist. @@ -100,4 +112,4 @@ def add_model(self, model_name: str, provider: str, is_think: bool, input_cost: "context_tokens": context_window } self._model_list.append(model_dict) - return True \ No newline at end of file + return True diff --git a/src/jrdev/services/streaming/openai_stream.py b/src/jrdev/services/streaming/openai_stream.py index 259fcf3..5fb9d3c 100644 --- a/src/jrdev/services/streaming/openai_stream.py +++ b/src/jrdev/services/streaming/openai_stream.py @@ -14,12 +14,14 @@ async def stream_openai_format(app, model, messages, task_id=None, print_stream= app.logger.info(log_msg) # Get the appropriate client + model_info = None model_provider = None # Find the model in AVAILABLE_MODELS available_models = app.get_models() for entry in available_models: if entry["name"] == model: + model_info = entry model_provider = entry["provider"] break @@ -61,6 +63,10 @@ async def stream_openai_format(app, model, messages, task_id=None, print_stream= kwargs["response_format"] = {"type": "json_object"} elif model_provider == "venice": kwargs["extra_body"] = {"venice_parameters": {"include_venice_system_prompt": False}} + elif model_provider == "open_router" and model_info: + quantizations = model_info.get("quantizations") + if isinstance(quantizations, list) and quantizations and all(isinstance(item, str) for item in quantizations): + kwargs["extra_body"] = {"provider": {"quantizations": quantizations}} stream = await client.chat.completions.create(**kwargs) @@ -127,4 +133,4 @@ async def stream_openai_format(app, model, messages, task_id=None, print_stream= elapsed_seconds = round(end_time - start_time, 2) stream_elapsed = end_time - stream_start_time app.logger.info(f"Response completed (no usage data in final chunk): {model}, {elapsed_seconds}s, {chunk_count} chunks, {round(chunk_count/stream_elapsed,2) if stream_elapsed > 0 else 0} chunks/sec") - await get_instance().add_use(model, input_token_estimate, output_tokens_estimate) \ No newline at end of file + await get_instance().add_use(model, input_token_estimate, output_tokens_estimate) diff --git a/tests/test_commands_model.py b/tests/test_commands_model.py index 5caf43e..f19e802 100644 --- a/tests/test_commands_model.py +++ b/tests/test_commands_model.py @@ -35,6 +35,7 @@ def __init__(self, models=None, model=None): self.added_models = [] self.removed_models = [] self.edited_models = [] + self.quantization_models = [] self.failed_remove = False def get_model_names(self): return [m["name"] for m in self._models] @@ -79,12 +80,21 @@ def edit_model(self, name, provider, is_think, input_cost, output_cost, context_ self.edited_models.append(name) return True return False + def set_model_quantizations(self, name, quantizations): + for m in self._models: + if m["name"] == name: + m["quantizations"] = quantizations + self.quantization_models.append(name) + return True + return False + class TestModelCommand(unittest.TestCase): def setUp(self): self.default_models = [ {"name": "gpt-4", "provider": "openai", "is_think": True, "input_cost": 1, "output_cost": 2, "context_tokens": 8192}, - {"name": "gpt-3.5", "provider": "openai", "is_think": False, "input_cost": 1, "output_cost": 2, "context_tokens": 4096} + {"name": "gpt-3.5", "provider": "openai", "is_think": False, "input_cost": 1, "output_cost": 2, "context_tokens": 4096}, + {"name": "openai/glm-5", "provider": "open_router", "is_think": True, "input_cost": 1, "output_cost": 2, "context_tokens": 128000} ] self.app = DummyApp(models=[m.copy() for m in self.default_models], model="gpt-4") @@ -93,7 +103,7 @@ def test_usage_message_and_current_model(self): out = "\n".join(msg for msg, _ in self.app.ui.printed) self.assertIn("Current chat model: gpt-4", out) self.assertIn("/model list", out) - self.assertIn("Available models: gpt-4, gpt-3.5", out) + self.assertIn("Available models: gpt-4, gpt-3.5, openai/glm-5", out) def test_list_models(self): run_async(model_cmd.handle_model(self.app, ["/model", "list"], "w1")) @@ -101,6 +111,7 @@ def test_list_models(self): self.assertIn("Available models (from your user_models.json):", out) self.assertIn(" - gpt-4", out) self.assertIn(" - gpt-3.5", out) + self.assertIn(" - openai/glm-5", out) def test_list_models_empty(self): app = DummyApp(models=[]) @@ -119,7 +130,7 @@ def test_set_model_missing_arg(self): run_async(model_cmd.handle_model(self.app, ["/model", "set"], "w1")) out = "\n".join(msg for msg, _ in self.app.ui.printed) self.assertIn("Usage: /model set ", out) - self.assertIn("Available models: gpt-4, gpt-3.5", out) + self.assertIn("Available models: gpt-4, gpt-3.5, openai/glm-5", out) def test_set_model_not_found(self): run_async(model_cmd.handle_model(self.app, ["/model", "set", "notfound"], "w1")) @@ -434,6 +445,31 @@ def test_edit_model_context_window_too_high(self): self.assertIn("context_window must be between 1 and 1,000,000,000", out) self.assertEqual(len(self.app.edited_models), 0) + def test_set_openrouter_quantizations_success(self): + args = ["/model", "openai/glm-5", "quantizations", "[int4,", "int8]"] + run_async(model_cmd.handle_model(self.app, args, "w1")) + model = next(m for m in self.app.get_models() if m["name"] == "openai/glm-5") + self.assertEqual(model["quantizations"], ["int4", "int8"]) + self.assertIn("openai/glm-5", self.app.quantization_models) + out = "\n".join(msg for msg, _ in self.app.ui.printed) + self.assertIn("Set quantizations for model 'openai/glm-5' to: int4, int8", out) + + def test_set_quantizations_rejects_non_openrouter_model(self): + args = ["/model", "gpt-4", "quantizations", "[int4,", "int8]"] + run_async(model_cmd.handle_model(self.app, args, "w1")) + model = next(m for m in self.app.get_models() if m["name"] == "gpt-4") + self.assertNotIn("quantizations", model) + out = "\n".join(msg for msg, _ in self.app.ui.printed) + self.assertIn("Quantizations can only be set for OpenRouter models", out) + + def test_set_quantizations_rejects_empty_list(self): + args = ["/model", "openai/glm-5", "quantizations", "[]"] + run_async(model_cmd.handle_model(self.app, args, "w1")) + model = next(m for m in self.app.get_models() if m["name"] == "openai/glm-5") + self.assertNotIn("quantizations", model) + out = "\n".join(msg for msg, _ in self.app.ui.printed) + self.assertIn("Invalid quantizations", out) + if __name__ == "__main__": - unittest.main() \ No newline at end of file + unittest.main() diff --git a/tests/test_openai_stream.py b/tests/test_openai_stream.py new file mode 100644 index 0000000..5f35bf3 --- /dev/null +++ b/tests/test_openai_stream.py @@ -0,0 +1,102 @@ +import asyncio +import os +import sys +import unittest +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../src"))) + +from jrdev.services.streaming.openai_stream import stream_openai_format + + +def run_async(coro): + return asyncio.get_event_loop().run_until_complete(coro) + + +class FakeUsage: + async def add_use(self, _model, _input_tokens, _output_tokens): + return None + + +class FakeStream: + def __aiter__(self): + self._chunks = iter( + [ + SimpleNamespace( + choices=[SimpleNamespace(delta=SimpleNamespace(content="ok"))], + usage=SimpleNamespace(prompt_tokens=1, completion_tokens=1), + ) + ] + ) + return self + + async def __anext__(self): + try: + return next(self._chunks) + except StopIteration as exc: + raise StopAsyncIteration from exc + + +class FakeCompletions: + def __init__(self): + self.kwargs = None + + async def create(self, **kwargs): + self.kwargs = kwargs + return FakeStream() + + +class FakeClient: + def __init__(self): + self.completions = FakeCompletions() + self.chat = SimpleNamespace(completions=self.completions) + + +class FakeClients: + def __init__(self, client): + self.client = client + + def get_all_clients(self): + return {"open_router": self.client} + + def get_client(self, provider): + if provider == "open_router": + return self.client + return None + + +class TestOpenAIStream(unittest.TestCase): + def test_openrouter_quantizations_sent_as_provider_routing(self): + client = FakeClient() + app = SimpleNamespace( + logger=MagicMock(), + ui=MagicMock(), + state=SimpleNamespace(clients=FakeClients(client)), + get_models=lambda: [ + { + "name": "openai/glm-5", + "provider": "open_router", + "quantizations": ["int4", "int8"], + } + ], + ) + + async def consume(): + chunks = [] + async for chunk in stream_openai_format(app, "openai/glm-5", [{"role": "user", "content": "hi"}]): + chunks.append(chunk) + return "".join(chunks) + + with patch("jrdev.services.streaming.openai_stream.get_instance", return_value=FakeUsage()): + response = run_async(consume()) + + self.assertEqual(response, "ok") + self.assertEqual( + client.completions.kwargs["extra_body"], + {"provider": {"quantizations": ["int4", "int8"]}}, + ) + + +if __name__ == "__main__": + unittest.main() From efa4eac372ae3ca840a5b4b33812e3280127d0aa Mon Sep 17 00:00:00 2001 From: presstab Date: Fri, 19 Jun 2026 15:19:10 -0600 Subject: [PATCH 2/4] feat: add model tracking to assistant messages and display model name in UI --- src/jrdev/agents/research_agent.py | 10 +++----- src/jrdev/agents/router_agent.py | 10 ++++---- src/jrdev/commands/compact.py | 2 +- src/jrdev/core/application.py | 3 ++- src/jrdev/messages/message_builder.py | 6 ++--- src/jrdev/messages/thread.py | 29 ++++++++++++++++------- src/jrdev/services/llm_requests.py | 12 +++++++++- src/jrdev/services/message_service.py | 9 +++---- src/jrdev/ui/cli_events.py | 5 ++-- src/jrdev/ui/tui/chat/chat_view_widget.py | 12 +++++++--- src/jrdev/ui/tui/chat/message_bubble.py | 18 ++++++++++++-- src/jrdev/ui/tui/textual_events.py | 7 +++--- src/jrdev/ui/ui_wrapper.py | 3 ++- 13 files changed, 85 insertions(+), 41 deletions(-) diff --git a/src/jrdev/agents/research_agent.py b/src/jrdev/agents/research_agent.py index 2a955cf..df5e411 100644 --- a/src/jrdev/agents/research_agent.py +++ b/src/jrdev/agents/research_agent.py @@ -94,9 +94,7 @@ async def interpret( summary_parts.append("No research actions were taken.") summary = "\n".join(summary_parts) - self.thread.messages.append( - {"role": "assistant", "content": summary} - ) + self.thread.add_response(summary, model=research_model) return {"type": "summary", "data": summary} if response_json is None: @@ -111,9 +109,7 @@ async def interpret( return None # Add the structured assistant response to history *after* successful parsing. - self.thread.messages.append( - {"role": "assistant", "content": json.dumps(response_json, indent=2)} - ) + self.thread.add_response(json.dumps(response_json, indent=2), model=research_model) decision = response_json.get("decision") @@ -147,7 +143,7 @@ async def interpret( if decision == "summary": summary = response_json.get("response", "") # Add summary to thread as final assistant message - self.thread.messages.append({"role": "assistant", "content": summary}) + self.thread.add_response(summary, model=research_model) return {"type": "summary", "data": summary} self.logger.error(f"Research agent returned an unknown decision: {decision}. Aborting.") diff --git a/src/jrdev/agents/router_agent.py b/src/jrdev/agents/router_agent.py index 7402df1..9125e1c 100644 --- a/src/jrdev/agents/router_agent.py +++ b/src/jrdev/agents/router_agent.py @@ -67,6 +67,7 @@ async def interpret( # Use a specific, fast model for this routing task router_model = self.app.profile_manager().get_model("intent_router") + response_model = router_model response_text = await generate_llm_response(self.app, router_model, messages, task_id=worker_id) # The user's input is part of the request, so add it to history. @@ -97,6 +98,7 @@ async def interpret( salvage_builder.add_user_message(response_text) salvage_model = self.app.profile_manager().get_model("quick_reasoning") + response_model = salvage_model response_text = await generate_llm_response(self.app, salvage_model, salvage_builder.build(), task_id=worker_id) try: @@ -105,15 +107,13 @@ async def interpret( except (json.JSONDecodeError, KeyError) as e2: self.logger.error(f"Failed to parse router JSON response{e2}\nResponse:\n {response_text}\n") msg = "Sorry, I had issues parsing my response. Do you want to try again?" - self.thread.messages.append({"role": "assistant", "content": msg}) + self.thread.add_response(msg, model=response_model) self.app.ui.print_text(msg, print_type=PrintType.ERROR) return None # Add the structured assistant response to history *after* successful parsing. # The content is the JSON string of the decision. - self.thread.messages.append( - {"role": "assistant", "content": json.dumps(response_json, indent=2)} - ) + self.thread.add_response(json.dumps(response_json, indent=2), model=response_model) decision = response_json.get("decision") @@ -152,4 +152,4 @@ async def interpret( self.app.ui.print_text(chat_response, print_type=PrintType.LLM) return None - return None \ No newline at end of file + return None diff --git a/src/jrdev/commands/compact.py b/src/jrdev/commands/compact.py index 8cd8ab1..a1d18b1 100644 --- a/src/jrdev/commands/compact.py +++ b/src/jrdev/commands/compact.py @@ -127,7 +127,7 @@ async def handle_compact(app: Any, args: List[str], worker_id: str) -> None: # Create new messages in the required format new_messages = [ {"role": "user", "content": compact_data["user"]}, - {"role": "assistant", "content": compact_data["assistant"]}, + {"role": "assistant", "content": compact_data["assistant"], "model": model}, ] # Replace the thread's messages with just these two diff --git a/src/jrdev/core/application.py b/src/jrdev/core/application.py index b27d438..4ccbf1d 100644 --- a/src/jrdev/core/application.py +++ b/src/jrdev/core/application.py @@ -584,9 +584,10 @@ async def process_chat_input(self, user_input, worker_id=None): # 3) stream the LLM response content = f"{USER_INPUT_PREFIX}{user_input}" + response_model = self.state.model async for chunk in self.message_service.stream_message(msg_thread, content, worker_id): # for each piece of text we hand it off to the UI - self.ui.stream_chunk(thread_id, chunk) + self.ui.stream_chunk(thread_id, chunk, response_model) # 4) at the end, notify UI to refresh thread list or button state self.ui.chat_thread_update(thread_id) diff --git a/src/jrdev/messages/message_builder.py b/src/jrdev/messages/message_builder.py index 82385bb..7ce1503 100644 --- a/src/jrdev/messages/message_builder.py +++ b/src/jrdev/messages/message_builder.py @@ -11,7 +11,7 @@ class MessageBuilder: def __init__(self, app: Any): self.app = app - self.messages: List[Dict[str, str]] = [] + self.messages: List[Dict[str, Any]] = [] self.files: Set[str] = set() self.project_files: Set[str] = set() self.include_tree: bool = False @@ -33,7 +33,7 @@ def add_assistant_message(self, content: str) -> None: """Add an assistant message to the conversation""" self.messages.append({"role": "assistant", "content": content}) - def add_historical_messages(self, messages: List[Dict[str, str]]) -> None: + def add_historical_messages(self, messages: List[Dict[str, Any]]) -> None: """Add historical message chain""" self.messages.extend(messages) @@ -176,7 +176,7 @@ def finalize_user_section(self) -> None: def clean(self) -> None: self.messages = [m for m in self.messages if m["content"] != ""] - def build(self) -> List[Dict[str, str]]: + def build(self) -> List[Dict[str, Any]]: """Return the fully constructed message list""" if not self.isUserSectionFinal: self.finalize_user_section() diff --git a/src/jrdev/messages/thread.py b/src/jrdev/messages/thread.py index 4795b02..d9b1a69 100644 --- a/src/jrdev/messages/thread.py +++ b/src/jrdev/messages/thread.py @@ -37,7 +37,7 @@ def __init__(self, thread_id: str): """ self.thread_id: str = thread_id self.name: Optional[str] = None - self.messages: List[Dict[str, str]] = [] + self.messages: List[Dict[str, Any]] = [] self.context: Set[str] = set() self.embedded_files: Set[str] = set() self.token_usage: Dict[str, int] = {"input": 0, "output": 0} @@ -171,31 +171,44 @@ def add_embedded_files(self, files: List[str]) -> None: self.metadata["last_modified"] = datetime.now() @auto_persist - def add_response(self, response: str) -> None: + def add_response(self, response: str, model: Optional[str] = None) -> None: """Add a complete assistant response to the thread history.""" - self.messages.append({"role": "assistant", "content": response}) + message = {"role": "assistant", "content": response} + if model: + message["model"] = model + self.messages.append(message) self.metadata["last_modified"] = datetime.now() @auto_persist - def add_response_partial(self, chunk: str) -> None: + def add_response_partial(self, chunk: str, model: Optional[str] = None) -> None: """Add a partial assistant response chunk to the thread history.""" if self.messages and self.messages[-1].get("role") == "assistant": self.messages[-1]["content"] += chunk + if model: + self.messages[-1]["model"] = model else: - self.messages.append({"role": "assistant", "content": chunk}) + message = {"role": "assistant", "content": chunk} + if model: + message["model"] = model + self.messages.append(message) self.metadata["last_modified"] = datetime.now() @auto_persist - def finalize_response(self, full_text: str) -> None: + def finalize_response(self, full_text: str, model: Optional[str] = None) -> None: """Finalize the assistant response, replacing partials with full text.""" if self.messages and self.messages[-1].get("role") == "assistant": self.messages[-1]["content"] = full_text + if model: + self.messages[-1]["model"] = model else: - self.messages.append({"role": "assistant", "content": full_text}) + message = {"role": "assistant", "content": full_text} + if model: + message["model"] = model + self.messages.append(message) self.metadata["last_modified"] = datetime.now() @auto_persist - def set_compacted(self, messages: List[Dict[str, str]]) -> None: + def set_compacted(self, messages: List[Dict[str, Any]]) -> None: """Replace the existing messages list and reset file states.""" self.messages = messages self.context = set() diff --git a/src/jrdev/services/llm_requests.py b/src/jrdev/services/llm_requests.py index 05e3612..0b753cd 100644 --- a/src/jrdev/services/llm_requests.py +++ b/src/jrdev/services/llm_requests.py @@ -6,8 +6,18 @@ from jrdev.services.streaming.openai_stream import stream_openai_format +def _provider_messages(messages): + """Return provider-safe chat messages without local persistence metadata.""" + return [ + {"role": msg["role"], "content": msg["content"]} + for msg in messages + if "role" in msg and "content" in msg + ] + + def stream_request(app, model, messages, task_id=None, print_stream=True, json_output=False, max_output_tokens=None) -> AsyncIterator[str]: """Route a streaming LLM request to the appropriate provider based on the model.""" + messages = _provider_messages(messages) model_provider = None for entry in app.get_models(): if entry["name"] == model: @@ -62,4 +72,4 @@ async def generate_llm_response(app, model, messages, task_id=None, print_stream # try again app.logger.info("Attempting LLM stream again") attempts += 1 - return await generate_llm_response(app, model, messages, task_id, print_stream, json_output, max_output_tokens, attempts) \ No newline at end of file + return await generate_llm_response(app, model, messages, task_id, print_stream, json_output, max_output_tokens, attempts) diff --git a/src/jrdev/services/message_service.py b/src/jrdev/services/message_service.py index 02dc421..94bca58 100644 --- a/src/jrdev/services/message_service.py +++ b/src/jrdev/services/message_service.py @@ -63,9 +63,10 @@ async def stream_message(self, msg_thread: MessageThread, content: str, task_id: response_accumulator = "" try: # stream_request returns an async generator directly as per refactoring note (b) + response_model = self.app.state.model llm_response_stream = stream_request( self.app, - self.app.state.model, + response_model, messages_for_llm, task_id ) @@ -81,18 +82,18 @@ async def stream_message(self, msg_thread: MessageThread, content: str, task_id: yield "Thinking..." else: response_accumulator += chunk - msg_thread.add_response_partial(chunk) # Update thread with partial assistant response + msg_thread.add_response_partial(chunk, model=response_model) # Update thread with partial assistant response yield chunk elif in_think: if chunk == "": in_think = False else: response_accumulator += chunk - msg_thread.add_response_partial(chunk) # Update thread with partial assistant response + msg_thread.add_response_partial(chunk, model=response_model) # Update thread with partial assistant response yield chunk # Finalize the full response in the message thread - msg_thread.finalize_response(response_accumulator.strip()) + msg_thread.finalize_response(response_accumulator.strip(), model=response_model) except Exception as e: logger.error("Message Service: %s", e) if task_id: diff --git a/src/jrdev/ui/cli_events.py b/src/jrdev/ui/cli_events.py index 953556d..258fe45 100644 --- a/src/jrdev/ui/cli_events.py +++ b/src/jrdev/ui/cli_events.py @@ -22,13 +22,14 @@ def print_stream(self, message: str): if self.capture_active: self.capture += message - def stream_chunk(self, thread_id: str, chunk: str) -> None: + def stream_chunk(self, thread_id: str, chunk: str, model: Optional[str] = None) -> None: """ Handle an incoming chunk of text from a streaming LLM response. Args: thread_id: The ID of the conversation thread this chunk belongs to. chunk: The piece of text from the AI's response. + model: Optional model name associated with the response. """ terminal_print(chunk, PrintType.LLM, end="", flush=True) @@ -272,4 +273,4 @@ def providers_updated(self) -> None: pass def model_list_updated(self) -> None: - pass \ No newline at end of file + pass diff --git a/src/jrdev/ui/tui/chat/chat_view_widget.py b/src/jrdev/ui/tui/chat/chat_view_widget.py index 4d6d9a6..eab43d2 100644 --- a/src/jrdev/ui/tui/chat/chat_view_widget.py +++ b/src/jrdev/ui/tui/chat/chat_view_widget.py @@ -288,7 +288,7 @@ async def _load_current_thread(self) -> None: else: display_content = body - bubble = MessageBubble(display_content, role=role) + bubble = MessageBubble(display_content, role=role, model=msg.get("model")) await self.message_scroller.mount(bubble) await self._prune_bubbles() @@ -318,11 +318,17 @@ async def handle_stream_chunk(self, event: TextualEvents.StreamChunk) -> None: bubbles = [child for child in self.message_scroller.children if isinstance(child, MessageBubble)] last_bubble = bubbles[-1] if bubbles else None + model = event.model + if not model and active_thread.messages: + latest_message = active_thread.messages[-1] + if latest_message.get("role") == "assistant": + model = latest_message.get("model") if last_bubble and last_bubble.role == "assistant": + last_bubble.set_model(model) last_bubble.append_chunk(event.chunk) else: - new_bubble = MessageBubble(event.chunk, role="assistant") + new_bubble = MessageBubble(event.chunk, role="assistant", model=model) await self.message_scroller.mount(new_bubble) await self._prune_bubbles() @@ -506,4 +512,4 @@ def update_models(self) -> None: def handle_external_update(self, is_enabled: bool) -> None: """Handles external updates to the project context state (e.g., from core app).""" if self.context_switch.value != is_enabled: - self.set_project_context_on(is_enabled) \ No newline at end of file + self.set_project_context_on(is_enabled) diff --git a/src/jrdev/ui/tui/chat/message_bubble.py b/src/jrdev/ui/tui/chat/message_bubble.py index 5c1fb11..08457d7 100644 --- a/src/jrdev/ui/tui/chat/message_bubble.py +++ b/src/jrdev/ui/tui/chat/message_bubble.py @@ -33,10 +33,17 @@ class MessageBubble(Vertical): } """ - def __init__(self, message_content: str, role: str, id: str | None = None) -> None: + def __init__( + self, + message_content: str, + role: str, + id: str | None = None, + model: str | None = None + ) -> None: super().__init__(id=id) self.message_content = message_content self.role = role + self.model = model self.is_thinking = message_content == "Thinking..." border_color_map = { @@ -68,7 +75,7 @@ def on_mount(self) -> None: self.border_title = "Me" else: self.styles.border = ("round", Color.parse("#27dfd0")) - self.border_title = "Assistant" + self.border_title = self.model or "Assistant" @on(Button.Pressed) async def handle_copy_button(self, event: Button.Pressed) -> None: @@ -101,3 +108,10 @@ def append_chunk(self, chunk: str) -> None: # add the new text self.text_area.append_text(chunk) + + def set_model(self, model: str | None) -> None: + """Set the displayed model name for assistant messages.""" + if self.role != "assistant" or not model: + return + self.model = model + self.border_title = model diff --git a/src/jrdev/ui/tui/textual_events.py b/src/jrdev/ui/tui/textual_events.py index 8da9b4e..49127d8 100644 --- a/src/jrdev/ui/tui/textual_events.py +++ b/src/jrdev/ui/tui/textual_events.py @@ -58,10 +58,11 @@ def __init__(self, is_enabled): class StreamChunk(Message): """Fired once per text-chunk from the LLM.""" - def __init__(self, thread_id: str, chunk: str): + def __init__(self, thread_id: str, chunk: str, model: Optional[str] = None): super().__init__() self.thread_id = thread_id self.chunk = chunk + self.model = model class ConfirmationRequest(Message): def __init__(self, prompt_text: str, future: asyncio.Future, diff_lines: Optional[List[str]] = None, error_msg: str = None): @@ -224,9 +225,9 @@ def project_context_changed(self, is_enabled: bool) -> None: """Signal to UI that project context has been toggled on or off""" self.app.post_message(self.ProjectContextUpdate(is_enabled)) - def stream_chunk(self, thread_id: str, chunk: str) -> None: + def stream_chunk(self, thread_id: str, chunk: str, model: Optional[str] = None) -> None: """Post a chunk event into Textual's event bus.""" - self.app.post_message(self.StreamChunk(thread_id, chunk)) + self.app.post_message(self.StreamChunk(thread_id, chunk, model)) def providers_updated(self) -> None: """Providers has been changed (add/delete/edit)""" diff --git a/src/jrdev/ui/ui_wrapper.py b/src/jrdev/ui/ui_wrapper.py index e713873..223e994 100644 --- a/src/jrdev/ui/ui_wrapper.py +++ b/src/jrdev/ui/ui_wrapper.py @@ -15,13 +15,14 @@ def print_stream(self, message: str): """print a stream of text""" raise NotImplementedError("Subclasses must implement print_stream()") - def stream_chunk(self, thread_id: str, chunk: str) -> None: + def stream_chunk(self, thread_id: str, chunk: str, model: Optional[str] = None) -> None: """ Handle an incoming chunk of text from a streaming LLM response. Args: thread_id: The ID of the conversation thread this chunk belongs to. chunk: The piece of text from the AI's response. + model: Optional model name associated with the response. """ raise NotImplementedError("Subclasses must implement stream_chunk()") From d9016a5f68b117e619a377f2f0d73b931926e4de Mon Sep 17 00:00:00 2001 From: presstab Date: Fri, 19 Jun 2026 15:26:18 -0600 Subject: [PATCH 3/4] refactor: use asyncio.run() in test helper functions --- tests/test_commands_model.py | 2 +- tests/test_openai_stream.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/test_commands_model.py b/tests/test_commands_model.py index f19e802..4cc25b6 100644 --- a/tests/test_commands_model.py +++ b/tests/test_commands_model.py @@ -15,7 +15,7 @@ # Helper: run async function in sync test def run_async(coro): - return asyncio.get_event_loop().run_until_complete(coro) + return asyncio.run(coro) class DummyUI: def __init__(self): diff --git a/tests/test_openai_stream.py b/tests/test_openai_stream.py index 5f35bf3..a2f6b70 100644 --- a/tests/test_openai_stream.py +++ b/tests/test_openai_stream.py @@ -11,7 +11,7 @@ def run_async(coro): - return asyncio.get_event_loop().run_until_complete(coro) + return asyncio.run(coro) class FakeUsage: From 5d62200074673635146a229bedf95a478954dc7a Mon Sep 17 00:00:00 2001 From: presstab Date: Fri, 19 Jun 2026 15:30:25 -0600 Subject: [PATCH 4/4] refactor: add from __future__ import annotations for py 3.9 compat --- src/jrdev/commands/model.py | 2 ++ src/jrdev/core/application.py | 2 ++ src/jrdev/ui/tui/chat/message_bubble.py | 2 ++ .../ui/tui/settings/model_management/base_model_modal.py | 2 ++ src/jrdev/ui/tui/terminal/bordered_switcher.py | 4 +++- src/jrdev/ui/tui/terminal/terminal_text_area.py | 2 ++ src/jrdev/ui/tui/textual_ui.py | 2 ++ 7 files changed, 15 insertions(+), 1 deletion(-) diff --git a/src/jrdev/commands/model.py b/src/jrdev/commands/model.py index 0c88c45..0f06ccb 100644 --- a/src/jrdev/commands/model.py +++ b/src/jrdev/commands/model.py @@ -4,6 +4,8 @@ Model command implementation for the JrDev terminal. Manages the user's list of available models and the active chat model. """ +from __future__ import annotations + from typing import Any, List, Union from jrdev.ui.ui import PrintType diff --git a/src/jrdev/core/application.py b/src/jrdev/core/application.py index 4ccbf1d..0e0485e 100644 --- a/src/jrdev/core/application.py +++ b/src/jrdev/core/application.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import asyncio import json import os diff --git a/src/jrdev/ui/tui/chat/message_bubble.py b/src/jrdev/ui/tui/chat/message_bubble.py index 08457d7..5e44ab2 100644 --- a/src/jrdev/ui/tui/chat/message_bubble.py +++ b/src/jrdev/ui/tui/chat/message_bubble.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import pyperclip import logging diff --git a/src/jrdev/ui/tui/settings/model_management/base_model_modal.py b/src/jrdev/ui/tui/settings/model_management/base_model_modal.py index a207d3a..b3cb4b3 100644 --- a/src/jrdev/ui/tui/settings/model_management/base_model_modal.py +++ b/src/jrdev/ui/tui/settings/model_management/base_model_modal.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from textual.screen import ModalScreen from textual.widgets import Input, Button, Label, Select from textual.containers import Vertical, Horizontal diff --git a/src/jrdev/ui/tui/terminal/bordered_switcher.py b/src/jrdev/ui/tui/terminal/bordered_switcher.py index 1326424..754eb35 100644 --- a/src/jrdev/ui/tui/terminal/bordered_switcher.py +++ b/src/jrdev/ui/tui/terminal/bordered_switcher.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from textual.widgets import ContentSwitcher @@ -9,4 +11,4 @@ def watch_current(self, old: str|None, new: str|None) -> None: # now tell the App that we flipped panels # we assume the App implements _on_panel_switched(old, new) if hasattr(self.app, "_on_panel_switched"): - self.app._on_panel_switched(old, new) \ No newline at end of file + self.app._on_panel_switched(old, new) diff --git a/src/jrdev/ui/tui/terminal/terminal_text_area.py b/src/jrdev/ui/tui/terminal/terminal_text_area.py index 1b0d630..61ccaee 100644 --- a/src/jrdev/ui/tui/terminal/terminal_text_area.py +++ b/src/jrdev/ui/tui/terminal/terminal_text_area.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import logging import re from typing import Dict diff --git a/src/jrdev/ui/tui/textual_ui.py b/src/jrdev/ui/tui/textual_ui.py index 7127c72..0df0770 100644 --- a/src/jrdev/ui/tui/textual_ui.py +++ b/src/jrdev/ui/tui/textual_ui.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from jrdev.ui.tui.model_listview import ModelListView from jrdev.ui.ui import printtype_to_string from textual import on, events