feat: enhance compatibility with more LLM models (#905)

* add try-except to avoid retry when using completion_cost

* feat: add JSON prompt injection and think tag removal handling

* refactor: simplify cost handling in LiteLLM backend and clean up style

* docs: clarify purpose of reasoning_think_rm in LLMSettings

* refactor: remove unused *args and redundant cost assignment

---------

Co-authored-by: Young <afe.young@gmail.com>
This commit is contained in:
Haoran Pan
2025-05-29 15:21:34 +08:00
committed by GitHub
parent 4ba28aa15f
commit b0e88c7375
3 changed files with 57 additions and 39 deletions
+40 -34
View File
@@ -330,6 +330,7 @@ class APIBackend(ABC):
*args,
**kwargs,
) -> str | list[list[float]]:
"""This function to share operation between embedding and chat completion"""
assert not (chat_completion and embedding), "chat_completion and embedding cannot be True at the same time"
max_retry = LLM_SETTINGS.max_retry if LLM_SETTINGS.max_retry is not None else max_retry
timeout_count = 0
@@ -397,38 +398,31 @@ class APIBackend(ABC):
error_message = f"Failed to create chat completion after {max_retry} retries."
raise RuntimeError(error_message)
def _create_chat_completion_add_json_in_prompt(
self,
messages: list[dict[str, Any]],
add_json_in_prompt: bool = False,
json_mode: bool = False,
*args: Any,
**kwargs: Any,
) -> tuple[str, str | None]:
def _add_json_in_prompt(self, messages: list[dict[str, Any]]) -> None:
"""
add json related content in the prompt if add_json_in_prompt is True
"""
if json_mode and add_json_in_prompt:
for message in messages[::-1]:
message["content"] = message["content"] + "\nPlease respond in json format."
if message["role"] == LLM_SETTINGS.system_prompt_role:
# NOTE: assumption: systemprompt is always the first message
break
return self._create_chat_completion_inner_function(messages=messages, json_mode=json_mode, *args, **kwargs) # type: ignore[misc]
for message in messages[::-1]:
message["content"] = message["content"] + "\nPlease respond in json format."
if message["role"] == LLM_SETTINGS.system_prompt_role:
# NOTE: assumption: systemprompt is always the first message
break
def _create_chat_completion_auto_continue(
self,
messages: list[dict[str, Any]],
*args: Any,
json_mode: bool = False,
chat_cache_prefix: str = "",
seed: Optional[int] = None,
json_target_type: Optional[str] = None,
add_json_in_prompt: bool = False,
**kwargs: Any,
) -> str:
"""
Call the chat completion function and automatically continue the conversation if the finish_reason is length.
"""
# 0) return directly if cache is hit
if seed is None and LLM_SETTINGS.use_auto_chat_cache_seed_gen:
seed = LLM_CACHE_SEED_GEN.get_next_seed()
input_content_json = json.dumps(messages)
@@ -443,31 +437,43 @@ class APIBackend(ABC):
logger.info(f"{LogColors.CYAN}Response:{cache_result}{LogColors.END}", tag="llm_messages")
return cache_result
# 1) get a full response
all_response = ""
new_messages = deepcopy(messages)
# Loop to get a full response
try_n = 6
for _ in range(try_n): # for some long code, 3 times may not enough for reasoning models
if "json_mode" in kwargs:
del kwargs["json_mode"]
response, finish_reason = self._create_chat_completion_add_json_in_prompt(
new_messages, json_mode=json_mode, *args, **kwargs
) # type: ignore[misc]
if json_mode and add_json_in_prompt:
self._add_json_in_prompt(new_messages)
response, finish_reason = self._create_chat_completion_inner_function(
messages=new_messages,
json_mode=json_mode,
**kwargs,
)
all_response += response
if finish_reason is None or finish_reason != "length":
if json_mode:
try:
json.loads(all_response)
except:
match = re.search(r"```json(.*)```", all_response, re.DOTALL)
all_response = match.groups()[0] if match else all_response
json.loads(all_response)
if json_target_type is not None:
TypeAdapter(json_target_type).validate_json(all_response)
if self.dump_chat_cache:
self.cache.chat_set(input_content_json, all_response)
return all_response
break # we get a full response now.
new_messages.append({"role": "assistant", "content": response})
raise RuntimeError(f"Failed to continue the conversation after {try_n} retries.")
else:
raise RuntimeError(f"Failed to continue the conversation after {try_n} retries.")
# 2) refine the response and return
if LLM_SETTINGS.reasoning_think_rm:
match = re.search(r"<think>(.*?)</think>(.*)", all_response, re.DOTALL)
_, all_response = match.groups() if match else ("", all_response)
if json_mode:
try:
json.loads(all_response)
except json.decoder.JSONDecodeError:
match = re.search(r"```json(.*)```", all_response, re.DOTALL)
all_response = match.groups()[0] if match else all_response
json.loads(all_response)
if json_target_type is not None:
TypeAdapter(json_target_type).validate_json(all_response)
if self.dump_chat_cache:
self.cache.chat_set(input_content_json, all_response)
return all_response
def _create_embedding_with_cache(
self, input_content_list: list[str], *args: Any, **kwargs: Any
+10 -5
View File
@@ -135,11 +135,16 @@ class LiteLLMAPIBackend(APIBackend):
logger.info(f"{LogColors.BLUE}assistant:{LogColors.END} {finish_reason_str}\n{content}", tag="llm_messages")
global ACC_COST
cost = completion_cost(model=model, messages=messages, completion=content)
ACC_COST += cost
logger.info(
f"Current Cost: ${float(cost):.10f}; Accumulated Cost: ${float(ACC_COST):.10f}; {finish_reason=}",
)
try:
cost = completion_cost(model=model, messages=messages, completion=content)
except Exception as e:
logger.warning(f"Cost calculation failed for model {model}: {e}. Skip cost statistics.")
else:
ACC_COST += cost
logger.info(
f"Current Cost: ${float(cost):.10f}; Accumulated Cost: ${float(ACC_COST):.10f}; {finish_reason=}",
)
prompt_tokens = token_counter(model=model, messages=messages)
completion_tokens = token_counter(model=model, text=content)
logger.log_object(
+7
View File
@@ -17,6 +17,13 @@ class LLMSettings(ExtendedBaseSettings):
reasoning_effort: Literal["low", "medium", "high"] | None = None
# Handling format
reasoning_think_rm: bool = False
"""
Some LLMs include <think>...</think> tags in their responses, which can interfere with the main output.
Set reasoning_think_rm to True to remove any <think>...</think> content from responses.
"""
# TODO: most of the settings are only used on deprec.DeprecBackend.
# So they should move the settings to that folder.