diff --git a/src/codegate/pipeline/extract_snippets/output.py b/src/codegate/pipeline/extract_snippets/output.py index cd7391f8..895eee7b 100644 --- a/src/codegate/pipeline/extract_snippets/output.py +++ b/src/codegate/pipeline/extract_snippets/output.py @@ -98,7 +98,7 @@ async def process_chunk( input_context: Optional[PipelineContext] = None, ) -> list[ModelResponse]: """Process a single chunk of the stream""" - if not chunk.choices[0].delta.content: + if len(chunk.choices) == 0 or not chunk.choices[0].delta.content: return [chunk] # Get current content plus this new chunk diff --git a/src/codegate/pipeline/secrets/secrets.py b/src/codegate/pipeline/secrets/secrets.py index f12eda5f..a1989024 100644 --- a/src/codegate/pipeline/secrets/secrets.py +++ b/src/codegate/pipeline/secrets/secrets.py @@ -262,7 +262,7 @@ async def process_chunk( if input_context.sensitive.session_id == "": raise ValueError("Session ID not found in input context") - if not chunk.choices[0].delta.content: + if len(chunk.choices) == 0 or not chunk.choices[0].delta.content: return [chunk] # Check the buffered content diff --git a/src/codegate/providers/copilot/provider.py b/src/codegate/providers/copilot/provider.py index cfbf76fe..33eb823d 100644 --- a/src/codegate/providers/copilot/provider.py +++ b/src/codegate/providers/copilot/provider.py @@ -1,5 +1,4 @@ import asyncio -import json import re import ssl from dataclasses import dataclass @@ -7,6 +6,7 @@ from urllib.parse import unquote, urljoin, urlparse import structlog +from litellm.types.utils import Delta, ModelResponse, StreamingChoices from codegate.ca.codegate_ca import CertificateAuthority from codegate.config import Config @@ -559,32 +559,75 @@ def __init__(self, proxy: CopilotProvider): self.headers_sent = False self.sse_processor: Optional[SSEProcessor] = None self.output_pipeline_instance: Optional[OutputPipelineInstance] = None + self.stream_queue: Optional[asyncio.Queue] = None def connection_made(self, transport: asyncio.Transport) -> None: """Handle successful connection to target""" self.transport = transport self.proxy.target_transport = transport - def _process_chunk(self, chunk: bytes): - records = self.sse_processor.process_chunk(chunk) + async def _process_stream(self): + try: - for record in records: - if record["type"] == "done": - sse_data = b"data: [DONE]\n\n" - # Add chunk size for DONE message too - chunk_size = hex(len(sse_data))[2:] + "\r\n" - self._proxy_transport_write(chunk_size.encode()) - self._proxy_transport_write(sse_data) - self._proxy_transport_write(b"\r\n") - # Now send the final zero chunk - self._proxy_transport_write(b"0\r\n\r\n") - else: - sse_data = f"data: {json.dumps(record['content'])}\n\n".encode("utf-8") + async def stream_iterator(): + while True: + incoming_record = await self.stream_queue.get() + record_content = incoming_record.get("content", {}) + + streaming_choices = [] + for choice in record_content.get("choices", []): + streaming_choices.append( + StreamingChoices( + finish_reason=choice.get("finish_reason", None), + index=0, + delta=Delta( + content=choice.get("delta", {}).get("content"), role="assistant" + ), + logprobs=None, + ) + ) + + # Convert record to ModelResponse + mr = ModelResponse( + id=record_content.get("id", ""), + choices=streaming_choices, + created=record_content.get("created", 0), + model=record_content.get("model", ""), + object="chat.completion.chunk", + ) + yield mr + + async for record in self.output_pipeline_instance.process_stream(stream_iterator()): + chunk = record.model_dump_json(exclude_none=True, exclude_unset=True) + sse_data = f"data:{chunk}\n\n".encode("utf-8") chunk_size = hex(len(sse_data))[2:] + "\r\n" self._proxy_transport_write(chunk_size.encode()) self._proxy_transport_write(sse_data) self._proxy_transport_write(b"\r\n") + sse_data = b"data: [DONE]\n\n" + # Add chunk size for DONE message too + chunk_size = hex(len(sse_data))[2:] + "\r\n" + self._proxy_transport_write(chunk_size.encode()) + self._proxy_transport_write(sse_data) + self._proxy_transport_write(b"\r\n") + # Now send the final zero chunk + self._proxy_transport_write(b"0\r\n\r\n") + + except Exception as e: + logger.error(f"Error processing stream: {e}") + + def _process_chunk(self, chunk: bytes): + records = self.sse_processor.process_chunk(chunk) + + for record in records: + if self.stream_queue is None: + # Initialize queue and start processing task on first record + self.stream_queue = asyncio.Queue() + self.processing_task = asyncio.create_task(self._process_stream()) + + self.stream_queue.put_nowait(record) + def _proxy_transport_write(self, data: bytes): self.proxy.transport.write(data) diff --git a/src/codegate/providers/copilot/streaming.py b/src/codegate/providers/copilot/streaming.py index f904be25..60e6ca71 100644 --- a/src/codegate/providers/copilot/streaming.py +++ b/src/codegate/providers/copilot/streaming.py @@ -13,9 +13,6 @@ def __init__(self): self.size_written = False def process_chunk(self, chunk: bytes) -> list: - print("BUFFER AT START") - print(self.buffer) - print("BUFFER AT START - END") # Skip any chunk size lines (hex number followed by \r\n) try: chunk_str = chunk.decode("utf-8") @@ -25,13 +22,12 @@ def process_chunk(self, chunk: bytes) -> list: continue self.buffer += line except UnicodeDecodeError: - print("Failed to decode chunk") + logger.error("Failed to decode chunk") records = [] while True: record_end = self.buffer.find("\n\n") if record_end == -1: - print(f"REMAINING BUFFER {self.buffer}") break record = self.buffer[:record_end]