File size: 6,631 Bytes
6b62834
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
"""Core domain models for Agentic RAG."""

import time
import uuid
from enum import Enum
from typing import Any, Optional

from pydantic import BaseModel, Field


# ──────────────────────────────────────────────
# Message & Conversation Models
# ──────────────────────────────────────────────

class MessageRole(str, Enum):
    SYSTEM = "system"
    USER = "user"
    ASSISTANT = "assistant"
    TOOL = "tool"


class ToolCall(BaseModel):
    """A tool call made by the LLM."""
    id: str = Field(default_factory=lambda: f"call_{uuid.uuid4().hex[:12]}")
    name: str
    arguments: dict[str, Any] = Field(default_factory=dict)


class ToolCallResult(BaseModel):
    """Result of a tool call."""
    call_id: str
    name: str
    result: Any
    error: Optional[str] = None


class Message(BaseModel):
    """A single message in a conversation."""
    role: MessageRole
    content: str | list[dict[str, Any]]  # text or multimodal content
    tool_calls: list[ToolCall] = Field(default_factory=list)
    tool_call_id: Optional[str] = None
    timestamp: float = Field(default_factory=time.time)

    @classmethod
    def system(cls, content: str) -> "Message":
        return cls(role=MessageRole.SYSTEM, content=content)

    @classmethod
    def user(cls, content: str) -> "Message":
        return cls(role=MessageRole.USER, content=content)

    @classmethod
    def assistant(cls, content: str, tool_calls: list[ToolCall] | None = None) -> "Message":
        return cls(role=MessageRole.ASSISTANT, content=content,
                   tool_calls=tool_calls or [])

    @classmethod
    def tool(cls, content: str, tool_call_id: str) -> "Message":
        return cls(role=MessageRole.TOOL, content=content, tool_call_id=tool_call_id)


# ──────────────────────────────────────────────
# LLM Models
# ──────────────────────────────────────────────

class ToolDefinition(BaseModel):
    """Tool definition for LLM function calling."""
    name: str
    description: str
    parameters: dict[str, Any]  # JSON Schema


class LLMResponse(BaseModel):
    """Response from an LLM provider."""
    content: str
    tool_calls: list[ToolCall] = Field(default_factory=list)
    stop_reason: Optional[str] = None
    usage: dict[str, int] = Field(default_factory=dict)


class LLMChunk(BaseModel):
    """Streaming chunk from an LLM provider."""
    content_delta: str = ""
    tool_call_delta: Optional[dict] = None
    stop_reason: Optional[str] = None


# ──────────────────────────────────────────────
# Multimodal Input Models
# ──────────────────────────────────────────────

class MediaType(str, Enum):
    IMAGE = "image"
    VIDEO = "video"
    AUDIO = "audio"
    TEXT = "text"


class MultimodalInput(BaseModel):
    """User input that may contain multiple modalities."""
    text: Optional[str] = None
    images: list[str] = Field(default_factory=list)    # base64 or file paths
    audio: Optional[str] = None                         # base64 or file path
    video: Optional[str] = None                         # file path
    metadata: dict[str, Any] = Field(default_factory=dict)


class ProcessedContent(BaseModel):
    """Content after multimodal processing."""
    text: str
    source_types: list[MediaType] = Field(default_factory=list)
    metadata: dict[str, Any] = Field(default_factory=dict)


# ──────────────────────────────────────────────
# Agent Models
# ──────────────────────────────────────────────

class AgentEventType(str, Enum):
    THOUGHT = "thought"
    ACTION = "action"
    OBSERVATION = "observation"
    TEXT_DELTA = "text_delta"
    TOOL_CALL_START = "tool_call_start"
    TOOL_CALL_RESULT = "tool_call_result"
    ERROR = "error"
    DONE = "done"


class AgentEvent(BaseModel):
    """An event emitted during agent execution."""
    event_type: AgentEventType
    data: dict[str, Any] = Field(default_factory=dict)
    turn_id: str = ""
    timestamp: float = Field(default_factory=time.time)


class AgentInput(BaseModel):
    """Input to an agent."""
    messages: list[Message] = Field(default_factory=list)
    query: str = ""
    multimodal: Optional[MultimodalInput] = None
    parameters: dict[str, Any] = Field(default_factory=dict)


class AgentOutput(BaseModel):
    """Output from an agent."""
    messages: list[Message] = Field(default_factory=list)
    final_answer: str = ""
    tool_calls_made: list[ToolCallResult] = Field(default_factory=list)
    usage: dict[str, int] = Field(default_factory=dict)
    iterations: int = 0


# ──────────────────────────────────────────────
# RAG Models
# ──────────────────────────────────────────────

class Document(BaseModel):
    """A document in the knowledge base."""
    id: str = Field(default_factory=lambda: uuid.uuid4().hex)
    text: str
    metadata: dict[str, Any] = Field(default_factory=dict)
    embedding: list[float] = Field(default_factory=list)


class RetrievalResult(BaseModel):
    """Result of a retrieval operation."""
    documents: list[Document] = Field(default_factory=list)
    scores: list[float] = Field(default_factory=list)
    query: str = ""


# ──────────────────────────────────────────────
# Session Models
# ──────────────────────────────────────────────

class Session(BaseModel):
    """A user session."""
    id: str = Field(default_factory=lambda: uuid.uuid4().hex)
    user_id: str = "default"
    created_at: float = Field(default_factory=time.time)
    expires_at: float = Field(default_factory=lambda: time.time() + 3600)
    metadata: dict[str, Any] = Field(default_factory=dict)