Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 13 additions & 0 deletions packages/gen/gen_ai_hub/orchestration_v2/models/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,3 +31,16 @@ def model_dump(self, **kwargs):
kwargs.setdefault("by_alias", True)
kwargs.setdefault("exclude_none", True)
return super().model_dump(**kwargs)

class ResponseBaseModel(BaseModel):
"""
Base model for API response models.

- `extra="allow"` allows unexpected fields in responses to be accepted,
since the external API might introduce new attributes in the response.
"""

model_config = ConfigDict(
extra="allow",
frozen=False,
)
10 changes: 5 additions & 5 deletions packages/gen/gen_ai_hub/orchestration_v2/models/embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@

from pydantic import Field

from gen_ai_hub.orchestration_v2.models.base import ABCBaseModel as BaseModel
from gen_ai_hub.orchestration_v2.models.base import ABCBaseModel as BaseModel, ResponseBaseModel
from gen_ai_hub.orchestration_v2.models.data_masking import MaskingModuleConfig


Expand Down Expand Up @@ -117,7 +117,7 @@ class EmbeddingsInput(BaseModel):
type_: Optional[EmbeddingsInputType] = Field(default=None, alias="type")


class EmbeddingsUsage(BaseModel):
class EmbeddingsUsage(ResponseBaseModel):
"""
Token usage information for the embeddings request.

Expand All @@ -129,7 +129,7 @@ class EmbeddingsUsage(BaseModel):
total_tokens: int


class EmbeddingResult(BaseModel):
class EmbeddingResult(ResponseBaseModel):
"""
A single embedding result.

Expand All @@ -143,7 +143,7 @@ class EmbeddingResult(BaseModel):
index: int


class EmbeddingsResponse(BaseModel):
class EmbeddingsResponse(ResponseBaseModel):
"""
The response from the embedding model, following OpenAI specification.

Expand All @@ -159,7 +159,7 @@ class EmbeddingsResponse(BaseModel):
usage: EmbeddingsUsage


class EmbeddingsPostResponse(BaseModel):
class EmbeddingsPostResponse(ResponseBaseModel):
"""
Response for an embeddings POST request.

Expand Down
9 changes: 5 additions & 4 deletions packages/gen/gen_ai_hub/orchestration_v2/models/message.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,11 +10,11 @@

from pydantic import field_validator, ValidationError

from gen_ai_hub.orchestration_v2.models.base import ABCBaseModel as BaseModel
from gen_ai_hub.orchestration_v2.models.base import ABCBaseModel as BaseModel, ResponseBaseModel
from gen_ai_hub.orchestration_v2.models.multimodal_items import ContentPart, ImageItem, TextPart, ImageUrl, ImagePart


class FunctionCall(BaseModel):
class FunctionCall(ResponseBaseModel):
"""
Represents a function call with its name and arguments.

Expand Down Expand Up @@ -44,7 +44,7 @@ def parse_arguments(self) -> dict:
return json.loads(self.arguments)


class MessageToolCall(BaseModel):
class MessageToolCall(ResponseBaseModel):
"""
The tool calls generated by the model, such as function calls.

Expand All @@ -71,6 +71,7 @@ class ReasoningBlock(BaseModel):
signature: str



class Role(str, Enum):
"""
Enumerates supported roles in LLM-based conversations.
Expand Down Expand Up @@ -180,7 +181,7 @@ class DeveloperChatMessage(BaseModel):
role: Role = Role.DEVELOPER
content: Union[str, List[TextPart]]

class ResponseChatMessage(BaseModel):
class ResponseChatMessage(ResponseBaseModel):
Comment thread
yamaceay marked this conversation as resolved.
Comment thread
yamaceay marked this conversation as resolved.
"""
Represents a response message in a conversation.

Expand Down
18 changes: 1 addition & 17 deletions packages/gen/gen_ai_hub/orchestration_v2/models/response.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,26 +6,10 @@
from pydantic import ConfigDict, Field

from gen_ai_hub.orchestration.models.response import ModuleResultsStreaming
from gen_ai_hub.orchestration_v2.models.base import ABCBaseModel as BaseModel
from gen_ai_hub.orchestration_v2.models.base import ResponseBaseModel
from gen_ai_hub.orchestration_v2.models.message import ChatMessage, FunctionCall, ResponseChatMessage


class ResponseBaseModel(BaseModel):
"""
Abstract base model that extends Pydantic's BaseModel and ABC.

- `extra="allow"` allows unexpected fields in responses to be accepted,
since the external API might introduce new attributes in the response.

This enforces consistent and safe serialization behavior across all
derived models.
"""

model_config = ConfigDict(
extra="allow",
frozen=False,
)


class CacheCreationTokenDetails(ResponseBaseModel):
"""
Expand Down
41 changes: 41 additions & 0 deletions packages/gen/tests/orchestration_v2/test_embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -734,6 +734,47 @@ def test_request_with_masking_allowlist(self):
allowlist = result["config"]["modules"]["masking"]["masking_providers"][0]["allowlist"]
self.assertEqual(allowlist, ["SAP", "Microsoft"])

class TestEmbeddingsResponseExtraFields(unittest.TestCase):
"""EmbeddingsUsage, EmbeddingResult, EmbeddingsResponse and EmbeddingsPostResponse
were switched to ResponseBaseModel and must silently accept unknown fields."""

def test_embeddings_usage_stores_extra_field(self):
usage = EmbeddingsUsage.model_validate({
"prompt_tokens": 10, "total_tokens": 10,
"extra_field": "extra",
})
self.assertEqual(usage.extra_field, "extra")

def test_embedding_result_stores_extra_field(self):
result = EmbeddingResult.model_validate({
"object": "embedding", "embedding": [0.1, 0.2], "index": 0,
"extra_field": "extra",
})
self.assertEqual(result.extra_field, "extra")

def test_embeddings_response_stores_extra_field(self):
response = EmbeddingsResponse.model_validate({
"object": "list",
"data": [{"object": "embedding", "embedding": [0.1], "index": 0}],
"model": "text-embedding-3-large",
"usage": {"prompt_tokens": 5, "total_tokens": 5},
"extra_field": "extra",
})
self.assertEqual(response.extra_field, "extra")

def test_embeddings_post_response_stores_extra_field(self):
response = EmbeddingsPostResponse.model_validate({
"request_id": "emb-req-1",
"final_result": {
"object": "list",
"data": [{"object": "embedding", "embedding": [0.1], "index": 0}],
"model": "text-embedding-3-large",
"usage": {"prompt_tokens": 5, "total_tokens": 5},
},
"extra_field": "extra",
})
self.assertEqual(response.extra_field, "extra")


if __name__ == "__main__":
unittest.main()
Loading
Loading