use crate::core::language_models::{BaseChatModel, BaseLanguageModel, LLMResult};
use crate::core::runnables::Runnable;
use crate::core::tools::ToolDefinition;
use crate::language_models::openai::{
OpenAIChat, OpenAIConfig, OpenAIError, StructuredOutputMethod,
};
use crate::schema::Message;
use crate::RunnableConfig;
use async_trait::async_trait;
use futures_util::Stream;
use schemars::JsonSchema;
use serde::de::DeserializeOwned;
use std::env;
use std::pin::Pin;
pub const QWEN_BASE_URL: &str = "https://dashscope.aliyuncs.com/compatible-mode/v1";
pub const QWEN_MODELS: [&str; 6] = [
"qwen-turbo", "qwen-plus", "qwen-max", "qwen-max-longcontext", "qwen2.5-72b-instruct", "qwen-coder-plus", ];
#[derive(Debug, Clone)]
pub struct QwenConfig {
pub api_key: String,
pub base_url: String,
pub model: String,
pub temperature: Option<f32>,
pub max_tokens: Option<usize>,
}
impl Default for QwenConfig {
fn default() -> Self {
Self {
api_key: String::new(),
base_url: QWEN_BASE_URL.to_string(),
model: "qwen-plus".to_string(),
temperature: None,
max_tokens: None,
}
}
}
impl QwenConfig {
pub fn new(api_key: impl Into<String>) -> Self {
Self {
api_key: api_key.into(),
..Default::default()
}
}
pub fn from_env() -> Result<Self, String> {
let api_key = env::var("QWEN_API_KEY")
.map_err(|_| "QWEN_API_KEY environment variable not set".to_string())?;
let base_url = env::var("QWEN_BASE_URL").unwrap_or_else(|_| QWEN_BASE_URL.to_string());
let model = env::var("QWEN_MODEL").unwrap_or_else(|_| "qwen-plus".to_string());
Ok(Self {
api_key,
base_url,
model,
..Default::default()
})
}
pub fn with_model(mut self, model: impl Into<String>) -> Self {
self.model = model.into();
self
}
pub fn with_temperature(mut self, temp: f32) -> Self {
self.temperature = Some(temp);
self
}
pub fn with_max_tokens(mut self, max: usize) -> Self {
self.max_tokens = Some(max);
self
}
pub fn into_openai_config(self) -> OpenAIConfig {
OpenAIConfig {
api_key: self.api_key,
base_url: self.base_url,
model: self.model,
temperature: self.temperature,
max_tokens: self.max_tokens,
top_p: None,
frequency_penalty: None,
presence_penalty: None,
streaming: false,
organization: None,
tools: None,
tool_choice: None,
}
}
}
pub struct QwenChat {
inner: OpenAIChat,
}
impl QwenChat {
pub fn new(config: QwenConfig) -> Self {
Self {
inner: OpenAIChat::new(config.into_openai_config()),
}
}
pub fn from_env() -> Result<Self, String> {
Ok(Self::new(QwenConfig::from_env()?))
}
pub fn with_model(model: impl Into<String>) -> Result<Self, String> {
let config = QwenConfig::from_env()?.with_model(model);
Ok(Self::new(config))
}
}
impl QwenChat {
pub async fn chat(
&self,
messages: Vec<Message>,
config: Option<RunnableConfig>,
) -> Result<LLMResult, OpenAIError> {
self.inner.chat(messages, config).await
}
pub async fn stream_chat(
&self,
messages: Vec<Message>,
config: Option<RunnableConfig>,
) -> Result<Pin<Box<dyn Stream<Item = Result<String, OpenAIError>> + Send>>, OpenAIError> {
self.inner.stream_chat(messages, config).await
}
pub fn bind_tools(&self, tools: Vec<ToolDefinition>) -> Self {
Self {
inner: self.inner.bind_tools(tools),
}
}
pub fn with_structured_output<T: DeserializeOwned + JsonSchema>(
&self,
) -> StructuredOutputMethod<T> {
self.inner.with_structured_output()
}
}
#[async_trait]
impl BaseLanguageModel<Vec<Message>, LLMResult> for QwenChat {
fn model_name(&self) -> &str {
self.inner.model_name()
}
fn get_num_tokens(&self, text: &str) -> usize {
self.inner.get_num_tokens(text)
}
fn temperature(&self) -> Option<f32> {
self.inner.temperature()
}
fn max_tokens(&self) -> Option<usize> {
self.inner.max_tokens()
}
fn with_temperature(self, temp: f32) -> Self {
Self {
inner: self.inner.with_temperature(temp),
}
}
fn with_max_tokens(self, max: usize) -> Self {
Self {
inner: self.inner.with_max_tokens(max),
}
}
}
#[async_trait]
impl Runnable<Vec<Message>, LLMResult> for QwenChat {
type Error = OpenAIError;
async fn invoke(
&self,
input: Vec<Message>,
config: Option<RunnableConfig>,
) -> Result<LLMResult, Self::Error> {
self.chat(input, config).await
}
}
#[async_trait]
impl BaseChatModel for QwenChat {
async fn chat(
&self,
messages: Vec<Message>,
config: Option<RunnableConfig>,
) -> Result<LLMResult, Self::Error> {
self.inner.chat(messages, config).await
}
async fn stream_chat(
&self,
messages: Vec<Message>,
config: Option<RunnableConfig>,
) -> Result<Pin<Box<dyn Stream<Item = Result<String, Self::Error>> + Send>>, Self::Error> {
self.inner.stream_chat(messages, config).await
}
}