diff --git a/src/codegate/dashboard/dashboard.py b/src/codegate/dashboard/dashboard.py index b844d224..fc146609 100644 --- a/src/codegate/dashboard/dashboard.py +++ b/src/codegate/dashboard/dashboard.py @@ -1,5 +1,5 @@ import asyncio -from typing import List, AsyncGenerator +from typing import AsyncGenerator, List import structlog from fastapi import APIRouter @@ -45,6 +45,7 @@ async def generate_sse_events() -> AsyncGenerator[str, None]: message = await alert_queue.get() yield f"data: {message}\n\n" + @dashboard_router.get("/dashboard/alerts_notification") async def stream_sse(): """ diff --git a/src/codegate/db/connection.py b/src/codegate/db/connection.py index e64bac21..3f6581ff 100644 --- a/src/codegate/db/connection.py +++ b/src/codegate/db/connection.py @@ -5,6 +5,7 @@ import uuid from pathlib import Path from typing import AsyncGenerator, AsyncIterator, List, Optional + import structlog from litellm import ChatCompletionRequest, ModelResponse from pydantic import BaseModel @@ -21,6 +22,7 @@ logger = structlog.get_logger("codegate") alert_queue = asyncio.Queue() + class DbCodeGate: def __init__(self, sqlite_path: Optional[str] = None): diff --git a/src/codegate/pipeline/codegate_context_retriever/codegate.py b/src/codegate/pipeline/codegate_context_retriever/codegate.py index f5eddc0f..0e49f01a 100644 --- a/src/codegate/pipeline/codegate_context_retriever/codegate.py +++ b/src/codegate/pipeline/codegate_context_retriever/codegate.py @@ -1,3 +1,5 @@ +import json + import structlog from litellm import ChatCompletionRequest @@ -34,7 +36,7 @@ async def get_objects_from_search( objects = await storage_engine.search(search, distance=0.8, packages=packages) return objects - def generate_context_str(self, objects: list[object]) -> str: + def generate_context_str(self, objects: list[object], context: PipelineContext) -> str: context_str = "" for obj in objects: # generate dictionary from object @@ -44,6 +46,12 @@ def generate_context_str(self, objects: list[object]) -> str: "status": obj.properties["status"], "description": obj.properties["description"], } + # Add one alert for each package found + context.add_alert( + self.name, + trigger_string=json.dumps(package_obj), + severity_category=AlertSeverity.CRITICAL, + ) package_str = generate_vector_string(package_obj) context_str += package_str + "\n" return context_str @@ -101,8 +109,8 @@ async def process( searched_objects = updated_searched_objects # Generate context string using the searched objects - logger.info(f"Adding {len(searched_objects)} packages to the context") - context_str = self.generate_context_str(searched_objects) + logger.info(f"Adding {len(updated_searched_objects)} packages to the context") + context_str = self.generate_context_str(updated_searched_objects, context) # Make a copy of the request new_request = request.copy() @@ -114,19 +122,11 @@ async def process( message = new_request["messages"][last_user_idx] if isinstance(message["content"], str): context_msg = f'Context: {context_str} \n\n Query: {message["content"]}' - context.add_alert( - self.name, trigger_string=context_msg, severity_category=AlertSeverity.CRITICAL - ) message["content"] = context_msg elif isinstance(message["content"], (list, tuple)): for item in message["content"]: if isinstance(item, dict) and item.get("type") == "text": context_msg = f'Context: {context_str} \n\n Query: {item["text"]}' - context.add_alert( - self.name, - trigger_string=context_msg, - severity_category=AlertSeverity.CRITICAL, - ) item["text"] = context_msg return PipelineResult(request=new_request, context=context) diff --git a/tests/test_cli.py b/tests/test_cli.py index bb0b32f3..73c12b80 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -79,11 +79,13 @@ def test_serve_default_options( } # Retrieve the actual call arguments - calls = [call[1]['extra'] for call in logger_instance.info.call_args_list] + calls = [call[1]["extra"] for call in logger_instance.info.call_args_list] # Check if one of the calls matches the expected subset - assert any(all(expected_extra[k] == actual_extra.get(k) - for k in expected_extra) for actual_extra in calls) + assert any( + all(expected_extra[k] == actual_extra.get(k) for k in expected_extra) + for actual_extra in calls + ) mock_run.assert_called_once() @@ -114,7 +116,7 @@ def test_serve_custom_options( mock_logging.assert_called_with("codegate") # Retrieve the actual call arguments - calls = [call[1]['extra'] for call in logger_instance.info.call_args_list] + calls = [call[1]["extra"] for call in logger_instance.info.call_args_list] expected_extra = { "host": "localhost", @@ -126,8 +128,10 @@ def test_serve_custom_options( } # Check if one of the calls matches the expected subset - assert any(all(expected_extra[k] == actual_extra.get(k) - for k in expected_extra) for actual_extra in calls) + assert any( + all(expected_extra[k] == actual_extra.get(k) for k in expected_extra) + for actual_extra in calls + ) mock_run.assert_called_once() @@ -159,7 +163,7 @@ def test_serve_with_config_file( mock_logging.assert_called_with("codegate") # Retrieve the actual call arguments - calls = [call[1]['extra'] for call in logger_instance.info.call_args_list] + calls = [call[1]["extra"] for call in logger_instance.info.call_args_list] expected_extra = { "host": "localhost", @@ -171,8 +175,10 @@ def test_serve_with_config_file( } # Check if one of the calls matches the expected subset - assert any(all(expected_extra[k] == actual_extra.get(k) - for k in expected_extra) for actual_extra in calls) + assert any( + all(expected_extra[k] == actual_extra.get(k) for k in expected_extra) + for actual_extra in calls + ) mock_run.assert_called_once() @@ -216,7 +222,7 @@ def test_serve_priority_resolution( mock_logging.assert_called_with("codegate") # Retrieve the actual call arguments - calls = [call[1]['extra'] for call in logger_instance.info.call_args_list] + calls = [call[1]["extra"] for call in logger_instance.info.call_args_list] expected_extra = { "host": "example.com", @@ -228,8 +234,10 @@ def test_serve_priority_resolution( } # Check if one of the calls matches the expected subset - assert any(all(expected_extra[k] == actual_extra.get(k) - for k in expected_extra) for actual_extra in calls) + assert any( + all(expected_extra[k] == actual_extra.get(k) for k in expected_extra) + for actual_extra in calls + ) mock_run.assert_called_once()