From 77311339e215f6eb8da58baedacef2817d6ea2a1 Mon Sep 17 00:00:00 2001 From: Alejandro Ponce Date: Wed, 18 Dec 2024 14:22:32 +0100 Subject: [PATCH 1/2] fix: FIM not caching correctly non-python files Closes: #405 The implemented cache was not working correctly because of the way the context is added for FIM requests in other programming languages which are not python. The context for the LLM is provided as single line comments. In Python this means lines which start with the character `#`. Other languages may have other starting sequence for single line comments, e.g. Javascript uses `//`. This PR changes the regex to detect the paths for other languages. --- src/codegate/codegate_logging.py | 6 +- src/codegate/db/connection.py | 89 ++--------- src/codegate/db/fim_cache.py | 136 +++++++++++++++++ src/codegate/providers/copilot/provider.py | 18 +-- tests/db/test_connection.py | 36 ----- tests/db/test_fim_cache.py | 169 +++++++++++++++++++++ 6 files changed, 331 insertions(+), 123 deletions(-) create mode 100644 src/codegate/db/fim_cache.py delete mode 100644 tests/db/test_connection.py create mode 100644 tests/db/test_fim_cache.py diff --git a/src/codegate/codegate_logging.py b/src/codegate/codegate_logging.py index 5a15dec1..5cb3c4a9 100644 --- a/src/codegate/codegate_logging.py +++ b/src/codegate/codegate_logging.py @@ -50,10 +50,10 @@ def _missing_(cls, value: str) -> Optional["LogFormat"]: def add_origin(logger, log_method, event_dict): # Add 'origin' if it's bound to the logger but not explicitly in the event dict - if 'origin' not in event_dict and hasattr(logger, '_context'): - origin = logger._context.get('origin') + if "origin" not in event_dict and hasattr(logger, "_context"): + origin = logger._context.get("origin") if origin: - event_dict['origin'] = origin + event_dict["origin"] = origin return event_dict diff --git a/src/codegate/db/connection.py b/src/codegate/db/connection.py index f4c0c064..f3c57c1a 100644 --- a/src/codegate/db/connection.py +++ b/src/codegate/db/connection.py @@ -1,8 +1,5 @@ import asyncio -import hashlib import json -import re -from datetime import timedelta from pathlib import Path from typing import List, Optional @@ -11,7 +8,7 @@ from sqlalchemy import text from sqlalchemy.ext.asyncio import create_async_engine -from codegate.config import Config +from codegate.db.fim_cache import FimCache from codegate.db.models import Alert, Output, Prompt from codegate.db.queries import ( AsyncQuerier, @@ -22,7 +19,7 @@ logger = structlog.get_logger("codegate") alert_queue = asyncio.Queue() -fim_entries = {} +fim_cache = FimCache() class DbCodeGate: @@ -183,47 +180,6 @@ async def record_alerts(self, alerts: List[Alert]) -> List[Alert]: logger.debug(f"Recorded alerts: {recorded_alerts}") return recorded_alerts - def _extract_request_message(self, request: str) -> Optional[dict]: - """Extract the user message from the FIM request""" - try: - parsed_request = json.loads(request) - except Exception as e: - logger.exception(f"Failed to extract request message: {request}", error=str(e)) - return None - - messages = [message for message in parsed_request["messages"] if message["role"] == "user"] - if len(messages) != 1: - logger.warning(f"Expected one user message, found {len(messages)}.") - return None - - content_message = messages[0].get("content") - return content_message - - def _create_hash_key(self, message: str, provider: str) -> str: - """Creates a hash key from the message and includes the provider""" - # Try to extract the path from the FIM message. The path is in FIM request in these formats: - # folder/testing_file.py - # Path: file3.py - pattern = r"^#.*?\b([a-zA-Z0-9_\-\/]+\.\w+)\b" - matches = re.findall(pattern, message, re.MULTILINE) - # If no path is found, hash the entire prompt message. - if not matches: - logger.warning("No path found in messages. Creating hash cache from message.") - message_to_hash = f"{message}-{provider}" - else: - # Copilot puts the path at the top of the file. Continue providers contain - # several paths, the one in which the fim is triggered is the last one. - if provider == "copilot": - filepath = matches[0] - else: - filepath = matches[-1] - message_to_hash = f"{filepath}-{provider}" - - logger.debug(f"Message to hash: {message_to_hash}") - hashed_content = hashlib.sha256(message_to_hash.encode("utf-8")).hexdigest() - logger.debug(f"Hashed contnet: {hashed_content}") - return hashed_content - def _should_record_context(self, context: Optional[PipelineContext]) -> bool: """Check if the context should be recorded in DB""" if context is None or context.metadata.get("stored_in_db", False): @@ -237,37 +193,22 @@ def _should_record_context(self, context: Optional[PipelineContext]) -> bool: if context.input_request.type != "fim": return True - # Couldn't process the user message. Skip creating a mapping entry. - message = self._extract_request_message(context.input_request.request) - if message is None: - logger.warning(f"Couldn't read FIM message: {message}. Will not record to DB.") - return False - - hash_key = self._create_hash_key(message, context.input_request.provider) - old_timestamp = fim_entries.get(hash_key, None) - if old_timestamp is None: - fim_entries[hash_key] = context.input_request.timestamp - return True + return fim_cache.could_store_fim_request(context) - elapsed_seconds = (context.input_request.timestamp - old_timestamp).total_seconds() - if elapsed_seconds < Config.get_config().max_fim_hash_lifetime: + async def record_context(self, context: Optional[PipelineContext]) -> None: + try: + if not self._should_record_context(context): + return + await self.record_request(context.input_request) + await self.record_outputs(context.output_responses) + await self.record_alerts(context.alerts_raised) + context.metadata["stored_in_db"] = True logger.info( - f"Skipping DB context recording. " - f"Elapsed time since last FIM cache: {timedelta(seconds=elapsed_seconds)}." + f"Recorded context in DB. Output chunks: {len(context.output_responses)}. " + f"Alerts: {len(context.alerts_raised)}." ) - return False - - async def record_context(self, context: Optional[PipelineContext]) -> None: - if not self._should_record_context(context): - return - await self.record_request(context.input_request) - await self.record_outputs(context.output_responses) - await self.record_alerts(context.alerts_raised) - context.metadata["stored_in_db"] = True - logger.info( - f"Recorded context in DB. Output chunks: {len(context.output_responses)}. " - f"Alerts: {len(context.alerts_raised)}." - ) + except Exception as e: + logger.error(f"Failed to record context: {context}.", error=str(e)) class DbReader(DbCodeGate): diff --git a/src/codegate/db/fim_cache.py b/src/codegate/db/fim_cache.py new file mode 100644 index 00000000..50ce5c6c --- /dev/null +++ b/src/codegate/db/fim_cache.py @@ -0,0 +1,136 @@ +import datetime +import hashlib +import json +import re +from typing import Dict, List, Optional + +import structlog +from pydantic import BaseModel + +from codegate.config import Config +from codegate.db.models import Alert +from codegate.pipeline.base import AlertSeverity, PipelineContext + +logger = structlog.get_logger("codegate") + + +class CachedFim(BaseModel): + + timestamp: datetime.datetime + critical_alerts: List[Alert] + + +class FimCache: + + def __init__(self): + self.cache: Dict[str, CachedFim] = {} + + def _extract_message_from_fim_request(self, request: str) -> Optional[str]: + """Extract the user message from the FIM request""" + try: + parsed_request = json.loads(request) + except Exception as e: + logger.error(f"Failed to extract request message: {request}", error=str(e)) + return None + + if not isinstance(parsed_request, dict): + logger.warning(f"Expected a dictionary, got {type(parsed_request)}.") + return None + + messages = [ + message + for message in parsed_request.get("messages", []) + if isinstance(message, dict) and message.get("role", "") == "user" + ] + if len(messages) != 1: + logger.warning(f"Expected one user message, found {len(messages)}.") + return None + + content_message = messages[0].get("content") + return content_message + + def _match_filepath(self, message: str, provider: str) -> Optional[str]: + # Try to extract the path from the FIM message. The path is in FIM request as a comment: + # folder/testing_file.py + # Path: file3.py + # // Path: file3.js <-- Javascript + pattern = r"^(#|//|