Skip to content
Open
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
5 changes: 5 additions & 0 deletions src/memos/configs/llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,11 @@ class OpenAILLMConfig(BaseLLMConfig):
default="https://api.openai.com/v1", description="Base URL for OpenAI API"
)
extra_body: Any = Field(default=None, description="extra body")
enable_thinking: bool | None = Field(
default=None,
description="Enable/disable thinking mode for models that support it (e.g. Qwen3, DeepSeek-R1). "
"When None (default), the provider's default behavior is preserved.",
)
backup_client: bool = Field(
default=False,
description="Whether to enable backup client for fallback on primary failure",
Expand Down
73 changes: 46 additions & 27 deletions src/memos/llms/openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,15 +61,7 @@ def _parse_response(self, response) -> str:
return reasoning_content + (response_content or "")
return response_content or ""

@timed_with_status(
log_prefix="OpenAI LLM",
log_extra_args=lambda self, messages, **kwargs: {
"model_name_or_path": kwargs.get("model_name_or_path", self.config.model_name_or_path),
"messages": messages,
},
)
def generate(self, messages: MessageList, **kwargs) -> str:
"""Generate a response from OpenAI LLM, optionally overriding generation params."""
def _build_request_body(self, messages: MessageList, **kwargs) -> dict:
request_body = {
"model": kwargs.get("model_name_or_path", self.config.model_name_or_path),
"messages": messages,
Expand All @@ -79,6 +71,21 @@ def generate(self, messages: MessageList, **kwargs) -> str:
"extra_body": kwargs.get("extra_body", self.config.extra_body),
"tools": kwargs.get("tools", NOT_GIVEN),
}
enable_thinking = kwargs.get("enable_thinking", self.config.enable_thinking)
if enable_thinking is not None:
request_body["enable_thinking"] = enable_thinking
return request_body

@timed_with_status(
log_prefix="OpenAI LLM",
log_extra_args=lambda self, messages, **kwargs: {
"model_name_or_path": kwargs.get("model_name_or_path", self.config.model_name_or_path),
"messages": messages,
},
)
def generate(self, messages: MessageList, **kwargs) -> str:
"""Generate a response from OpenAI LLM, optionally overriding generation params."""
request_body = self._build_request_body(messages, **kwargs)
start_time = time.perf_counter()
logger.info(f"OpenAI LLM Request body: {request_body}")

Expand Down Expand Up @@ -132,6 +139,10 @@ def generate_stream(self, messages: MessageList, **kwargs) -> Generator[str, Non
"tools": kwargs.get("tools", NOT_GIVEN),
}

enable_thinking = kwargs.get("enable_thinking", self.config.enable_thinking)
if enable_thinking is not None:
request_body["enable_thinking"] = enable_thinking

logger.info(f"OpenAI LLM Stream Request body: {request_body}")
response = self.client.chat.completions.create(**request_body)

Expand Down Expand Up @@ -184,15 +195,19 @@ def __init__(self, config: AzureLLMConfig):

def generate(self, messages: MessageList, **kwargs) -> str:
"""Generate a response from Azure OpenAI LLM."""
response = self.client.chat.completions.create(
model=self.config.model_name_or_path,
messages=messages,
temperature=kwargs.get("temperature", self.config.temperature),
max_tokens=kwargs.get("max_tokens", self.config.max_tokens),
top_p=kwargs.get("top_p", self.config.top_p),
tools=kwargs.get("tools", NOT_GIVEN),
extra_body=kwargs.get("extra_body", self.config.extra_body),
)
request_body = {
"model": self.config.model_name_or_path,
"messages": messages,
"temperature": kwargs.get("temperature", self.config.temperature),
"max_tokens": kwargs.get("max_tokens", self.config.max_tokens),
"top_p": kwargs.get("top_p", self.config.top_p),
"tools": kwargs.get("tools", NOT_GIVEN),
"extra_body": kwargs.get("extra_body", self.config.extra_body),
}
enable_thinking = kwargs.get("enable_thinking", getattr(self.config, "enable_thinking", None))
if enable_thinking is not None:
request_body["enable_thinking"] = enable_thinking
response = self.client.chat.completions.create(**request_body)
logger.info(f"Response from Azure OpenAI: {response.model_dump_json()}")
if not response.choices:
logger.warning("Azure OpenAI response has no choices")
Expand All @@ -212,15 +227,19 @@ def generate_stream(self, messages: MessageList, **kwargs) -> Generator[str, Non
logger.info("stream api not support tools")
return

response = self.client.chat.completions.create(
model=self.config.model_name_or_path,
messages=messages,
stream=True,
temperature=kwargs.get("temperature", self.config.temperature),
max_tokens=kwargs.get("max_tokens", self.config.max_tokens),
top_p=kwargs.get("top_p", self.config.top_p),
extra_body=kwargs.get("extra_body", self.config.extra_body),
)
request_body = {
"model": self.config.model_name_or_path,
"messages": messages,
"stream": True,
"temperature": kwargs.get("temperature", self.config.temperature),
"max_tokens": kwargs.get("max_tokens", self.config.max_tokens),
"top_p": kwargs.get("top_p", self.config.top_p),
"extra_body": kwargs.get("extra_body", self.config.extra_body),
}
enable_thinking = kwargs.get("enable_thinking", getattr(self.config, "enable_thinking", None))
if enable_thinking is not None:
request_body["enable_thinking"] = enable_thinking
response = self.client.chat.completions.create(**request_body)

reasoning_started = False

Expand Down
75 changes: 75 additions & 0 deletions tests/llms/test_enable_thinking.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,75 @@
"""Tests for enable_thinking parameter in OpenAILLM."""

from memos.configs.llm import OpenAILLMConfig
from memos.llms.openai import OpenAILLM


def _make_config(**overrides):
defaults = {
"provider": "openai",
"api_key": "test-key",
"model_name_or_path": "gpt-4o",
}
defaults.update(overrides)
return OpenAILLMConfig(**defaults)


class TestEnableThinkingConfig:
def test_enable_thinking_defaults_to_none(self):
config = _make_config()
assert config.enable_thinking is None

def test_enable_thinking_can_be_set(self):
config = _make_config(enable_thinking=True)
assert config.enable_thinking is True

def test_enable_thinking_false(self):
config = _make_config(enable_thinking=False)
assert config.enable_thinking is False


class TestEnableThinkingInRequest:
def test_build_request_body_without_enable_thinking(self):
config = _make_config()
embedder = OpenAILLM.__new__(OpenAILLM)
embedder.config = config

body = embedder._build_request_body([{"role": "user", "content": "hi"}])
assert "enable_thinking" not in body

def test_build_request_body_with_enable_thinking_from_config(self):
config = _make_config(enable_thinking=True)
embedder = OpenAILLM.__new__(OpenAILLM)
embedder.config = config

body = embedder._build_request_body([{"role": "user", "content": "hi"}])
assert body["enable_thinking"] is True

def test_build_request_body_with_enable_thinking_from_kwargs(self):
config = _make_config(enable_thinking=True)
embedder = OpenAILLM.__new__(OpenAILLM)
embedder.config = config

body = embedder._build_request_body(
[{"role": "user", "content": "hi"}], enable_thinking=False
)
assert body["enable_thinking"] is False

def test_build_request_body_with_enable_thinking_false(self):
config = _make_config(enable_thinking=False)
embedder = OpenAILLM.__new__(OpenAILLM)
embedder.config = config

body = embedder._build_request_body([{"role": "user", "content": "hi"}])
assert body["enable_thinking"] is False

def test_build_request_body_preserves_other_params(self):
config = _make_config(enable_thinking=True, temperature=0.5, max_tokens=100)
embedder = OpenAILLM.__new__(OpenAILLM)
embedder.config = config

body = embedder._build_request_body([{"role": "user", "content": "hi"}])
assert body["enable_thinking"] is True
assert body["temperature"] == 0.5
assert body["max_tokens"] == 100
assert body["model"] == "gpt-4o"
Loading