add FileAPI to gemini.py
This commit is contained in:
@@ -71,6 +71,8 @@ from pipecat.utils.time import time_now_iso8601
|
|||||||
from pipecat.utils.tracing.service_decorators import traced_gemini_live, traced_stt
|
from pipecat.utils.tracing.service_decorators import traced_gemini_live, traced_stt
|
||||||
|
|
||||||
from . import events
|
from . import events
|
||||||
|
from .audio_transcriber import AudioTranscriber
|
||||||
|
from .file_api import GeminiFileAPI
|
||||||
|
|
||||||
try:
|
try:
|
||||||
import websockets
|
import websockets
|
||||||
@@ -218,6 +220,29 @@ class GeminiMultimodalLiveContext(OpenAILLMContext):
|
|||||||
system_instruction += str(content)
|
system_instruction += str(content)
|
||||||
return system_instruction
|
return system_instruction
|
||||||
|
|
||||||
|
def add_file_reference(self, file_uri: str, mime_type: str, text: Optional[str] = None):
|
||||||
|
"""Add a file reference to the context.
|
||||||
|
|
||||||
|
This adds a user message with a file reference that will be sent during context initialization.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
file_uri: URI of the uploaded file
|
||||||
|
mime_type: MIME type of the file
|
||||||
|
text: Optional text prompt to accompany the file
|
||||||
|
"""
|
||||||
|
# Create parts list with file reference
|
||||||
|
parts = []
|
||||||
|
if text:
|
||||||
|
parts.append({"type": "text", "text": text})
|
||||||
|
|
||||||
|
# Add file reference part
|
||||||
|
parts.append({"type": "file_data", "file_data": {"mime_type": mime_type, "file_uri": file_uri}})
|
||||||
|
|
||||||
|
# Add to messages
|
||||||
|
message = {"role": "user", "content": parts}
|
||||||
|
self.messages.append(message)
|
||||||
|
logger.info(f"Added file reference to context: {file_uri}")
|
||||||
|
|
||||||
def get_messages_for_initializing_history(self):
|
def get_messages_for_initializing_history(self):
|
||||||
"""Get messages formatted for Gemini history initialization.
|
"""Get messages formatted for Gemini history initialization.
|
||||||
|
|
||||||
@@ -242,6 +267,14 @@ class GeminiMultimodalLiveContext(OpenAILLMContext):
|
|||||||
for part in content:
|
for part in content:
|
||||||
if part.get("type") == "text":
|
if part.get("type") == "text":
|
||||||
parts.append({"text": part.get("text")})
|
parts.append({"text": part.get("text")})
|
||||||
|
elif part.get("type") == "file_data":
|
||||||
|
file_data = part.get("file_data", {})
|
||||||
|
parts.append({
|
||||||
|
"fileData": {
|
||||||
|
"mimeType": file_data.get("mime_type"),
|
||||||
|
"fileUri": file_data.get("file_uri")
|
||||||
|
}
|
||||||
|
})
|
||||||
else:
|
else:
|
||||||
logger.warning(f"Unsupported content type: {str(part)[:80]}")
|
logger.warning(f"Unsupported content type: {str(part)[:80]}")
|
||||||
else:
|
else:
|
||||||
@@ -432,6 +465,62 @@ class GeminiMultimodalLiveLLMService(LLMService):
|
|||||||
# Overriding the default adapter to use the Gemini one.
|
# Overriding the default adapter to use the Gemini one.
|
||||||
adapter_class = GeminiLLMAdapter
|
adapter_class = GeminiLLMAdapter
|
||||||
|
|
||||||
|
"""Gemini Live LLM Service with multimodal capabilities including File API support.
|
||||||
|
|
||||||
|
This service implements the Gemini Multimodal Live API with support for:
|
||||||
|
- Audio input and output
|
||||||
|
- Image/video input
|
||||||
|
- File API (upload, reference, and management)
|
||||||
|
- Tools/function calling
|
||||||
|
|
||||||
|
Example usage of File API:
|
||||||
|
```python
|
||||||
|
# Initialize the service
|
||||||
|
gemini_service = GeminiMultimodalLiveLLMService(api_key="YOUR_API_KEY")
|
||||||
|
|
||||||
|
# Upload a file from the client
|
||||||
|
file_path = "/path/to/user_uploaded_file.pdf"
|
||||||
|
file_info = await gemini_service.file_api.upload_file(file_path)
|
||||||
|
|
||||||
|
# Get file URI and mime type from response
|
||||||
|
file_uri = file_info["file"]["uri"]
|
||||||
|
mime_type = "application/pdf" # Set appropriate MIME type
|
||||||
|
|
||||||
|
# When starting a new bot session:
|
||||||
|
# 1. Initialize the context
|
||||||
|
context = GeminiMultimodalLiveContext()
|
||||||
|
|
||||||
|
# 2. Add file reference to context BEFORE starting the conversation
|
||||||
|
context.add_file_reference(
|
||||||
|
file_uri=file_uri,
|
||||||
|
mime_type=mime_type,
|
||||||
|
text="Please analyze this document"
|
||||||
|
)
|
||||||
|
|
||||||
|
# 3. Now set the context to start the conversation with file reference included
|
||||||
|
await gemini_service.set_context(context)
|
||||||
|
|
||||||
|
# Gemini now has access to the file reference in its context window
|
||||||
|
# The file URI remains valid for 48 hours before Google deletes it
|
||||||
|
|
||||||
|
# Optional: List all files for this user
|
||||||
|
files = await gemini_service.file_api.list_files()
|
||||||
|
|
||||||
|
# Optional: Get metadata for a specific file
|
||||||
|
file_metadata = await gemini_service.file_api.get_file(file_info["file"]["name"])
|
||||||
|
|
||||||
|
# Optional: Delete a file when no longer needed
|
||||||
|
await gemini_service.file_api.delete_file(file_info["file"]["name"])
|
||||||
|
```
|
||||||
|
|
||||||
|
Notes:
|
||||||
|
- Files are stored for 48 hours on Google's servers
|
||||||
|
- Maximum file size is 2GB
|
||||||
|
- Total storage per project is 20GB
|
||||||
|
- File references should be added to the context BEFORE starting the conversation
|
||||||
|
- The same file reference can be reused for multiple sessions within the 48-hour window
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -445,6 +534,7 @@ class GeminiMultimodalLiveLLMService(LLMService):
|
|||||||
tools: Optional[Union[List[dict], ToolsSchema]] = None,
|
tools: Optional[Union[List[dict], ToolsSchema]] = None,
|
||||||
params: Optional[InputParams] = None,
|
params: Optional[InputParams] = None,
|
||||||
inference_on_context_initialization: bool = True,
|
inference_on_context_initialization: bool = True,
|
||||||
|
file_api_base_url: str = "https://generativelanguage.googleapis.com/v1beta/files",
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
"""Initialize the Gemini Multimodal Live LLM service.
|
"""Initialize the Gemini Multimodal Live LLM service.
|
||||||
@@ -523,6 +613,9 @@ class GeminiMultimodalLiveLLMService(LLMService):
|
|||||||
"extra": params.extra if isinstance(params.extra, dict) else {},
|
"extra": params.extra if isinstance(params.extra, dict) else {},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# Initialize the File API client
|
||||||
|
self.file_api = GeminiFileAPI(api_key=api_key, base_url=file_api_base_url)
|
||||||
|
|
||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
"""Check if the service can generate usage metrics.
|
"""Check if the service can generate usage metrics.
|
||||||
|
|
||||||
@@ -938,7 +1031,7 @@ class GeminiMultimodalLiveLLMService(LLMService):
|
|||||||
self._needs_turn_complete_message = True
|
self._needs_turn_complete_message = True
|
||||||
|
|
||||||
async def _create_single_response(self, messages_list):
|
async def _create_single_response(self, messages_list):
|
||||||
# refactor to combine this logic with same logic in GeminiMultimodalLiveContext
|
# Refactor to combine this logic with same logic in GeminiMultimodalLiveContext
|
||||||
messages = []
|
messages = []
|
||||||
for item in messages_list:
|
for item in messages_list:
|
||||||
role = item.get("role")
|
role = item.get("role")
|
||||||
@@ -957,6 +1050,14 @@ class GeminiMultimodalLiveLLMService(LLMService):
|
|||||||
for part in content:
|
for part in content:
|
||||||
if part.get("type") == "text":
|
if part.get("type") == "text":
|
||||||
parts.append({"text": part.get("text")})
|
parts.append({"text": part.get("text")})
|
||||||
|
elif part.get("type") == "file_data":
|
||||||
|
file_data = part.get("file_data", {})
|
||||||
|
parts.append({
|
||||||
|
"fileData": {
|
||||||
|
"mimeType": file_data.get("mime_type"),
|
||||||
|
"fileUri": file_data.get("file_uri")
|
||||||
|
}
|
||||||
|
})
|
||||||
else:
|
else:
|
||||||
logger.warning(f"Unsupported content type: {str(part)[:80]}")
|
logger.warning(f"Unsupported content type: {str(part)[:80]}")
|
||||||
else:
|
else:
|
||||||
|
|||||||
Reference in New Issue
Block a user