fix(oci): preserve tool results, embed URL path, and generic finish reason

- Use SerializeAsAny on CohereChatRequest.chatHistory so subclass-specific
  fields like CohereToolMessage.toolResults are not dropped during Pydantic
  v2 serialization.
- Make OCIEmbedConfig.get_complete_url append the /20231130/actions/embedText
  action path consistently with chat, so setting litellm.api_base to the
  region inference base URL no longer posts to the bare hostname.
- Map OCI finishReason (COMPLETE / MAX_TOKENS / TOOL_CALLS) to OpenAI
  finish_reason values in handle_generic_response, mirroring the streaming
  handler and the Cohere non-streaming handler.

Co-authored-by: Yassin Kortam <yassin@berri.ai>
This commit is contained in:
Cursor Agent 2026-05-19 07:23:46 +00:00
parent 167a629267
commit d2701fc7b8
No known key found for this signature in database
6 changed files with 26 additions and 16 deletions

View file

@ -309,8 +309,9 @@ def handle_generic_response(
model_response.created = int(dt.timestamp())
model_response.model = completion_response.modelId
response_choice = completion_response.chatResponse.choices[0]
message = model_response.choices[0].message # type: ignore
response_message = completion_response.chatResponse.choices[0].message
response_message = response_choice.message
if response_message is not None:
if (
response_message.content
@ -323,6 +324,16 @@ def handle_generic_response(
response_message.toolCalls
)
oci_finish_reason = response_choice.finishReason
if oci_finish_reason == "COMPLETE":
model_response.choices[0].finish_reason = "stop" # type: ignore[union-attr]
elif oci_finish_reason == "MAX_TOKENS":
model_response.choices[0].finish_reason = "length" # type: ignore[union-attr]
elif oci_finish_reason == "TOOL_CALLS":
model_response.choices[0].finish_reason = "tool_calls" # type: ignore[union-attr]
elif oci_finish_reason is not None:
model_response.choices[0].finish_reason = oci_finish_reason # type: ignore[union-attr]
oci_usage = completion_response.chatResponse.usage
reasoning_tokens: Optional[int] = None
if (

View file

@ -124,12 +124,7 @@ class OCIEmbedConfig(BaseEmbeddingConfig):
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
# If the caller provides a full endpoint URL, use it as-is.
# Otherwise construct the standard OCI GenAI embedText endpoint from the region.
resolved_base = api_base or litellm.api_base
if resolved_base:
return resolved_base.rstrip("/")
base = get_oci_base_url(optional_params, None)
base = get_oci_base_url(optional_params, api_base or litellm.api_base)
return f"{base}/{OCI_API_VERSION}/actions/embedText"
def sign_request(

View file

@ -3,7 +3,7 @@ from __future__ import annotations
from enum import Enum
from typing import Any, Dict, List, Literal, Optional, Union
from pydantic import BaseModel
from pydantic import BaseModel, SerializeAsAny
OCIRoles = Literal["SYSTEM", "USER", "ASSISTANT", "TOOL"]
@ -318,7 +318,11 @@ class CohereChatRequest(BaseModel):
apiFormat: Literal["COHERE"] = "COHERE"
# Optional fields
chatHistory: Optional[List[CohereMessage]] = None
# ``SerializeAsAny`` preserves subclass-specific fields (e.g. ``toolResults``
# on ``CohereToolMessage``) when this request is serialized via ``model_dump``.
# Without it, Pydantic v2 would serialize each element using the declared
# ``CohereMessage`` schema and silently drop subclass fields.
chatHistory: Optional[List[SerializeAsAny[CohereMessage]]] = None
maxTokens: Optional[int] = None
temperature: Optional[float] = None
topP: Optional[float] = None

View file

@ -443,7 +443,7 @@ class TestOCIChatConfig:
{"type": "TEXT", "text": "I am doing well, thank you!"}
],
},
"finishReason": "STOP",
"finishReason": "COMPLETE",
}
],
"timeCreated": created_time,

View file

@ -73,7 +73,7 @@ class TestOCIEmbedConfig:
)
def test_get_complete_url_respects_api_base(self):
"""api_base is returned as-is (caller supplies complete URL for dedicated/custom endpoints)."""
"""api_base is treated as a base URL — the action path is appended."""
cfg = self._config()
url = cfg.get_complete_url(
api_base="https://custom.endpoint.example.com",
@ -82,10 +82,10 @@ class TestOCIEmbedConfig:
optional_params={},
litellm_params={},
)
assert url == "https://custom.endpoint.example.com"
assert url == "https://custom.endpoint.example.com/20231130/actions/embedText"
def test_get_complete_url_strips_trailing_slash(self):
"""Trailing slash is stripped from api_base."""
"""Trailing slash is stripped from api_base before appending the action path."""
cfg = self._config()
url = cfg.get_complete_url(
api_base="https://custom.endpoint.example.com/",
@ -94,7 +94,7 @@ class TestOCIEmbedConfig:
optional_params={},
litellm_params={},
)
assert url == "https://custom.endpoint.example.com"
assert url == "https://custom.endpoint.example.com/20231130/actions/embedText"
# ------------------------------------------------------------------
# transform_embedding_request

View file

@ -76,7 +76,7 @@ class TestOCIEmbeddingConfig:
assert "embedText" in url
def test_get_complete_url_custom_api_base(self):
"""test_get_complete_url returns api_base as-is when provided."""
"""test_get_complete_url treats api_base as a base URL and appends the embedText path."""
config = OCIEmbeddingConfig()
custom_base = "https://custom.oci.example.com/embed"
url = config.get_complete_url(
@ -86,7 +86,7 @@ class TestOCIEmbeddingConfig:
optional_params={},
litellm_params={},
)
assert url == custom_base
assert url == f"{custom_base}/20231130/actions/embedText"
def test_get_supported_openai_params(self):
"""test_get_supported_openai_params returns expected params list."""