bits of pydantic

This commit is contained in:
Kwindla Hultman Kramer
2024-10-06 08:13:37 -07:00
parent 7ae3c420f4
commit 71fe09f7f0
4 changed files with 47 additions and 24 deletions

View File

@@ -17,8 +17,8 @@ from pipecat.processors.aggregators.openai_llm_context import (
) )
from pipecat.services.openai_realtime_beta import ( from pipecat.services.openai_realtime_beta import (
OpenAILLMServiceRealtimeBeta, OpenAILLMServiceRealtimeBeta,
OpenAITurnDetection, TurnDetection,
RealtimeSessionProperties, SessionProperties,
) )
from pipecat.transports.services.daily import DailyParams, DailyTransport from pipecat.transports.services.daily import DailyParams, DailyTransport
from pipecat.vad.silero import SileroVADAnalyzer from pipecat.vad.silero import SileroVADAnalyzer
@@ -83,8 +83,8 @@ async def main():
), ),
) )
session_properties = RealtimeSessionProperties( session_properties = SessionProperties(
turn_detection=OpenAITurnDetection(silence_duration_ms=1000), turn_detection=TurnDetection(silence_duration_ms=1000),
tools=tools, tools=tools,
instructions=""" instructions="""
Your knowledge cutoff is 2023-10. You are a helpful and friendly AI. Your knowledge cutoff is 2023-10. You are a helpful and friendly AI.

View File

@@ -0,0 +1,2 @@
from .llm_and_context import OpenAILLMServiceRealtimeBeta
from .client_events import SessionProperties, TurnDetection

View File

@@ -0,0 +1,33 @@
from pydantic import BaseModel
from typing import Dict, List, Optional, Literal
class InputAudioTranscription(BaseModel):
model: Optional[str] = "whisper-1"
class TurnDetection(BaseModel):
type: Optional[Literal["server_vad"]] = "server_vad"
threshold: Optional[float] = 0.5
prefix_padding_ms: Optional[int] = 300
silence_duration_ms: Optional[int] = 800
class SessionProperties(BaseModel):
modalities: Optional[List[Literal["text", "audio"]]] = ["text", "audio"]
instructions: Optional[str] = None
voice: Optional[str] = "alloy"
input_audio_format: Optional[Literal["pcm16", "g711_ulaw", "g711_alaw"]] = "pcm16"
output_audio_format: Optional[Literal["pcm16", "g711_ulaw", "g711_alaw"]] = "pcm16"
input_audio_transcription: Optional[InputAudioTranscription] = InputAudioTranscription()
turn_detection: Optional[TurnDetection] = TurnDetection()
tools: Optional[List[Dict]] = []
tool_choice: Optional[Literal["auto", "none", "required"]] = "auto"
temperature: Optional[float] = 0.8
max_response_output_tokens: Optional[int] = 4096
class SessionUpdateEvent(BaseModel):
event_id: str
type: Literal["session.update"]
session: SessionProperties

View File

@@ -1,11 +1,11 @@
import asyncio import asyncio
import base64 import base64
import random import random
import traceback
import json import json
import websockets import websockets
from copy import deepcopy from copy import deepcopy
from typing import List, Optional
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
@@ -37,6 +37,8 @@ from pipecat.processors.aggregators.openai_llm_context import (
OpenAILLMContextFrame, OpenAILLMContextFrame,
) )
from . import client_events as events
from loguru import logger from loguru import logger
# temp: websocket logger # temp: websocket logger
@@ -112,29 +114,13 @@ class OpenAITurnDetection(BaseModel):
) )
class RealtimeSessionProperties(BaseModel):
modalities: List[str] = Field(default=["text", "audio"])
instructions: str = Field(default="")
voice: str = Field(default="alloy")
input_audio_format: str = Field(default="pcm16")
output_audio_format: str = Field(default="pcm16")
input_audio_transcription: Optional[OpenAIInputTranscription] = Field(
default=OpenAIInputTranscription()
)
turn_detection: Optional[OpenAITurnDetection] = Field(default=None)
tools: List[dict] = Field(default=[])
tool_choice: str = Field(default="auto")
temperature: float = Field(default=0.8)
max_response_output_tokens: int = Field(default=4096)
class OpenAILLMServiceRealtimeBeta(LLMService): class OpenAILLMServiceRealtimeBeta(LLMService):
def __init__( def __init__(
self, self,
*, *,
api_key: str, api_key: str,
base_url="wss://api.openai.com/v1/realtime?model=gpt-4o-realtime-preview-2024-10-01", base_url="wss://api.openai.com/v1/realtime?model=gpt-4o-realtime-preview-2024-10-01",
session_properties: RealtimeSessionProperties = RealtimeSessionProperties(), session_properties: events.SessionProperties = events.SessionProperties(),
**kwargs, **kwargs,
): ):
super().__init__(base_url=base_url, **kwargs) super().__init__(base_url=base_url, **kwargs)
@@ -175,7 +161,7 @@ class OpenAILLMServiceRealtimeBeta(LLMService):
await self._ws_send( await self._ws_send(
{ {
"type": "session.update", "type": "session.update",
"session": self._session_properties.dict(), "session": self._session_properties.dict(exclude_none=True),
} }
) )
@@ -223,6 +209,8 @@ class OpenAILLMServiceRealtimeBeta(LLMService):
if not msg: if not msg:
continue continue
if msg["type"] == "session.created": if msg["type"] == "session.created":
# session.created is received right after connecting. send a message
# to configure the session properties.
await self.update_session_properties() await self.update_session_properties()
elif msg["type"] == "session.updated": elif msg["type"] == "session.updated":
self._session_properties = msg["session"] self._session_properties = msg["session"]
@@ -326,7 +314,7 @@ class OpenAILLMServiceRealtimeBeta(LLMService):
except asyncio.CancelledError: except asyncio.CancelledError:
pass pass
except Exception as e: except Exception as e:
logger.error(f"{self} exception: {e}") logger.error(f"{self} exception: {e}\n\nStack trace:\n{traceback.format_exc()}")
async def _handle_function_call_items(self, items): async def _handle_function_call_items(self, items):
total_items = len(items) total_items = len(items)