diff --git a/src/codegate/pipeline/base.py b/src/codegate/pipeline/base.py index b8875dda..dc7bac53 100644 --- a/src/codegate/pipeline/base.py +++ b/src/codegate/pipeline/base.py @@ -15,15 +15,13 @@ class CodeSnippet: code: The actual code content """ - language: str + language: Optional[str] + filepath: Optional[str] code: str def __post_init__(self): - if not self.language or not self.language.strip(): - raise ValueError("Language must not be empty") - if not self.code or not self.code.strip(): - raise ValueError("Code must not be empty") - self.language = self.language.strip().lower() + if self.language is not None: + self.language = self.language.strip().lower() @dataclass @@ -57,6 +55,7 @@ class PipelineResult: request: Optional[ChatCompletionRequest] = None response: Optional[PipelineResponse] = None + context: Optional[PipelineContext] = None error_message: Optional[str] = None def shortcuts_processing(self) -> bool: @@ -165,4 +164,7 @@ async def process_request( if result.request is not None: current_request = result.request - return PipelineResult(request=current_request) + if result.context is not None: + context = result.context + + return PipelineResult(request=current_request, context=context) diff --git a/src/codegate/pipeline/extract_snippets/__init__.py b/src/codegate/pipeline/extract_snippets/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/src/codegate/pipeline/extract_snippets/extract_snippets.py b/src/codegate/pipeline/extract_snippets/extract_snippets.py new file mode 100644 index 00000000..a50460ee --- /dev/null +++ b/src/codegate/pipeline/extract_snippets/extract_snippets.py @@ -0,0 +1,131 @@ +import os +import re +from typing import List, Optional + +import structlog +from litellm.types.llms.openai import ChatCompletionRequest + +from codegate.pipeline.base import CodeSnippet, PipelineContext, PipelineResult, PipelineStep + +CODE_BLOCK_PATTERN = re.compile( + r"```(?:(?P\w+)\s+)?(?P[^\s\(]+)?(?:\s*\((?P[^)]+)\))?\n(?P(?:.|\n)*?)```" +) + +logger = structlog.get_logger("codegate") + +def ecosystem_from_filepath(filepath: str) -> Optional[str]: + """ + Determine language from filepath. + + Args: + filepath: Path to the file + + Returns: + Determined language based on file extension + """ + # Implement file extension to language mapping + extension_mapping = { + ".py": "python", + ".js": "javascript", + ".ts": "typescript", + ".tsx": "typescript", + ".go": "go", + ".rs": "rust", + ".java": "java", + } + + # Get the file extension + ext = os.path.splitext(filepath)[1].lower() + return extension_mapping.get(ext, None) + + +def ecosystem_from_message(message: str) -> Optional[str]: + """ + Determine language from message. + + Args: + message: The language from the message. Some extensions send a different + format where the language is present in the snippet, + e.g. "py /path/to/file (lineFrom-lineTo)" + + Returns: + Determined language based on message content + """ + language_mapping = { + "py": "python", + "js": "javascript", + "ts": "typescript", + "tsx": "typescript", + "go": "go", + } + return language_mapping.get(message, None) + + +def extract_snippets(message: str) -> List[CodeSnippet]: + """ + Extract code snippets from a message. + + Args: + message: Input text containing code snippets + + Returns: + List of extracted code snippets + """ + # Regular expression to find code blocks + + snippets: List[CodeSnippet] = [] + + # Find all code block matches + for match in CODE_BLOCK_PATTERN.finditer(message): + filename = match.group("filename") + content = match.group("content") + matched_language = match.group("language") + + # Determine language + lang = None + if matched_language: + lang = ecosystem_from_message(matched_language.strip()) + if lang is None and filename: + filename = filename.strip() + # Determine language from the filename + lang = ecosystem_from_filepath(filename) + + snippets.append(CodeSnippet(filepath=filename, code=content, language=lang)) + + return snippets + + +class CodeSnippetExtractor(PipelineStep): + """ + Pipeline step that merely extracts code snippets from the user message. + """ + + def __init__(self): + """Initialize the CodeSnippetExtractor pipeline step.""" + super().__init__() + + @property + def name(self) -> str: + return "code-snippet-extractor" + + async def process( + self, + request: ChatCompletionRequest, + context: PipelineContext, + ) -> PipelineResult: + last_user_message = self.get_last_user_message(request) + if not last_user_message: + return PipelineResult(request=request, context=context) + msg_content, _ = last_user_message + snippets = extract_snippets(msg_content) + + logger.info(f"Extracted {len(snippets)} code snippets from the user message") + + if len(snippets) > 0: + for snippet in snippets: + logger.debug(f"Code snippet: {snippet}") + context.add_code_snippet(snippet) + + return PipelineResult( + context=context, + ) diff --git a/src/codegate/server.py b/src/codegate/server.py index 92d673d6..631824bf 100644 --- a/src/codegate/server.py +++ b/src/codegate/server.py @@ -6,6 +6,7 @@ from codegate.config import Config from codegate.pipeline.base import PipelineStep, SequentialPipelineProcessor from codegate.pipeline.codegate_system_prompt.codegate import CodegateSystemPrompt +from codegate.pipeline.extract_snippets.extract_snippets import CodeSnippetExtractor from codegate.pipeline.version.version import CodegateVersion from codegate.providers.anthropic.provider import AnthropicProvider from codegate.providers.llamacpp.provider import LlamaCppProvider @@ -23,6 +24,7 @@ def init_app() -> FastAPI: steps: List[PipelineStep] = [ CodegateVersion(), + CodeSnippetExtractor(), CodegateSystemPrompt(Config.get_config().prompts.codegate_chat), # CodegateSecrets(), ] diff --git a/tests/pipeline/extract_snippets/test_extract_snippets.py b/tests/pipeline/extract_snippets/test_extract_snippets.py new file mode 100644 index 00000000..c6cd2411 --- /dev/null +++ b/tests/pipeline/extract_snippets/test_extract_snippets.py @@ -0,0 +1,245 @@ +from typing import List, NamedTuple + +import pytest +from litellm.types.llms.openai import ChatCompletionRequest + +from codegate.pipeline.base import CodeSnippet, PipelineContext +from codegate.pipeline.extract_snippets.extract_snippets import ( + CodeSnippetExtractor, + ecosystem_from_filepath, + extract_snippets, +) + + +class CodeSnippetTest(NamedTuple): + input_message: str + expected_count: int + expected: List[CodeSnippet] + + +@pytest.mark.parametrize( + "test_case", + [ + # Single Python code block without filename + CodeSnippetTest( + input_message=""": + Here's a Python snippet: + ``` + def hello(): + print("Hello, world!") + ``` + """, + expected_count=1, + expected=[ + CodeSnippet(language=None, filepath=None, code='print("Hello, world!")'), + ], + ), + # Single Python code block + CodeSnippetTest( + input_message=""" + Here's a Python snippet: + ```hello_world.py (8-13) + def hello(): + print("Hello, world!") + ``` + """, + expected_count=1, + expected=[ + CodeSnippet( + language="python", filepath="hello_world.py", code='print("Hello, world!")' + ), + ], + ), + # Single Python code block with a language identifier + CodeSnippetTest( + input_message=""" + Here's a Python snippet: + ```py goodbye_world.py (8-13) + def hello(): + print("Goodbye, world!") + ``` + """, + expected_count=1, + expected=[ + CodeSnippet( + language="python", filepath="goodbye_world.py", code='print("Goodbye, world!")' + ), + ], + ), + # Multiple code blocks with different languages + CodeSnippetTest( + input_message=""" + Python snippet: + ```main.py + def hello(): + print("Hello") + ``` + + JavaScript snippet: + ```script.js (1-3) + function greet() { + console.log("Hi"); + } + ``` + """, + expected_count=2, + expected=[ + CodeSnippet(language="python", filepath="main.py", code='print("Hello")'), + CodeSnippet( + language="javascript", + filepath="script.js", + code='console.log("Hi");', + ), + ], + ), + # No code blocks + CodeSnippetTest( + input_message="Just a plain text message", + expected_count=0, + expected=[], + ), + # unknown language + CodeSnippetTest( + input_message=""": + Here's a Perl snippet: + ```hello_world.pl + I'm a Perl script + ``` + """, + expected_count=1, + expected=[ + CodeSnippet( + language=None, + filepath="hello_world.pl", + code="I'm a Perl script", + ), + ], + ), + ], +) +def test_extract_snippets(test_case): + snippets = extract_snippets(test_case.input_message) + + assert len(snippets) == test_case.expected_count + + for expected, actual in zip(test_case.expected, snippets): + assert actual.language == expected.language + assert actual.filepath == expected.filepath + assert expected.code in actual.code + + +@pytest.mark.parametrize( + "filepath,expected", + [ + # Standard extensions + ("file.py", "python"), + ("script.js", "javascript"), + ("code.go", "go"), + ("app.ts", "typescript"), + ("component.tsx", "typescript"), + ("program.rs", "rust"), + ("App.java", "java"), + # Case insensitive + ("FILE.PY", "python"), + ("SCRIPT.JS", "javascript"), + # Full paths + ("/path/to/file.rs", "rust"), + ("C:\\Users\\name\\file.java", "java"), + ], +) +def test_valid_extensions(filepath, expected): + assert ecosystem_from_filepath(filepath) == expected + + +@pytest.mark.parametrize( + "filepath", + [ + # No extension + "README", + "script", + "README.txt", + # Unknown extensions + "file.xyz", + "unknown.extension", + ], +) +def test_no_or_unknown_extensions(filepath): + assert ecosystem_from_filepath(filepath) is None + + +@pytest.mark.asyncio +async def test_code_snippet_extractor(): + # Create a mock ChatCompletionRequest with a code snippet in the message + mock_request = { + "messages": [ + { + "role": "user", + "content": """ + Here's a Python snippet: + ```main.py + def hello(): + print("Hello, world!") + ``` + """, + } + ] + } + + # Create a pipeline context + context = PipelineContext() + + # Instantiate the extractor + extractor = CodeSnippetExtractor() + + # Process the request + result = await extractor.process(ChatCompletionRequest(**mock_request), context) + + # Assertions + assert result.context is not None + assert len(result.context.code_snippets) == 1 + + # Verify the extracted snippet + snippet = result.context.code_snippets[0] + assert snippet.language == "python" + assert snippet.filepath == "main.py" + assert 'print("Hello, world!")' in snippet.code + + +@pytest.mark.asyncio +async def test_code_snippet_extractor_no_snippets(): + # Create a mock ChatCompletionRequest without code snippets + mock_request = { + "messages": [{"role": "user", "content": "Just a plain text message with no code"}] + } + + # Create a pipeline context + context = PipelineContext() + + # Instantiate the extractor + extractor = CodeSnippetExtractor() + + # Process the request + result = await extractor.process(ChatCompletionRequest(**mock_request), context) + + # Assertions + assert result.context is not None + assert len(result.context.code_snippets) == 0 + + +@pytest.mark.asyncio +async def test_code_snippet_extractor_no_messages(): + # Create a mock ChatCompletionRequest with no messages + mock_request = {} + + # Create a pipeline context + context = PipelineContext() + + # Instantiate the extractor + extractor = CodeSnippetExtractor() + + # Process the request + result = await extractor.process(ChatCompletionRequest(**mock_request), context) + + # Assertions + assert result.context is not None + assert len(result.context.code_snippets) == 0