Merge 521009048a into sapling-pr-archive-ehhuang

This commit is contained in:
ehhuang 2025-10-08 13:54:26 -07:00 committed by GitHub
commit ee0152fc07
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 16 additions and 0 deletions

View file

@ -1052,6 +1052,7 @@ class InferenceProvider(Protocol):
prompt_logprobs: int | None = None,
# for fill-in-the-middle type completion
suffix: str | None = None,
**kwargs: Any,
) -> OpenAICompletion:
"""Create completion.
@ -1075,6 +1076,7 @@ class InferenceProvider(Protocol):
:param top_p: (Optional) The top p to use.
:param user: (Optional) The user to use.
:param suffix: (Optional) The suffix that should be appended to the completion.
:param kwargs: (Optional) Additional provider-specific parameters to pass through as extra_body.
:returns: An OpenAICompletion.
"""
...

View file

@ -201,6 +201,7 @@ class InferenceRouter(Inference):
guided_choice: list[str] | None = None,
prompt_logprobs: int | None = None,
suffix: str | None = None,
**kwargs: Any,
) -> OpenAICompletion:
logger.debug(
f"InferenceRouter.openai_completion: {model=}, {stream=}, {prompt=}",
@ -227,6 +228,7 @@ class InferenceRouter(Inference):
guided_choice=guided_choice,
prompt_logprobs=prompt_logprobs,
suffix=suffix,
**kwargs,
)
provider = await self.routing_table.get_provider_impl(model_obj.identifier)
if stream:

View file

@ -96,6 +96,7 @@ class SentenceTransformersInferenceImpl(
prompt_logprobs: int | None = None,
# for fill-in-the-middle type completion
suffix: str | None = None,
**kwargs: Any,
) -> OpenAICompletion:
raise NotImplementedError("OpenAI completion not supported by sentence transformers provider")

View file

@ -158,6 +158,7 @@ class BedrockInferenceAdapter(
prompt_logprobs: int | None = None,
# for fill-in-the-middle type completion
suffix: str | None = None,
**kwargs: Any,
) -> OpenAICompletion:
raise NotImplementedError("OpenAI completion not supported by the Bedrock provider")

View file

@ -63,5 +63,6 @@ class DatabricksInferenceAdapter(OpenAIMixin):
guided_choice: list[str] | None = None,
prompt_logprobs: int | None = None,
suffix: str | None = None,
**kwargs: Any,
) -> OpenAICompletion:
raise NotImplementedError()

View file

@ -54,6 +54,7 @@ class LlamaCompatInferenceAdapter(OpenAIMixin):
guided_choice: list[str] | None = None,
prompt_logprobs: int | None = None,
suffix: str | None = None,
**kwargs: Any,
) -> OpenAICompletion:
raise NotImplementedError()

View file

@ -100,6 +100,7 @@ class PassthroughInferenceAdapter(Inference):
guided_choice: list[str] | None = None,
prompt_logprobs: int | None = None,
suffix: str | None = None,
**kwargs: Any,
) -> OpenAICompletion:
client = self._get_client()
model_obj = await self.model_store.get_model(model)
@ -124,6 +125,7 @@ class PassthroughInferenceAdapter(Inference):
user=user,
guided_choice=guided_choice,
prompt_logprobs=prompt_logprobs,
**kwargs,
)
return await client.inference.openai_completion(**params)

View file

@ -247,6 +247,7 @@ class LiteLLMOpenAIMixin(
guided_choice: list[str] | None = None,
prompt_logprobs: int | None = None,
suffix: str | None = None,
**kwargs: Any,
) -> OpenAICompletion:
model_obj = await self.model_store.get_model(model)
params = await prepare_openai_completion_params(
@ -271,6 +272,7 @@ class LiteLLMOpenAIMixin(
prompt_logprobs=prompt_logprobs,
api_key=self.get_api_key(),
api_base=self.api_base,
**kwargs,
)
return await litellm.atext_completion(**params)

View file

@ -247,6 +247,7 @@ class OpenAIMixin(NeedsRequestProviderData, ABC, BaseModel):
guided_choice: list[str] | None = None,
prompt_logprobs: int | None = None,
suffix: str | None = None,
**kwargs: Any,
) -> OpenAICompletion:
"""
Direct OpenAI completion API call.
@ -261,6 +262,9 @@ class OpenAIMixin(NeedsRequestProviderData, ABC, BaseModel):
if guided_choice:
extra_body["guided_choice"] = guided_choice
# Merge any additional kwargs into extra_body
extra_body.update(kwargs)
# TODO: fix openai_completion to return type compatible with OpenAI's API response
resp = await self.client.completions.create(
**await prepare_openai_completion_params(