logiguard fork v3: full patch set on verified 8c74db0 tree
Includes prior-session patches (carry forward so the app compiles): - crates/gpui/build.rs: cross-compile manifest fix - crates/gpui/src/platform.rs: PlatformWindow::activate_with_token trait method - crates/gpui/src/window.rs: Window::activate_with_token public API - crates/gpui_linux/src/linux/wayland/window.rs: WaylandWindow::activate_with_token + activate() keyboard-serial fix Plus the focus-serial fix: - serial.rs: SerialKind::KeyboardEnter - client.rs: store wl_keyboard.enter serial; latest_serial_of() Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
38
crates/language_model/Cargo.toml
Normal file
38
crates/language_model/Cargo.toml
Normal file
@@ -0,0 +1,38 @@
|
||||
[package]
|
||||
name = "language_model"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
publish.workspace = true
|
||||
license = "GPL-3.0-or-later"
|
||||
|
||||
[lints]
|
||||
workspace = true
|
||||
|
||||
[lib]
|
||||
path = "src/language_model.rs"
|
||||
doctest = false
|
||||
|
||||
[features]
|
||||
test-support = []
|
||||
|
||||
[dependencies]
|
||||
anyhow.workspace = true
|
||||
credentials_provider.workspace = true
|
||||
base64.workspace = true
|
||||
collections.workspace = true
|
||||
env_var.workspace = true
|
||||
futures.workspace = true
|
||||
gpui.workspace = true
|
||||
http_client.workspace = true
|
||||
icons.workspace = true
|
||||
image.workspace = true
|
||||
language_model_core.workspace = true
|
||||
log.workspace = true
|
||||
parking_lot.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
thiserror.workspace = true
|
||||
util.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
gpui = { workspace = true, features = ["test-support"] }
|
||||
1
crates/language_model/LICENSE-GPL
Symbolic link
1
crates/language_model/LICENSE-GPL
Symbolic link
@@ -0,0 +1 @@
|
||||
../../LICENSE-GPL
|
||||
298
crates/language_model/src/api_key.rs
Normal file
298
crates/language_model/src/api_key.rs
Normal file
@@ -0,0 +1,298 @@
|
||||
use anyhow::{Result, anyhow};
|
||||
use credentials_provider::CredentialsProvider;
|
||||
use env_var::EnvVar;
|
||||
use futures::{FutureExt, future};
|
||||
use gpui::{AsyncApp, Context, SharedString, Task};
|
||||
use std::{
|
||||
fmt::{Display, Formatter},
|
||||
sync::Arc,
|
||||
};
|
||||
use util::ResultExt as _;
|
||||
|
||||
use crate::AuthenticateError;
|
||||
|
||||
/// Manages a single API key for a language model provider. API keys either come from environment
|
||||
/// variables or the system keychain.
|
||||
///
|
||||
/// Keys from the system keychain are associated with a provider URL, and this ensures that they are
|
||||
/// only used with that URL.
|
||||
pub struct ApiKeyState {
|
||||
pub url: SharedString,
|
||||
env_var: EnvVar,
|
||||
load_status: LoadStatus,
|
||||
load_task: Option<future::Shared<Task<()>>>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum LoadStatus {
|
||||
NotPresent,
|
||||
Error(String),
|
||||
Loaded(ApiKey),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ApiKey {
|
||||
source: ApiKeySource,
|
||||
key: Arc<str>,
|
||||
}
|
||||
|
||||
impl ApiKeyState {
|
||||
pub fn new(url: SharedString, env_var: EnvVar) -> Self {
|
||||
Self {
|
||||
url,
|
||||
env_var,
|
||||
load_status: LoadStatus::NotPresent,
|
||||
load_task: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn has_key(&self) -> bool {
|
||||
matches!(self.load_status, LoadStatus::Loaded { .. })
|
||||
}
|
||||
|
||||
pub fn env_var_name(&self) -> &SharedString {
|
||||
&self.env_var.name
|
||||
}
|
||||
|
||||
pub fn is_from_env_var(&self) -> bool {
|
||||
match &self.load_status {
|
||||
LoadStatus::Loaded(ApiKey {
|
||||
source: ApiKeySource::EnvVar { .. },
|
||||
..
|
||||
}) => true,
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
/// Get the stored API key, verifying that it is associated with the URL. Returns `None` if
|
||||
/// there is no key or for URL mismatches, and the mismatch case is logged.
|
||||
///
|
||||
/// To avoid URL mismatches, expects that `load_if_needed` or `handle_url_change` has been
|
||||
/// called with this URL.
|
||||
pub fn key(&self, url: &str) -> Option<Arc<str>> {
|
||||
let api_key = match &self.load_status {
|
||||
LoadStatus::Loaded(api_key) => api_key,
|
||||
_ => return None,
|
||||
};
|
||||
if url == self.url.as_str() {
|
||||
Some(api_key.key.clone())
|
||||
} else if let ApiKeySource::EnvVar(var_name) = &api_key.source {
|
||||
log::warn!(
|
||||
"{} is now being used with URL {}, when initially it was used with URL {}",
|
||||
var_name,
|
||||
url,
|
||||
self.url
|
||||
);
|
||||
Some(api_key.key.clone())
|
||||
} else {
|
||||
// bug case because load_if_needed should be called whenever the url may have changed
|
||||
log::error!(
|
||||
"bug: Attempted to use API key associated with URL {} instead with URL {}",
|
||||
self.url,
|
||||
url
|
||||
);
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
/// Set or delete the API key in the system keychain.
|
||||
pub fn store<Ent: 'static>(
|
||||
&mut self,
|
||||
url: SharedString,
|
||||
key: Option<String>,
|
||||
get_this: impl Fn(&mut Ent) -> &mut Self + 'static,
|
||||
provider: Arc<dyn CredentialsProvider>,
|
||||
cx: &Context<Ent>,
|
||||
) -> Task<Result<()>> {
|
||||
if self.is_from_env_var() {
|
||||
return Task::ready(Err(anyhow!(
|
||||
"bug: attempted to store API key in system keychain when API key is from env var",
|
||||
)));
|
||||
}
|
||||
cx.spawn(async move |ent, cx| {
|
||||
if let Some(key) = &key {
|
||||
provider
|
||||
.write_credentials(&url, "Bearer", key.as_bytes(), cx)
|
||||
.await
|
||||
.log_err();
|
||||
} else {
|
||||
provider.delete_credentials(&url, cx).await.log_err();
|
||||
}
|
||||
ent.update(cx, |ent, cx| {
|
||||
let this = get_this(ent);
|
||||
this.url = url;
|
||||
this.load_status = match &key {
|
||||
Some(key) => LoadStatus::Loaded(ApiKey {
|
||||
source: ApiKeySource::SystemKeychain,
|
||||
key: key.as_str().into(),
|
||||
}),
|
||||
None => LoadStatus::NotPresent,
|
||||
};
|
||||
cx.notify();
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
/// Reloads the API key if the current API key is associated with a different URL.
|
||||
///
|
||||
/// Note that it is not efficient to use this or `load_if_needed` with multiple URLs
|
||||
/// interchangeably - URL change should correspond to some user initiated change.
|
||||
pub fn handle_url_change<Ent: 'static>(
|
||||
&mut self,
|
||||
url: SharedString,
|
||||
get_this: impl Fn(&mut Ent) -> &mut Self + Clone + 'static,
|
||||
provider: Arc<dyn CredentialsProvider>,
|
||||
cx: &mut Context<Ent>,
|
||||
) {
|
||||
if url != self.url {
|
||||
if !self.is_from_env_var() {
|
||||
// loading will continue even though this result task is dropped
|
||||
let _task = self.load_if_needed(url, get_this, provider, cx);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// If needed, loads the API key associated with the given URL from the system keychain. When a
|
||||
/// non-empty environment variable is provided, it will be used instead. If called when an API
|
||||
/// key was already loaded for a different URL, that key will be cleared before loading.
|
||||
///
|
||||
/// Dropping the returned Task does not cancel key loading.
|
||||
pub fn load_if_needed<Ent: 'static>(
|
||||
&mut self,
|
||||
url: SharedString,
|
||||
get_this: impl Fn(&mut Ent) -> &mut Self + Clone + 'static,
|
||||
provider: Arc<dyn CredentialsProvider>,
|
||||
cx: &mut Context<Ent>,
|
||||
) -> Task<Result<(), AuthenticateError>> {
|
||||
if let LoadStatus::Loaded { .. } = &self.load_status
|
||||
&& self.url == url
|
||||
{
|
||||
return Task::ready(Ok(()));
|
||||
}
|
||||
|
||||
if let Some(key) = &self.env_var.value
|
||||
&& !key.is_empty()
|
||||
{
|
||||
let api_key = ApiKey::from_env(self.env_var.name.clone(), key);
|
||||
self.url = url;
|
||||
self.load_status = LoadStatus::Loaded(api_key);
|
||||
self.load_task = None;
|
||||
cx.notify();
|
||||
return Task::ready(Ok(()));
|
||||
}
|
||||
|
||||
let task = if let Some(load_task) = &self.load_task {
|
||||
load_task.clone()
|
||||
} else {
|
||||
let load_task = Self::load(url.clone(), get_this.clone(), provider, cx).shared();
|
||||
self.url = url;
|
||||
self.load_status = LoadStatus::NotPresent;
|
||||
self.load_task = Some(load_task.clone());
|
||||
cx.notify();
|
||||
load_task
|
||||
};
|
||||
|
||||
cx.spawn(async move |ent, cx| {
|
||||
task.await;
|
||||
ent.update(cx, |ent, _cx| {
|
||||
get_this(ent).load_status.clone().into_authenticate_result()
|
||||
})
|
||||
.ok();
|
||||
Ok(())
|
||||
})
|
||||
}
|
||||
|
||||
fn load<Ent: 'static>(
|
||||
url: SharedString,
|
||||
get_this: impl Fn(&mut Ent) -> &mut Self + 'static,
|
||||
provider: Arc<dyn CredentialsProvider>,
|
||||
cx: &Context<Ent>,
|
||||
) -> Task<()> {
|
||||
cx.spawn({
|
||||
async move |ent, cx| {
|
||||
let load_status =
|
||||
ApiKey::load_from_system_keychain_impl(&url, provider.as_ref(), cx).await;
|
||||
ent.update(cx, |ent, cx| {
|
||||
let this = get_this(ent);
|
||||
this.url = url;
|
||||
this.load_status = load_status;
|
||||
this.load_task = None;
|
||||
cx.notify();
|
||||
})
|
||||
.ok();
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl ApiKey {
|
||||
pub fn key(&self) -> &str {
|
||||
&self.key
|
||||
}
|
||||
|
||||
pub fn from_env(env_var_name: SharedString, key: &str) -> Self {
|
||||
Self {
|
||||
source: ApiKeySource::EnvVar(env_var_name),
|
||||
key: key.into(),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn load_from_system_keychain(
|
||||
url: &str,
|
||||
credentials_provider: &dyn CredentialsProvider,
|
||||
cx: &AsyncApp,
|
||||
) -> Result<Self, AuthenticateError> {
|
||||
Self::load_from_system_keychain_impl(url, credentials_provider, cx)
|
||||
.await
|
||||
.into_authenticate_result()
|
||||
}
|
||||
|
||||
async fn load_from_system_keychain_impl(
|
||||
url: &str,
|
||||
credentials_provider: &dyn CredentialsProvider,
|
||||
cx: &AsyncApp,
|
||||
) -> LoadStatus {
|
||||
if url.is_empty() {
|
||||
return LoadStatus::NotPresent;
|
||||
}
|
||||
let read_result = credentials_provider.read_credentials(&url, cx).await;
|
||||
let api_key = match read_result {
|
||||
Ok(Some((_, api_key))) => api_key,
|
||||
Ok(None) => return LoadStatus::NotPresent,
|
||||
Err(err) => return LoadStatus::Error(err.to_string()),
|
||||
};
|
||||
let key = match str::from_utf8(&api_key) {
|
||||
Ok(key) => key,
|
||||
Err(_) => return LoadStatus::Error(format!("API key for URL {url} is not utf8")),
|
||||
};
|
||||
LoadStatus::Loaded(Self {
|
||||
source: ApiKeySource::SystemKeychain,
|
||||
key: key.into(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl LoadStatus {
|
||||
fn into_authenticate_result(self) -> Result<ApiKey, AuthenticateError> {
|
||||
match self {
|
||||
LoadStatus::Loaded(api_key) => Ok(api_key),
|
||||
LoadStatus::NotPresent => Err(AuthenticateError::CredentialsNotFound),
|
||||
LoadStatus::Error(err) => Err(AuthenticateError::Other(anyhow!(err))),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
enum ApiKeySource {
|
||||
EnvVar(SharedString),
|
||||
SystemKeychain,
|
||||
}
|
||||
|
||||
impl Display for ApiKeySource {
|
||||
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
ApiKeySource::EnvVar(var) => write!(f, "environment variable {}", var),
|
||||
ApiKeySource::SystemKeychain => write!(f, "system keychain"),
|
||||
}
|
||||
}
|
||||
}
|
||||
336
crates/language_model/src/fake_provider.rs
Normal file
336
crates/language_model/src/fake_provider.rs
Normal file
@@ -0,0 +1,336 @@
|
||||
use crate::{
|
||||
AuthenticateError, ConfigurationViewTargetAgent, LanguageModel, LanguageModelCompletionError,
|
||||
LanguageModelCompletionEvent, LanguageModelId, LanguageModelName, LanguageModelProvider,
|
||||
LanguageModelProviderId, LanguageModelProviderName, LanguageModelProviderState,
|
||||
LanguageModelRequest, LanguageModelToolChoice,
|
||||
};
|
||||
use anyhow::anyhow;
|
||||
use futures::{FutureExt, channel::mpsc, future::BoxFuture, stream::BoxStream, stream::StreamExt};
|
||||
use gpui::{AnyView, App, AsyncApp, Entity, Task, Window};
|
||||
use http_client::Result;
|
||||
use parking_lot::Mutex;
|
||||
use std::sync::{
|
||||
Arc,
|
||||
atomic::{AtomicBool, Ordering::SeqCst},
|
||||
};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct FakeLanguageModelProvider {
|
||||
id: LanguageModelProviderId,
|
||||
name: LanguageModelProviderName,
|
||||
models: Vec<Arc<dyn LanguageModel>>,
|
||||
}
|
||||
|
||||
impl Default for FakeLanguageModelProvider {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
id: LanguageModelProviderId::from("fake".to_string()),
|
||||
name: LanguageModelProviderName::from("Fake".to_string()),
|
||||
models: vec![Arc::new(FakeLanguageModel::default())],
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl LanguageModelProviderState for FakeLanguageModelProvider {
|
||||
type ObservableEntity = ();
|
||||
|
||||
fn observable_entity(&self) -> Option<Entity<Self::ObservableEntity>> {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
impl LanguageModelProvider for FakeLanguageModelProvider {
|
||||
fn id(&self) -> LanguageModelProviderId {
|
||||
self.id.clone()
|
||||
}
|
||||
|
||||
fn name(&self) -> LanguageModelProviderName {
|
||||
self.name.clone()
|
||||
}
|
||||
|
||||
fn default_model(&self, _cx: &App) -> Option<Arc<dyn LanguageModel>> {
|
||||
self.models.first().cloned()
|
||||
}
|
||||
|
||||
fn default_fast_model(&self, _cx: &App) -> Option<Arc<dyn LanguageModel>> {
|
||||
self.models.first().cloned()
|
||||
}
|
||||
|
||||
fn provided_models(&self, _: &App) -> Vec<Arc<dyn LanguageModel>> {
|
||||
self.models.clone()
|
||||
}
|
||||
|
||||
fn is_authenticated(&self, _: &App) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn authenticate(&self, _: &mut App) -> Task<Result<(), AuthenticateError>> {
|
||||
Task::ready(Ok(()))
|
||||
}
|
||||
|
||||
fn configuration_view(
|
||||
&self,
|
||||
_target_agent: ConfigurationViewTargetAgent,
|
||||
_window: &mut Window,
|
||||
_: &mut App,
|
||||
) -> AnyView {
|
||||
unimplemented!()
|
||||
}
|
||||
|
||||
fn reset_credentials(&self, _: &mut App) -> Task<Result<()>> {
|
||||
Task::ready(Ok(()))
|
||||
}
|
||||
}
|
||||
|
||||
impl FakeLanguageModelProvider {
|
||||
pub fn new(id: LanguageModelProviderId, name: LanguageModelProviderName) -> Self {
|
||||
Self {
|
||||
id,
|
||||
name,
|
||||
models: vec![Arc::new(FakeLanguageModel::default())],
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_models(mut self, models: Vec<Arc<dyn LanguageModel>>) -> Self {
|
||||
self.models = models;
|
||||
self
|
||||
}
|
||||
|
||||
pub fn test_model(&self) -> FakeLanguageModel {
|
||||
FakeLanguageModel::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, PartialEq)]
|
||||
pub struct ToolUseRequest {
|
||||
pub request: LanguageModelRequest,
|
||||
pub name: String,
|
||||
pub description: String,
|
||||
pub schema: serde_json::Value,
|
||||
}
|
||||
|
||||
pub struct FakeLanguageModel {
|
||||
id: LanguageModelId,
|
||||
name: LanguageModelName,
|
||||
provider_id: LanguageModelProviderId,
|
||||
provider_name: LanguageModelProviderName,
|
||||
current_completion_txs: Mutex<
|
||||
Vec<(
|
||||
LanguageModelRequest,
|
||||
mpsc::UnboundedSender<
|
||||
Result<LanguageModelCompletionEvent, LanguageModelCompletionError>,
|
||||
>,
|
||||
)>,
|
||||
>,
|
||||
forbid_requests: AtomicBool,
|
||||
supports_thinking: AtomicBool,
|
||||
supports_streaming_tools: AtomicBool,
|
||||
supports_images: AtomicBool,
|
||||
}
|
||||
|
||||
impl Default for FakeLanguageModel {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
id: LanguageModelId::from("fake".to_string()),
|
||||
name: LanguageModelName::from("Fake".to_string()),
|
||||
provider_id: LanguageModelProviderId::from("fake".to_string()),
|
||||
provider_name: LanguageModelProviderName::from("Fake".to_string()),
|
||||
current_completion_txs: Mutex::new(Vec::new()),
|
||||
forbid_requests: AtomicBool::new(false),
|
||||
supports_thinking: AtomicBool::new(false),
|
||||
supports_streaming_tools: AtomicBool::new(false),
|
||||
supports_images: AtomicBool::new(false),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl FakeLanguageModel {
|
||||
pub fn with_id_and_thinking(
|
||||
provider_id: &str,
|
||||
id: &str,
|
||||
name: &str,
|
||||
supports_thinking: bool,
|
||||
) -> Self {
|
||||
Self {
|
||||
id: LanguageModelId::from(id.to_string()),
|
||||
name: LanguageModelName::from(name.to_string()),
|
||||
provider_id: LanguageModelProviderId::from(provider_id.to_string()),
|
||||
supports_thinking: AtomicBool::new(supports_thinking),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
pub fn allow_requests(&self) {
|
||||
self.forbid_requests.store(false, SeqCst);
|
||||
}
|
||||
|
||||
pub fn forbid_requests(&self) {
|
||||
self.forbid_requests.store(true, SeqCst);
|
||||
}
|
||||
|
||||
pub fn set_supports_thinking(&self, supports: bool) {
|
||||
self.supports_thinking.store(supports, SeqCst);
|
||||
}
|
||||
|
||||
pub fn set_supports_streaming_tools(&self, supports: bool) {
|
||||
self.supports_streaming_tools.store(supports, SeqCst);
|
||||
}
|
||||
|
||||
pub fn set_supports_images(&self, supports: bool) {
|
||||
self.supports_images.store(supports, SeqCst);
|
||||
}
|
||||
|
||||
pub fn pending_completions(&self) -> Vec<LanguageModelRequest> {
|
||||
self.current_completion_txs
|
||||
.lock()
|
||||
.iter()
|
||||
.map(|(request, _)| request.clone())
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn completion_count(&self) -> usize {
|
||||
self.current_completion_txs.lock().len()
|
||||
}
|
||||
|
||||
pub fn send_completion_stream_text_chunk(
|
||||
&self,
|
||||
request: &LanguageModelRequest,
|
||||
chunk: impl Into<String>,
|
||||
) {
|
||||
self.send_completion_stream_event(
|
||||
request,
|
||||
LanguageModelCompletionEvent::Text(chunk.into()),
|
||||
);
|
||||
}
|
||||
|
||||
pub fn send_completion_stream_event(
|
||||
&self,
|
||||
request: &LanguageModelRequest,
|
||||
event: impl Into<LanguageModelCompletionEvent>,
|
||||
) {
|
||||
let current_completion_txs = self.current_completion_txs.lock();
|
||||
let tx = current_completion_txs
|
||||
.iter()
|
||||
.find(|(req, _)| req == request)
|
||||
.map(|(_, tx)| tx)
|
||||
.unwrap();
|
||||
tx.unbounded_send(Ok(event.into())).unwrap();
|
||||
}
|
||||
|
||||
pub fn send_completion_stream_error(
|
||||
&self,
|
||||
request: &LanguageModelRequest,
|
||||
error: impl Into<LanguageModelCompletionError>,
|
||||
) {
|
||||
let current_completion_txs = self.current_completion_txs.lock();
|
||||
let tx = current_completion_txs
|
||||
.iter()
|
||||
.find(|(req, _)| req == request)
|
||||
.map(|(_, tx)| tx)
|
||||
.unwrap();
|
||||
tx.unbounded_send(Err(error.into())).unwrap();
|
||||
}
|
||||
|
||||
pub fn end_completion_stream(&self, request: &LanguageModelRequest) {
|
||||
self.current_completion_txs
|
||||
.lock()
|
||||
.retain(|(req, _)| req != request);
|
||||
}
|
||||
|
||||
pub fn send_last_completion_stream_text_chunk(&self, chunk: impl Into<String>) {
|
||||
self.send_completion_stream_text_chunk(self.pending_completions().last().unwrap(), chunk);
|
||||
}
|
||||
|
||||
pub fn send_last_completion_stream_event(
|
||||
&self,
|
||||
event: impl Into<LanguageModelCompletionEvent>,
|
||||
) {
|
||||
self.send_completion_stream_event(self.pending_completions().last().unwrap(), event);
|
||||
}
|
||||
|
||||
pub fn send_last_completion_stream_error(
|
||||
&self,
|
||||
error: impl Into<LanguageModelCompletionError>,
|
||||
) {
|
||||
self.send_completion_stream_error(self.pending_completions().last().unwrap(), error);
|
||||
}
|
||||
|
||||
pub fn end_last_completion_stream(&self) {
|
||||
self.end_completion_stream(self.pending_completions().last().unwrap());
|
||||
}
|
||||
}
|
||||
|
||||
impl LanguageModel for FakeLanguageModel {
|
||||
fn id(&self) -> LanguageModelId {
|
||||
self.id.clone()
|
||||
}
|
||||
|
||||
fn name(&self) -> LanguageModelName {
|
||||
self.name.clone()
|
||||
}
|
||||
|
||||
fn provider_id(&self) -> LanguageModelProviderId {
|
||||
self.provider_id.clone()
|
||||
}
|
||||
|
||||
fn provider_name(&self) -> LanguageModelProviderName {
|
||||
self.provider_name.clone()
|
||||
}
|
||||
|
||||
fn supports_tools(&self) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
fn supports_tool_choice(&self, _choice: LanguageModelToolChoice) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
fn supports_images(&self) -> bool {
|
||||
self.supports_images.load(SeqCst)
|
||||
}
|
||||
|
||||
fn supports_thinking(&self) -> bool {
|
||||
self.supports_thinking.load(SeqCst)
|
||||
}
|
||||
|
||||
fn supports_streaming_tools(&self) -> bool {
|
||||
self.supports_streaming_tools.load(SeqCst)
|
||||
}
|
||||
|
||||
fn telemetry_id(&self) -> String {
|
||||
"fake".to_string()
|
||||
}
|
||||
|
||||
fn max_token_count(&self) -> u64 {
|
||||
1000000
|
||||
}
|
||||
|
||||
fn stream_completion(
|
||||
&self,
|
||||
request: LanguageModelRequest,
|
||||
_: &AsyncApp,
|
||||
) -> BoxFuture<
|
||||
'static,
|
||||
Result<
|
||||
BoxStream<'static, Result<LanguageModelCompletionEvent, LanguageModelCompletionError>>,
|
||||
LanguageModelCompletionError,
|
||||
>,
|
||||
> {
|
||||
if self.forbid_requests.load(SeqCst) {
|
||||
async move {
|
||||
Err(LanguageModelCompletionError::Other(anyhow!(
|
||||
"requests are forbidden"
|
||||
)))
|
||||
}
|
||||
.boxed()
|
||||
} else {
|
||||
let (tx, rx) = mpsc::unbounded();
|
||||
self.current_completion_txs.lock().push((request, tx));
|
||||
async move { Ok(rx.boxed()) }.boxed()
|
||||
}
|
||||
}
|
||||
|
||||
fn as_fake(&self) -> &Self {
|
||||
self
|
||||
}
|
||||
}
|
||||
358
crates/language_model/src/language_model.rs
Normal file
358
crates/language_model/src/language_model.rs
Normal file
@@ -0,0 +1,358 @@
|
||||
mod api_key;
|
||||
mod model;
|
||||
mod registry;
|
||||
mod request;
|
||||
|
||||
#[cfg(any(test, feature = "test-support"))]
|
||||
pub mod fake_provider;
|
||||
|
||||
pub use language_model_core::*;
|
||||
|
||||
use anyhow::Result;
|
||||
use futures::FutureExt;
|
||||
use futures::{StreamExt, future::BoxFuture, stream::BoxStream};
|
||||
use gpui::{AnyView, App, AsyncApp, Task, Window};
|
||||
use icons::IconName;
|
||||
use parking_lot::Mutex;
|
||||
use std::sync::Arc;
|
||||
|
||||
pub use crate::api_key::{ApiKey, ApiKeyState};
|
||||
pub use crate::model::*;
|
||||
pub use crate::registry::*;
|
||||
pub use crate::request::{LanguageModelImageExt, gpui_size_to_image_size, image_size_to_gpui};
|
||||
pub use env_var::{EnvVar, env_var};
|
||||
|
||||
pub fn init(cx: &mut App) {
|
||||
registry::init(cx);
|
||||
}
|
||||
|
||||
pub struct LanguageModelTextStream {
|
||||
pub message_id: Option<String>,
|
||||
pub stream: BoxStream<'static, Result<String, LanguageModelCompletionError>>,
|
||||
// Has complete token usage after the stream has finished
|
||||
pub last_token_usage: Arc<Mutex<TokenUsage>>,
|
||||
}
|
||||
|
||||
impl Default for LanguageModelTextStream {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
message_id: None,
|
||||
stream: Box::pin(futures::stream::empty()),
|
||||
last_token_usage: Arc::new(Mutex::new(TokenUsage::default())),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub trait LanguageModel: Send + Sync {
|
||||
fn id(&self) -> LanguageModelId;
|
||||
fn name(&self) -> LanguageModelName;
|
||||
fn provider_id(&self) -> LanguageModelProviderId;
|
||||
fn provider_name(&self) -> LanguageModelProviderName;
|
||||
fn upstream_provider_id(&self) -> LanguageModelProviderId {
|
||||
self.provider_id()
|
||||
}
|
||||
fn upstream_provider_name(&self) -> LanguageModelProviderName {
|
||||
self.provider_name()
|
||||
}
|
||||
|
||||
/// Returns whether this model is the "latest", so we can highlight it in the UI.
|
||||
fn is_latest(&self) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
fn telemetry_id(&self) -> String;
|
||||
|
||||
fn api_key(&self, _cx: &App) -> Option<String> {
|
||||
None
|
||||
}
|
||||
|
||||
/// Information about the cost of using this model, if available.
|
||||
fn model_cost_info(&self) -> Option<LanguageModelCostInfo> {
|
||||
None
|
||||
}
|
||||
|
||||
/// Whether this model supports thinking.
|
||||
fn supports_thinking(&self) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
fn supports_fast_mode(&self) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
/// Returns the list of supported effort levels that can be used when thinking.
|
||||
fn supported_effort_levels(&self) -> Vec<LanguageModelEffortLevel> {
|
||||
Vec::new()
|
||||
}
|
||||
|
||||
/// Returns the default effort level to use when thinking.
|
||||
fn default_effort_level(&self) -> Option<LanguageModelEffortLevel> {
|
||||
self.supported_effort_levels()
|
||||
.into_iter()
|
||||
.find(|effort_level| effort_level.is_default)
|
||||
}
|
||||
|
||||
/// Whether this model supports images
|
||||
fn supports_images(&self) -> bool;
|
||||
|
||||
/// Whether this model supports tools.
|
||||
fn supports_tools(&self) -> bool;
|
||||
|
||||
/// Whether this model supports choosing which tool to use.
|
||||
fn supports_tool_choice(&self, choice: LanguageModelToolChoice) -> bool;
|
||||
|
||||
/// Returns whether this model or provider supports streaming tool calls;
|
||||
fn supports_streaming_tools(&self) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
/// Returns whether this model/provider reports accurate split input/output token counts.
|
||||
/// When true, the UI may show separate input/output token indicators.
|
||||
fn supports_split_token_display(&self) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
fn tool_input_format(&self) -> LanguageModelToolSchemaFormat {
|
||||
LanguageModelToolSchemaFormat::JsonSchema
|
||||
}
|
||||
|
||||
fn max_token_count(&self) -> u64;
|
||||
fn max_output_tokens(&self) -> Option<u64> {
|
||||
None
|
||||
}
|
||||
|
||||
fn stream_completion(
|
||||
&self,
|
||||
request: LanguageModelRequest,
|
||||
cx: &AsyncApp,
|
||||
) -> BoxFuture<
|
||||
'static,
|
||||
Result<
|
||||
BoxStream<'static, Result<LanguageModelCompletionEvent, LanguageModelCompletionError>>,
|
||||
LanguageModelCompletionError,
|
||||
>,
|
||||
>;
|
||||
|
||||
fn stream_completion_text(
|
||||
&self,
|
||||
request: LanguageModelRequest,
|
||||
cx: &AsyncApp,
|
||||
) -> BoxFuture<'static, Result<LanguageModelTextStream, LanguageModelCompletionError>> {
|
||||
let future = self.stream_completion(request, cx);
|
||||
|
||||
async move {
|
||||
let events = future.await?;
|
||||
let mut events = events.fuse();
|
||||
let mut message_id = None;
|
||||
let mut first_item_text = None;
|
||||
let last_token_usage = Arc::new(Mutex::new(TokenUsage::default()));
|
||||
|
||||
if let Some(first_event) = events.next().await {
|
||||
match first_event {
|
||||
Ok(LanguageModelCompletionEvent::StartMessage { message_id: id }) => {
|
||||
message_id = Some(id);
|
||||
}
|
||||
Ok(LanguageModelCompletionEvent::Text(text)) => {
|
||||
first_item_text = Some(text);
|
||||
}
|
||||
_ => (),
|
||||
}
|
||||
}
|
||||
|
||||
let stream = futures::stream::iter(first_item_text.map(Ok))
|
||||
.chain(events.filter_map({
|
||||
let last_token_usage = last_token_usage.clone();
|
||||
move |result| {
|
||||
let last_token_usage = last_token_usage.clone();
|
||||
async move {
|
||||
match result {
|
||||
Ok(LanguageModelCompletionEvent::Queued { .. }) => None,
|
||||
Ok(LanguageModelCompletionEvent::Started) => None,
|
||||
Ok(LanguageModelCompletionEvent::StartMessage { .. }) => None,
|
||||
Ok(LanguageModelCompletionEvent::Text(text)) => Some(Ok(text)),
|
||||
Ok(LanguageModelCompletionEvent::Thinking { .. }) => None,
|
||||
Ok(LanguageModelCompletionEvent::RedactedThinking { .. }) => None,
|
||||
Ok(LanguageModelCompletionEvent::ReasoningDetails(_)) => None,
|
||||
Ok(LanguageModelCompletionEvent::Stop(_)) => None,
|
||||
Ok(LanguageModelCompletionEvent::ToolUse(_)) => None,
|
||||
Ok(LanguageModelCompletionEvent::ToolUseJsonParseError {
|
||||
..
|
||||
}) => None,
|
||||
Ok(LanguageModelCompletionEvent::UsageUpdate(token_usage)) => {
|
||||
*last_token_usage.lock() = token_usage;
|
||||
None
|
||||
}
|
||||
Err(err) => Some(Err(err)),
|
||||
}
|
||||
}
|
||||
}
|
||||
}))
|
||||
.boxed();
|
||||
|
||||
Ok(LanguageModelTextStream {
|
||||
message_id,
|
||||
stream,
|
||||
last_token_usage,
|
||||
})
|
||||
}
|
||||
.boxed()
|
||||
}
|
||||
|
||||
fn stream_completion_tool(
|
||||
&self,
|
||||
request: LanguageModelRequest,
|
||||
cx: &AsyncApp,
|
||||
) -> BoxFuture<'static, Result<LanguageModelToolUse, LanguageModelCompletionError>> {
|
||||
let future = self.stream_completion(request, cx);
|
||||
|
||||
async move {
|
||||
let events = future.await?;
|
||||
let mut events = events.fuse();
|
||||
|
||||
// Iterate through events until we find a complete ToolUse
|
||||
while let Some(event) = events.next().await {
|
||||
match event {
|
||||
Ok(LanguageModelCompletionEvent::ToolUse(tool_use))
|
||||
if tool_use.is_input_complete =>
|
||||
{
|
||||
return Ok(tool_use);
|
||||
}
|
||||
Err(err) => {
|
||||
return Err(err);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
// Stream ended without a complete tool use
|
||||
Err(LanguageModelCompletionError::Other(anyhow::anyhow!(
|
||||
"Stream ended without receiving a complete tool use"
|
||||
)))
|
||||
}
|
||||
.boxed()
|
||||
}
|
||||
|
||||
fn cache_configuration(&self) -> Option<LanguageModelCacheConfiguration> {
|
||||
None
|
||||
}
|
||||
|
||||
#[cfg(any(test, feature = "test-support"))]
|
||||
fn as_fake(&self) -> &fake_provider::FakeLanguageModel {
|
||||
unimplemented!()
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for dyn LanguageModel {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("<dyn LanguageModel>")
|
||||
.field("id", &self.id())
|
||||
.field("name", &self.name())
|
||||
.field("provider_id", &self.provider_id())
|
||||
.field("provider_name", &self.provider_name())
|
||||
.field("upstream_provider_name", &self.upstream_provider_name())
|
||||
.field("upstream_provider_id", &self.upstream_provider_id())
|
||||
.field("upstream_provider_id", &self.upstream_provider_id())
|
||||
.field("supports_streaming_tools", &self.supports_streaming_tools())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
/// Either a built-in icon name or a path to an external SVG.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum IconOrSvg {
|
||||
/// A built-in icon from Zed's icon set.
|
||||
Icon(IconName),
|
||||
/// Path to a custom SVG icon file.
|
||||
Svg(SharedString),
|
||||
}
|
||||
|
||||
impl Default for IconOrSvg {
|
||||
fn default() -> Self {
|
||||
Self::Icon(IconName::ZedAssistant)
|
||||
}
|
||||
}
|
||||
|
||||
pub trait LanguageModelProvider: 'static {
|
||||
fn id(&self) -> LanguageModelProviderId;
|
||||
fn name(&self) -> LanguageModelProviderName;
|
||||
fn icon(&self) -> IconOrSvg {
|
||||
IconOrSvg::default()
|
||||
}
|
||||
fn default_model(&self, cx: &App) -> Option<Arc<dyn LanguageModel>>;
|
||||
fn default_fast_model(&self, cx: &App) -> Option<Arc<dyn LanguageModel>>;
|
||||
fn provided_models(&self, cx: &App) -> Vec<Arc<dyn LanguageModel>>;
|
||||
fn recommended_models(&self, _cx: &App) -> Vec<Arc<dyn LanguageModel>> {
|
||||
Vec::new()
|
||||
}
|
||||
fn is_authenticated(&self, cx: &App) -> bool;
|
||||
fn authenticate(&self, cx: &mut App) -> Task<Result<(), AuthenticateError>>;
|
||||
fn configuration_view(
|
||||
&self,
|
||||
target_agent: ConfigurationViewTargetAgent,
|
||||
window: &mut Window,
|
||||
cx: &mut App,
|
||||
) -> AnyView;
|
||||
fn reset_credentials(&self, cx: &mut App) -> Task<Result<()>>;
|
||||
}
|
||||
|
||||
#[derive(Default, Clone, PartialEq, Eq)]
|
||||
pub enum ConfigurationViewTargetAgent {
|
||||
#[default]
|
||||
ZedAgent,
|
||||
Other(SharedString),
|
||||
}
|
||||
|
||||
pub trait LanguageModelProviderState: 'static {
|
||||
type ObservableEntity;
|
||||
|
||||
fn observable_entity(&self) -> Option<gpui::Entity<Self::ObservableEntity>>;
|
||||
|
||||
fn subscribe<T: 'static>(
|
||||
&self,
|
||||
cx: &mut gpui::Context<T>,
|
||||
callback: impl Fn(&mut T, &mut gpui::Context<T>) + 'static,
|
||||
) -> Option<gpui::Subscription> {
|
||||
let entity = self.observable_entity()?;
|
||||
Some(cx.observe(&entity, move |this, _, cx| {
|
||||
callback(this, cx);
|
||||
}))
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub enum LanguageModelCostInfo {
|
||||
/// Cost per 1,000 input and output tokens
|
||||
TokenCost {
|
||||
input_token_cost_per_1m: f64,
|
||||
output_token_cost_per_1m: f64,
|
||||
},
|
||||
/// Cost per request
|
||||
RequestCost { cost_per_request: f64 },
|
||||
}
|
||||
|
||||
impl LanguageModelCostInfo {
|
||||
pub fn to_shared_string(&self) -> SharedString {
|
||||
match self {
|
||||
LanguageModelCostInfo::RequestCost { cost_per_request } => {
|
||||
let cost_str = format!("{}×", Self::cost_value_to_string(cost_per_request));
|
||||
SharedString::from(cost_str)
|
||||
}
|
||||
LanguageModelCostInfo::TokenCost {
|
||||
input_token_cost_per_1m,
|
||||
output_token_cost_per_1m,
|
||||
} => {
|
||||
let input_cost = Self::cost_value_to_string(input_token_cost_per_1m);
|
||||
let output_cost = Self::cost_value_to_string(output_token_cost_per_1m);
|
||||
SharedString::from(format!("{}$/{}$", input_cost, output_cost))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn cost_value_to_string(cost: &f64) -> SharedString {
|
||||
if (cost.fract() - 0.0).abs() < std::f64::EPSILON {
|
||||
SharedString::from(format!("{:.0}", cost))
|
||||
} else {
|
||||
SharedString::from(format!("{:.2}", cost))
|
||||
}
|
||||
}
|
||||
}
|
||||
3
crates/language_model/src/model.rs
Normal file
3
crates/language_model/src/model.rs
Normal file
@@ -0,0 +1,3 @@
|
||||
pub mod cloud_model;
|
||||
|
||||
pub use cloud_model::*;
|
||||
15
crates/language_model/src/model/cloud_model.rs
Normal file
15
crates/language_model/src/model/cloud_model.rs
Normal file
@@ -0,0 +1,15 @@
|
||||
use std::fmt;
|
||||
|
||||
use thiserror::Error;
|
||||
|
||||
#[derive(Error, Debug)]
|
||||
pub struct PaymentRequiredError;
|
||||
|
||||
impl fmt::Display for PaymentRequiredError {
|
||||
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
|
||||
write!(
|
||||
f,
|
||||
"Payment required to use this language model. Please upgrade your account."
|
||||
)
|
||||
}
|
||||
}
|
||||
660
crates/language_model/src/registry.rs
Normal file
660
crates/language_model/src/registry.rs
Normal file
@@ -0,0 +1,660 @@
|
||||
use crate::{
|
||||
LanguageModel, LanguageModelId, LanguageModelProvider, LanguageModelProviderId,
|
||||
LanguageModelProviderState, ZED_CLOUD_PROVIDER_ID,
|
||||
};
|
||||
use collections::{BTreeMap, HashSet};
|
||||
use gpui::{App, Context, Entity, EventEmitter, Global, prelude::*};
|
||||
use std::{str::FromStr, sync::Arc};
|
||||
use thiserror::Error;
|
||||
|
||||
/// Function type for checking if a built-in provider should be hidden.
|
||||
/// Returns Some(extension_id) if the provider should be hidden when that extension is installed.
|
||||
pub type BuiltinProviderHidingFn = Box<dyn Fn(&str) -> Option<&'static str> + Send + Sync>;
|
||||
|
||||
pub fn init(cx: &mut App) {
|
||||
let registry = cx.new(|_cx| LanguageModelRegistry::default());
|
||||
cx.set_global(GlobalLanguageModelRegistry(registry));
|
||||
}
|
||||
|
||||
struct GlobalLanguageModelRegistry(Entity<LanguageModelRegistry>);
|
||||
|
||||
impl Global for GlobalLanguageModelRegistry {}
|
||||
|
||||
#[derive(Error)]
|
||||
pub enum ConfigurationError {
|
||||
#[error("Configure at least one LLM provider to start using the panel.")]
|
||||
NoProvider,
|
||||
#[error("LLM provider is not configured or does not support the configured model.")]
|
||||
ModelNotFound,
|
||||
#[error("{} LLM provider is not configured.", .0.name().0)]
|
||||
ProviderNotAuthenticated(Arc<dyn LanguageModelProvider>),
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for ConfigurationError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::NoProvider => write!(f, "NoProvider"),
|
||||
Self::ModelNotFound => write!(f, "ModelNotFound"),
|
||||
Self::ProviderNotAuthenticated(provider) => {
|
||||
write!(f, "ProviderNotAuthenticated({})", provider.id())
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
pub struct LanguageModelRegistry {
|
||||
default_model: Option<ConfiguredModel>,
|
||||
/// This model is automatically configured by a user's environment after
|
||||
/// authenticating all providers. It's only used when `default_model` is not set.
|
||||
available_fallback_model: Option<ConfiguredModel>,
|
||||
inline_assistant_model: Option<ConfiguredModel>,
|
||||
commit_message_model: Option<ConfiguredModel>,
|
||||
thread_summary_model: Option<ConfiguredModel>,
|
||||
providers: BTreeMap<LanguageModelProviderId, Arc<dyn LanguageModelProvider>>,
|
||||
inline_alternatives: Vec<Arc<dyn LanguageModel>>,
|
||||
/// Set of installed extension IDs that provide language models.
|
||||
/// Used to determine which built-in providers should be hidden.
|
||||
installed_llm_extension_ids: HashSet<Arc<str>>,
|
||||
/// Function to check if a built-in provider should be hidden by an extension.
|
||||
builtin_provider_hiding_fn: Option<BuiltinProviderHidingFn>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct SelectedModel {
|
||||
pub provider: LanguageModelProviderId,
|
||||
pub model: LanguageModelId,
|
||||
}
|
||||
|
||||
impl FromStr for SelectedModel {
|
||||
type Err = String;
|
||||
|
||||
/// Parse string identifiers like `provider_id/model_id` into a `SelectedModel`
|
||||
fn from_str(id: &str) -> Result<SelectedModel, Self::Err> {
|
||||
let parts: Vec<&str> = id.split('/').collect();
|
||||
let [provider_id, model_id] = parts.as_slice() else {
|
||||
return Err(format!(
|
||||
"Invalid model identifier format: `{}`. Expected `provider_id/model_id`",
|
||||
id
|
||||
));
|
||||
};
|
||||
|
||||
if provider_id.is_empty() || model_id.is_empty() {
|
||||
return Err(format!("Provider and model ids can't be empty: `{}`", id));
|
||||
}
|
||||
|
||||
Ok(SelectedModel {
|
||||
provider: LanguageModelProviderId(provider_id.to_string().into()),
|
||||
model: LanguageModelId(model_id.to_string().into()),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ConfiguredModel {
|
||||
pub provider: Arc<dyn LanguageModelProvider>,
|
||||
pub model: Arc<dyn LanguageModel>,
|
||||
}
|
||||
|
||||
impl ConfiguredModel {
|
||||
pub fn is_same_as(&self, other: &ConfiguredModel) -> bool {
|
||||
self.model.id() == other.model.id() && self.provider.id() == other.provider.id()
|
||||
}
|
||||
|
||||
pub fn is_provided_by_zed(&self) -> bool {
|
||||
self.provider.id() == ZED_CLOUD_PROVIDER_ID
|
||||
}
|
||||
}
|
||||
|
||||
pub enum Event {
|
||||
DefaultModelChanged,
|
||||
InlineAssistantModelChanged,
|
||||
CommitMessageModelChanged,
|
||||
ThreadSummaryModelChanged,
|
||||
ProviderStateChanged(LanguageModelProviderId),
|
||||
AddedProvider(LanguageModelProviderId),
|
||||
RemovedProvider(LanguageModelProviderId),
|
||||
/// Emitted when provider visibility changes due to extension install/uninstall.
|
||||
ProvidersChanged,
|
||||
}
|
||||
|
||||
impl EventEmitter<Event> for LanguageModelRegistry {}
|
||||
|
||||
impl LanguageModelRegistry {
|
||||
pub fn global(cx: &App) -> Entity<Self> {
|
||||
cx.global::<GlobalLanguageModelRegistry>().0.clone()
|
||||
}
|
||||
|
||||
pub fn read_global(cx: &App) -> &Self {
|
||||
cx.global::<GlobalLanguageModelRegistry>().0.read(cx)
|
||||
}
|
||||
|
||||
#[cfg(any(test, feature = "test-support"))]
|
||||
pub fn test(cx: &mut App) -> Arc<crate::fake_provider::FakeLanguageModelProvider> {
|
||||
let fake_provider = Arc::new(crate::fake_provider::FakeLanguageModelProvider::default());
|
||||
let registry = cx.new(|cx| {
|
||||
let mut registry = Self::default();
|
||||
registry.register_provider(fake_provider.clone(), cx);
|
||||
let model = fake_provider.provided_models(cx)[0].clone();
|
||||
let configured_model = ConfiguredModel {
|
||||
provider: fake_provider.clone(),
|
||||
model,
|
||||
};
|
||||
registry.set_default_model(Some(configured_model), cx);
|
||||
registry
|
||||
});
|
||||
cx.set_global(GlobalLanguageModelRegistry(registry));
|
||||
fake_provider
|
||||
}
|
||||
|
||||
#[cfg(any(test, feature = "test-support"))]
|
||||
pub fn fake_model(&self) -> Arc<dyn LanguageModel> {
|
||||
self.default_model.as_ref().unwrap().model.clone()
|
||||
}
|
||||
|
||||
pub fn register_provider<T: LanguageModelProvider + LanguageModelProviderState>(
|
||||
&mut self,
|
||||
provider: Arc<T>,
|
||||
cx: &mut Context<Self>,
|
||||
) {
|
||||
let id = provider.id();
|
||||
|
||||
let subscription = provider.subscribe(cx, {
|
||||
let id = id.clone();
|
||||
move |_, cx| {
|
||||
cx.emit(Event::ProviderStateChanged(id.clone()));
|
||||
}
|
||||
});
|
||||
if let Some(subscription) = subscription {
|
||||
subscription.detach();
|
||||
}
|
||||
|
||||
self.providers.insert(id.clone(), provider);
|
||||
cx.emit(Event::AddedProvider(id));
|
||||
}
|
||||
|
||||
pub fn unregister_provider(&mut self, id: LanguageModelProviderId, cx: &mut Context<Self>) {
|
||||
if self.providers.remove(&id).is_some() {
|
||||
cx.emit(Event::RemovedProvider(id));
|
||||
}
|
||||
}
|
||||
|
||||
pub fn providers(&self) -> Vec<Arc<dyn LanguageModelProvider>> {
|
||||
let zed_provider_id = LanguageModelProviderId("zed.dev".into());
|
||||
let mut providers = Vec::with_capacity(self.providers.len());
|
||||
if let Some(provider) = self.providers.get(&zed_provider_id) {
|
||||
providers.push(provider.clone());
|
||||
}
|
||||
providers.extend(self.providers.values().filter_map(|p| {
|
||||
if p.id() != zed_provider_id {
|
||||
Some(p.clone())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}));
|
||||
providers
|
||||
}
|
||||
|
||||
/// Returns providers, filtering out hidden built-in providers.
|
||||
pub fn visible_providers(&self) -> Vec<Arc<dyn LanguageModelProvider>> {
|
||||
self.providers()
|
||||
.into_iter()
|
||||
.filter(|p| !self.should_hide_provider(&p.id()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Sets the function used to check if a built-in provider should be hidden.
|
||||
pub fn set_builtin_provider_hiding_fn(&mut self, hiding_fn: BuiltinProviderHidingFn) {
|
||||
self.builtin_provider_hiding_fn = Some(hiding_fn);
|
||||
}
|
||||
|
||||
/// Called when an extension is installed/loaded.
|
||||
/// If the extension provides language models, track it so we can hide the corresponding built-in.
|
||||
pub fn extension_installed(&mut self, extension_id: Arc<str>, cx: &mut Context<Self>) {
|
||||
if self.installed_llm_extension_ids.insert(extension_id) {
|
||||
cx.emit(Event::ProvidersChanged);
|
||||
cx.notify();
|
||||
}
|
||||
}
|
||||
|
||||
/// Called when an extension is uninstalled/unloaded.
|
||||
pub fn extension_uninstalled(&mut self, extension_id: &str, cx: &mut Context<Self>) {
|
||||
if self.installed_llm_extension_ids.remove(extension_id) {
|
||||
cx.emit(Event::ProvidersChanged);
|
||||
cx.notify();
|
||||
}
|
||||
}
|
||||
|
||||
/// Sync the set of installed LLM extension IDs.
|
||||
pub fn sync_installed_llm_extensions(
|
||||
&mut self,
|
||||
extension_ids: HashSet<Arc<str>>,
|
||||
cx: &mut Context<Self>,
|
||||
) {
|
||||
if extension_ids != self.installed_llm_extension_ids {
|
||||
self.installed_llm_extension_ids = extension_ids;
|
||||
cx.emit(Event::ProvidersChanged);
|
||||
cx.notify();
|
||||
}
|
||||
}
|
||||
|
||||
/// Returns true if a provider should be hidden from the UI.
|
||||
/// Built-in providers are hidden when their corresponding extension is installed.
|
||||
pub fn should_hide_provider(&self, provider_id: &LanguageModelProviderId) -> bool {
|
||||
if let Some(ref hiding_fn) = self.builtin_provider_hiding_fn {
|
||||
if let Some(extension_id) = hiding_fn(&provider_id.0) {
|
||||
return self.installed_llm_extension_ids.contains(extension_id);
|
||||
}
|
||||
}
|
||||
false
|
||||
}
|
||||
|
||||
pub fn configuration_error(
|
||||
&self,
|
||||
model: Option<ConfiguredModel>,
|
||||
cx: &App,
|
||||
) -> Option<ConfigurationError> {
|
||||
let Some(model) = model else {
|
||||
if !self.has_authenticated_provider(cx) {
|
||||
return Some(ConfigurationError::NoProvider);
|
||||
}
|
||||
return Some(ConfigurationError::ModelNotFound);
|
||||
};
|
||||
|
||||
if !model.provider.is_authenticated(cx) {
|
||||
return Some(ConfigurationError::ProviderNotAuthenticated(model.provider));
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
/// Returns `true` if at least one provider that is authenticated.
|
||||
pub fn has_authenticated_provider(&self, cx: &App) -> bool {
|
||||
self.providers.values().any(|p| p.is_authenticated(cx))
|
||||
}
|
||||
|
||||
pub fn available_models<'a>(
|
||||
&'a self,
|
||||
cx: &'a App,
|
||||
) -> impl Iterator<Item = Arc<dyn LanguageModel>> + 'a {
|
||||
self.providers
|
||||
.values()
|
||||
.filter(|provider| provider.is_authenticated(cx))
|
||||
.flat_map(|provider| provider.provided_models(cx))
|
||||
}
|
||||
|
||||
pub fn provider(&self, id: &LanguageModelProviderId) -> Option<Arc<dyn LanguageModelProvider>> {
|
||||
self.providers.get(id).cloned()
|
||||
}
|
||||
|
||||
pub fn select_default_model(&mut self, model: Option<&SelectedModel>, cx: &mut Context<Self>) {
|
||||
let configured_model = model.and_then(|model| self.select_model(model, cx));
|
||||
self.set_default_model(configured_model, cx);
|
||||
}
|
||||
|
||||
pub fn select_inline_assistant_model(
|
||||
&mut self,
|
||||
model: Option<&SelectedModel>,
|
||||
cx: &mut Context<Self>,
|
||||
) {
|
||||
let configured_model = model.and_then(|model| self.select_model(model, cx));
|
||||
self.set_inline_assistant_model(configured_model, cx);
|
||||
}
|
||||
|
||||
pub fn select_commit_message_model(
|
||||
&mut self,
|
||||
model: Option<&SelectedModel>,
|
||||
cx: &mut Context<Self>,
|
||||
) {
|
||||
let configured_model = model.and_then(|model| self.select_model(model, cx));
|
||||
self.set_commit_message_model(configured_model, cx);
|
||||
}
|
||||
|
||||
pub fn select_thread_summary_model(
|
||||
&mut self,
|
||||
model: Option<&SelectedModel>,
|
||||
cx: &mut Context<Self>,
|
||||
) {
|
||||
let configured_model = model.and_then(|model| self.select_model(model, cx));
|
||||
self.set_thread_summary_model(configured_model, cx);
|
||||
}
|
||||
|
||||
/// Selects and sets the inline alternatives for language models based on
|
||||
/// provider name and id.
|
||||
pub fn select_inline_alternative_models(
|
||||
&mut self,
|
||||
alternatives: impl IntoIterator<Item = SelectedModel>,
|
||||
cx: &mut Context<Self>,
|
||||
) {
|
||||
self.inline_alternatives = alternatives
|
||||
.into_iter()
|
||||
.flat_map(|alternative| {
|
||||
self.select_model(&alternative, cx)
|
||||
.map(|configured_model| configured_model.model)
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
}
|
||||
|
||||
pub fn select_model(
|
||||
&mut self,
|
||||
selected_model: &SelectedModel,
|
||||
cx: &mut Context<Self>,
|
||||
) -> Option<ConfiguredModel> {
|
||||
let provider = self.provider(&selected_model.provider)?;
|
||||
let model = provider
|
||||
.provided_models(cx)
|
||||
.iter()
|
||||
.find(|model| model.id() == selected_model.model)?
|
||||
.clone();
|
||||
Some(ConfiguredModel { provider, model })
|
||||
}
|
||||
|
||||
pub fn set_default_model(&mut self, model: Option<ConfiguredModel>, cx: &mut Context<Self>) {
|
||||
match (self.default_model(), model.as_ref()) {
|
||||
(Some(old), Some(new)) if old.is_same_as(new) => {}
|
||||
(None, None) => {}
|
||||
_ => cx.emit(Event::DefaultModelChanged),
|
||||
}
|
||||
self.default_model = model;
|
||||
}
|
||||
|
||||
pub fn set_environment_fallback_model(
|
||||
&mut self,
|
||||
model: Option<ConfiguredModel>,
|
||||
cx: &mut Context<Self>,
|
||||
) {
|
||||
if self.default_model.is_none() {
|
||||
match (self.available_fallback_model.as_ref(), model.as_ref()) {
|
||||
(Some(old), Some(new)) if old.is_same_as(new) => {}
|
||||
(None, None) => {}
|
||||
_ => cx.emit(Event::DefaultModelChanged),
|
||||
}
|
||||
}
|
||||
self.available_fallback_model = model;
|
||||
}
|
||||
|
||||
pub fn set_inline_assistant_model(
|
||||
&mut self,
|
||||
model: Option<ConfiguredModel>,
|
||||
cx: &mut Context<Self>,
|
||||
) {
|
||||
match (self.inline_assistant_model.as_ref(), model.as_ref()) {
|
||||
(Some(old), Some(new)) if old.is_same_as(new) => {}
|
||||
(None, None) => {}
|
||||
_ => cx.emit(Event::InlineAssistantModelChanged),
|
||||
}
|
||||
self.inline_assistant_model = model;
|
||||
}
|
||||
|
||||
pub fn set_commit_message_model(
|
||||
&mut self,
|
||||
model: Option<ConfiguredModel>,
|
||||
cx: &mut Context<Self>,
|
||||
) {
|
||||
match (self.commit_message_model.as_ref(), model.as_ref()) {
|
||||
(Some(old), Some(new)) if old.is_same_as(new) => {}
|
||||
(None, None) => {}
|
||||
_ => cx.emit(Event::CommitMessageModelChanged),
|
||||
}
|
||||
self.commit_message_model = model;
|
||||
}
|
||||
|
||||
pub fn set_thread_summary_model(
|
||||
&mut self,
|
||||
model: Option<ConfiguredModel>,
|
||||
cx: &mut Context<Self>,
|
||||
) {
|
||||
match (self.thread_summary_model.as_ref(), model.as_ref()) {
|
||||
(Some(old), Some(new)) if old.is_same_as(new) => {}
|
||||
(None, None) => {}
|
||||
_ => cx.emit(Event::ThreadSummaryModelChanged),
|
||||
}
|
||||
self.thread_summary_model = model;
|
||||
}
|
||||
|
||||
pub fn default_model(&self) -> Option<ConfiguredModel> {
|
||||
#[cfg(debug_assertions)]
|
||||
if std::env::var("ZED_SIMULATE_NO_LLM_PROVIDER").is_ok() {
|
||||
return None;
|
||||
}
|
||||
|
||||
self.default_model
|
||||
.clone()
|
||||
.or_else(|| self.available_fallback_model.clone())
|
||||
}
|
||||
|
||||
pub fn default_fast_model(&self, cx: &App) -> Option<ConfiguredModel> {
|
||||
let configured = self.default_model()?;
|
||||
let fast_model = configured.provider.default_fast_model(cx)?;
|
||||
Some(ConfiguredModel {
|
||||
provider: configured.provider,
|
||||
model: fast_model,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn inline_assistant_model(&self) -> Option<ConfiguredModel> {
|
||||
#[cfg(debug_assertions)]
|
||||
if std::env::var("ZED_SIMULATE_NO_LLM_PROVIDER").is_ok() {
|
||||
return None;
|
||||
}
|
||||
|
||||
self.inline_assistant_model
|
||||
.clone()
|
||||
.or_else(|| self.default_model.clone())
|
||||
}
|
||||
|
||||
pub fn commit_message_model(&self, cx: &App) -> Option<ConfiguredModel> {
|
||||
#[cfg(debug_assertions)]
|
||||
if std::env::var("ZED_SIMULATE_NO_LLM_PROVIDER").is_ok() {
|
||||
return None;
|
||||
}
|
||||
|
||||
self.commit_message_model
|
||||
.clone()
|
||||
.or_else(|| self.default_fast_model(cx))
|
||||
.or_else(|| self.default_model())
|
||||
}
|
||||
|
||||
pub fn thread_summary_model(&self, cx: &App) -> Option<ConfiguredModel> {
|
||||
#[cfg(debug_assertions)]
|
||||
if std::env::var("ZED_SIMULATE_NO_LLM_PROVIDER").is_ok() {
|
||||
return None;
|
||||
}
|
||||
|
||||
self.thread_summary_model
|
||||
.clone()
|
||||
.or_else(|| self.default_fast_model(cx))
|
||||
.or_else(|| self.default_model())
|
||||
}
|
||||
|
||||
/// The models to use for inline assists. Returns the union of the active
|
||||
/// model and all inline alternatives. When there are multiple models, the
|
||||
/// user will be able to cycle through results.
|
||||
pub fn inline_alternative_models(&self) -> &[Arc<dyn LanguageModel>] {
|
||||
&self.inline_alternatives
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::fake_provider::FakeLanguageModelProvider;
|
||||
|
||||
#[gpui::test]
|
||||
fn test_register_providers(cx: &mut App) {
|
||||
let registry = cx.new(|_| LanguageModelRegistry::default());
|
||||
|
||||
let provider = Arc::new(FakeLanguageModelProvider::default());
|
||||
registry.update(cx, |registry, cx| {
|
||||
registry.register_provider(provider.clone(), cx);
|
||||
});
|
||||
|
||||
let providers = registry.read(cx).providers();
|
||||
assert_eq!(providers.len(), 1);
|
||||
assert_eq!(providers[0].id(), provider.id());
|
||||
|
||||
registry.update(cx, |registry, cx| {
|
||||
registry.unregister_provider(provider.id(), cx);
|
||||
});
|
||||
|
||||
let providers = registry.read(cx).providers();
|
||||
assert!(providers.is_empty());
|
||||
}
|
||||
|
||||
#[gpui::test]
|
||||
fn test_provider_hiding_on_extension_install(cx: &mut App) {
|
||||
let registry = cx.new(|_| LanguageModelRegistry::default());
|
||||
|
||||
let provider = Arc::new(FakeLanguageModelProvider::default());
|
||||
let provider_id = provider.id();
|
||||
|
||||
registry.update(cx, |registry, cx| {
|
||||
registry.register_provider(provider.clone(), cx);
|
||||
|
||||
registry.set_builtin_provider_hiding_fn(Box::new(|id| {
|
||||
if id == "fake" {
|
||||
Some("fake-extension")
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}));
|
||||
});
|
||||
|
||||
let visible = registry.read(cx).visible_providers();
|
||||
assert_eq!(visible.len(), 1);
|
||||
assert_eq!(visible[0].id(), provider_id);
|
||||
|
||||
registry.update(cx, |registry, cx| {
|
||||
registry.extension_installed("fake-extension".into(), cx);
|
||||
});
|
||||
|
||||
let visible = registry.read(cx).visible_providers();
|
||||
assert!(visible.is_empty());
|
||||
|
||||
let all = registry.read(cx).providers();
|
||||
assert_eq!(all.len(), 1);
|
||||
}
|
||||
|
||||
#[gpui::test]
|
||||
fn test_provider_unhiding_on_extension_uninstall(cx: &mut App) {
|
||||
let registry = cx.new(|_| LanguageModelRegistry::default());
|
||||
|
||||
let provider = Arc::new(FakeLanguageModelProvider::default());
|
||||
let provider_id = provider.id();
|
||||
|
||||
registry.update(cx, |registry, cx| {
|
||||
registry.register_provider(provider.clone(), cx);
|
||||
|
||||
registry.set_builtin_provider_hiding_fn(Box::new(|id| {
|
||||
if id == "fake" {
|
||||
Some("fake-extension")
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}));
|
||||
|
||||
registry.extension_installed("fake-extension".into(), cx);
|
||||
});
|
||||
|
||||
let visible = registry.read(cx).visible_providers();
|
||||
assert!(visible.is_empty());
|
||||
|
||||
registry.update(cx, |registry, cx| {
|
||||
registry.extension_uninstalled("fake-extension", cx);
|
||||
});
|
||||
|
||||
let visible = registry.read(cx).visible_providers();
|
||||
assert_eq!(visible.len(), 1);
|
||||
assert_eq!(visible[0].id(), provider_id);
|
||||
}
|
||||
|
||||
#[gpui::test]
|
||||
fn test_should_hide_provider(cx: &mut App) {
|
||||
let registry = cx.new(|_| LanguageModelRegistry::default());
|
||||
|
||||
registry.update(cx, |registry, cx| {
|
||||
registry.set_builtin_provider_hiding_fn(Box::new(|id| {
|
||||
if id == "anthropic" {
|
||||
Some("anthropic")
|
||||
} else if id == "openai" {
|
||||
Some("openai")
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}));
|
||||
|
||||
registry.extension_installed("anthropic".into(), cx);
|
||||
});
|
||||
|
||||
let registry_read = registry.read(cx);
|
||||
|
||||
assert!(registry_read.should_hide_provider(&LanguageModelProviderId("anthropic".into())));
|
||||
|
||||
assert!(!registry_read.should_hide_provider(&LanguageModelProviderId("openai".into())));
|
||||
|
||||
assert!(!registry_read.should_hide_provider(&LanguageModelProviderId("unknown".into())));
|
||||
}
|
||||
|
||||
#[gpui::test]
|
||||
async fn test_configure_environment_fallback_model(cx: &mut gpui::TestAppContext) {
|
||||
let registry = cx.new(|_| LanguageModelRegistry::default());
|
||||
|
||||
let provider = Arc::new(FakeLanguageModelProvider::default());
|
||||
registry.update(cx, |registry, cx| {
|
||||
registry.register_provider(provider.clone(), cx);
|
||||
});
|
||||
|
||||
cx.update(|cx| provider.authenticate(cx)).await.unwrap();
|
||||
|
||||
registry.update(cx, |registry, cx| {
|
||||
let provider = registry.provider(&provider.id()).unwrap();
|
||||
let model = provider.default_model(cx).unwrap();
|
||||
|
||||
registry.set_environment_fallback_model(
|
||||
Some(ConfiguredModel {
|
||||
provider: provider.clone(),
|
||||
model: model.clone(),
|
||||
}),
|
||||
cx,
|
||||
);
|
||||
|
||||
let default_model = registry.default_model().unwrap();
|
||||
assert_eq!(default_model.model.id(), model.id());
|
||||
assert_eq!(default_model.provider.id(), provider.id());
|
||||
});
|
||||
}
|
||||
|
||||
#[gpui::test]
|
||||
fn test_sync_installed_llm_extensions(cx: &mut App) {
|
||||
let registry = cx.new(|_| LanguageModelRegistry::default());
|
||||
|
||||
let provider = Arc::new(FakeLanguageModelProvider::default());
|
||||
|
||||
registry.update(cx, |registry, cx| {
|
||||
registry.register_provider(provider.clone(), cx);
|
||||
|
||||
registry.set_builtin_provider_hiding_fn(Box::new(|id| {
|
||||
if id == "fake" {
|
||||
Some("fake-extension")
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}));
|
||||
});
|
||||
|
||||
let mut extension_ids = HashSet::default();
|
||||
extension_ids.insert(Arc::from("fake-extension"));
|
||||
|
||||
registry.update(cx, |registry, cx| {
|
||||
registry.sync_installed_llm_extensions(extension_ids, cx);
|
||||
});
|
||||
|
||||
assert!(registry.read(cx).visible_providers().is_empty());
|
||||
|
||||
registry.update(cx, |registry, cx| {
|
||||
registry.sync_installed_llm_extensions(HashSet::default(), cx);
|
||||
});
|
||||
|
||||
assert_eq!(registry.read(cx).visible_providers().len(), 1);
|
||||
}
|
||||
}
|
||||
243
crates/language_model/src/request.rs
Normal file
243
crates/language_model/src/request.rs
Normal file
@@ -0,0 +1,243 @@
|
||||
use std::io::{Cursor, Write};
|
||||
use std::sync::Arc;
|
||||
|
||||
use anyhow::Result;
|
||||
use base64::write::EncoderWriter;
|
||||
use gpui::{
|
||||
App, AppContext as _, DevicePixels, Image, ImageFormat, ObjectFit, Size, Task, point, px, size,
|
||||
};
|
||||
use image::GenericImageView as _;
|
||||
use image::codecs::png::PngEncoder;
|
||||
use util::ResultExt;
|
||||
|
||||
use language_model_core::{ImageSize, LanguageModelImage};
|
||||
|
||||
/// Anthropic wants uploaded images to be smaller than this in both dimensions.
|
||||
const ANTHROPIC_SIZE_LIMIT: f32 = 1568.;
|
||||
|
||||
/// Default per-image hard limit (in bytes) for the encoded image payload we send upstream.
|
||||
///
|
||||
/// NOTE: `LanguageModelImage.source` is base64-encoded PNG bytes (without the `data:` prefix).
|
||||
/// This limit is enforced on the encoded PNG bytes *before* base64 encoding.
|
||||
const DEFAULT_IMAGE_MAX_BYTES: usize = 5 * 1024 * 1024;
|
||||
|
||||
/// Conservative cap on how many times we'll attempt to shrink/re-encode an image to fit
|
||||
/// `DEFAULT_IMAGE_MAX_BYTES`.
|
||||
const MAX_IMAGE_DOWNSCALE_PASSES: usize = 8;
|
||||
|
||||
/// Extension trait for `LanguageModelImage` that provides GPUI-dependent functionality.
|
||||
pub trait LanguageModelImageExt {
|
||||
const FORMAT: ImageFormat;
|
||||
fn from_image(data: Arc<Image>, cx: &mut App) -> Task<Option<LanguageModelImage>>;
|
||||
}
|
||||
|
||||
impl LanguageModelImageExt for LanguageModelImage {
|
||||
const FORMAT: ImageFormat = ImageFormat::Png;
|
||||
|
||||
fn from_image(data: Arc<Image>, cx: &mut App) -> Task<Option<LanguageModelImage>> {
|
||||
cx.background_spawn(async move {
|
||||
let image_bytes = Cursor::new(data.bytes());
|
||||
let dynamic_image = match data.format() {
|
||||
ImageFormat::Png => image::codecs::png::PngDecoder::new(image_bytes)
|
||||
.and_then(image::DynamicImage::from_decoder),
|
||||
ImageFormat::Jpeg => image::codecs::jpeg::JpegDecoder::new(image_bytes)
|
||||
.and_then(image::DynamicImage::from_decoder),
|
||||
ImageFormat::Webp => image::codecs::webp::WebPDecoder::new(image_bytes)
|
||||
.and_then(image::DynamicImage::from_decoder),
|
||||
ImageFormat::Gif => image::codecs::gif::GifDecoder::new(image_bytes)
|
||||
.and_then(image::DynamicImage::from_decoder),
|
||||
ImageFormat::Bmp => image::codecs::bmp::BmpDecoder::new(image_bytes)
|
||||
.and_then(image::DynamicImage::from_decoder),
|
||||
ImageFormat::Tiff => image::codecs::tiff::TiffDecoder::new(image_bytes)
|
||||
.and_then(image::DynamicImage::from_decoder),
|
||||
_ => return None,
|
||||
}
|
||||
.log_err()?;
|
||||
|
||||
let width = dynamic_image.width();
|
||||
let height = dynamic_image.height();
|
||||
let image_size = size(DevicePixels(width as i32), DevicePixels(height as i32));
|
||||
|
||||
// First apply any provider-specific dimension constraints we know about (Anthropic).
|
||||
let mut processed_image = if image_size.width.0 > ANTHROPIC_SIZE_LIMIT as i32
|
||||
|| image_size.height.0 > ANTHROPIC_SIZE_LIMIT as i32
|
||||
{
|
||||
let new_bounds = ObjectFit::ScaleDown.get_bounds(
|
||||
gpui::Bounds {
|
||||
origin: point(px(0.0), px(0.0)),
|
||||
size: size(px(ANTHROPIC_SIZE_LIMIT), px(ANTHROPIC_SIZE_LIMIT)),
|
||||
},
|
||||
image_size,
|
||||
);
|
||||
dynamic_image.resize(
|
||||
new_bounds.size.width.into(),
|
||||
new_bounds.size.height.into(),
|
||||
image::imageops::FilterType::Triangle,
|
||||
)
|
||||
} else {
|
||||
dynamic_image
|
||||
};
|
||||
|
||||
// Then enforce a default per-image size cap on the encoded PNG bytes.
|
||||
//
|
||||
// We always send PNG bytes (either original PNG bytes, or re-encoded PNG) base64'd.
|
||||
// The upstream provider limit we want to respect is effectively on the binary image
|
||||
// payload size, so we enforce against the encoded PNG bytes before base64 encoding.
|
||||
let mut encoded_png = encode_png_bytes(&processed_image).log_err()?;
|
||||
for _pass in 0..MAX_IMAGE_DOWNSCALE_PASSES {
|
||||
if encoded_png.len() <= DEFAULT_IMAGE_MAX_BYTES {
|
||||
break;
|
||||
}
|
||||
|
||||
// Scale down geometrically to converge quickly. We don't know the final PNG size
|
||||
// as a function of pixels, so we iteratively shrink.
|
||||
let (w, h) = processed_image.dimensions();
|
||||
if w <= 1 || h <= 1 {
|
||||
break;
|
||||
}
|
||||
|
||||
// Shrink by ~15% each pass (0.85). This is a compromise between speed and
|
||||
// preserving image detail.
|
||||
let new_w = ((w as f32) * 0.85).round().max(1.0) as u32;
|
||||
let new_h = ((h as f32) * 0.85).round().max(1.0) as u32;
|
||||
|
||||
processed_image =
|
||||
processed_image.resize(new_w, new_h, image::imageops::FilterType::Triangle);
|
||||
encoded_png = encode_png_bytes(&processed_image).log_err()?;
|
||||
}
|
||||
|
||||
if encoded_png.len() > DEFAULT_IMAGE_MAX_BYTES {
|
||||
// Still too large after multiple passes; treat as non-convertible for now.
|
||||
// (Provider-specific handling can be introduced later.)
|
||||
return None;
|
||||
}
|
||||
|
||||
// Now base64 encode the PNG bytes.
|
||||
let base64_image = encode_bytes_as_base64(encoded_png.as_slice()).log_err()?;
|
||||
|
||||
// SAFETY: The base64 encoder should not produce non-UTF8.
|
||||
let source = unsafe { String::from_utf8_unchecked(base64_image) };
|
||||
|
||||
let (final_width, final_height) = processed_image.dimensions();
|
||||
|
||||
Some(LanguageModelImage {
|
||||
size: Some(ImageSize {
|
||||
width: final_width as i32,
|
||||
height: final_height as i32,
|
||||
}),
|
||||
source: source.into(),
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn encode_png_bytes(image: &image::DynamicImage) -> Result<Vec<u8>> {
|
||||
let mut png = Vec::new();
|
||||
image.write_with_encoder(PngEncoder::new(&mut png))?;
|
||||
Ok(png)
|
||||
}
|
||||
|
||||
fn encode_bytes_as_base64(bytes: &[u8]) -> Result<Vec<u8>> {
|
||||
let mut base64_image = Vec::new();
|
||||
{
|
||||
let mut base64_encoder = EncoderWriter::new(
|
||||
Cursor::new(&mut base64_image),
|
||||
&base64::engine::general_purpose::STANDARD,
|
||||
);
|
||||
base64_encoder.write_all(bytes)?;
|
||||
}
|
||||
Ok(base64_image)
|
||||
}
|
||||
|
||||
/// Convert a core `ImageSize` to a gpui `Size<DevicePixels>`.
|
||||
pub fn image_size_to_gpui(size: ImageSize) -> Size<DevicePixels> {
|
||||
Size {
|
||||
width: DevicePixels(size.width),
|
||||
height: DevicePixels(size.height),
|
||||
}
|
||||
}
|
||||
|
||||
/// Convert a gpui `Size<DevicePixels>` to a core `ImageSize`.
|
||||
pub fn gpui_size_to_image_size(size: Size<DevicePixels>) -> ImageSize {
|
||||
ImageSize {
|
||||
width: size.width.0,
|
||||
height: size.height.0,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use base64::Engine as _;
|
||||
use gpui::TestAppContext;
|
||||
|
||||
fn base64_to_png_bytes(base64: &str) -> Vec<u8> {
|
||||
base64::engine::general_purpose::STANDARD
|
||||
.decode(base64)
|
||||
.expect("valid base64")
|
||||
}
|
||||
|
||||
fn png_dimensions(png_bytes: &[u8]) -> (u32, u32) {
|
||||
let img = image::load_from_memory(png_bytes).expect("valid png");
|
||||
(img.width(), img.height())
|
||||
}
|
||||
|
||||
fn make_noisy_png_bytes(width: u32, height: u32) -> Vec<u8> {
|
||||
use image::{ImageBuffer, Rgba};
|
||||
use std::hash::{Hash, Hasher};
|
||||
|
||||
let img = ImageBuffer::from_fn(width, height, |x, y| {
|
||||
let mut hasher = std::hash::DefaultHasher::new();
|
||||
(x, y, width, height).hash(&mut hasher);
|
||||
let h = hasher.finish();
|
||||
Rgba([h as u8, (h >> 8) as u8, (h >> 16) as u8, 255])
|
||||
});
|
||||
|
||||
let mut buf = Cursor::new(Vec::new());
|
||||
img.write_with_encoder(PngEncoder::new(&mut buf))
|
||||
.expect("encode");
|
||||
buf.into_inner()
|
||||
}
|
||||
|
||||
#[gpui::test]
|
||||
async fn test_from_image_downscales_to_default_5mb_limit(cx: &mut TestAppContext) {
|
||||
let raw_png = make_noisy_png_bytes(4096, 4096);
|
||||
assert!(
|
||||
raw_png.len() > DEFAULT_IMAGE_MAX_BYTES,
|
||||
"Test image should exceed the 5 MB limit (actual: {} bytes)",
|
||||
raw_png.len()
|
||||
);
|
||||
|
||||
let image = Arc::new(gpui::Image::from_bytes(ImageFormat::Png, raw_png));
|
||||
let lm_image = cx
|
||||
.update(|cx| LanguageModelImage::from_image(Arc::clone(&image), cx))
|
||||
.await
|
||||
.expect("from_image should succeed");
|
||||
|
||||
let decoded_png = base64_to_png_bytes(lm_image.source.as_ref());
|
||||
assert!(
|
||||
decoded_png.len() <= DEFAULT_IMAGE_MAX_BYTES,
|
||||
"Encoded PNG should be ≤ {} bytes after downscale, but was {} bytes",
|
||||
DEFAULT_IMAGE_MAX_BYTES,
|
||||
decoded_png.len()
|
||||
);
|
||||
|
||||
let (w, h) = png_dimensions(&decoded_png);
|
||||
assert!(
|
||||
w < 4096 && h < 4096,
|
||||
"Dimensions should have shrunk: got {}×{}",
|
||||
w,
|
||||
h
|
||||
);
|
||||
|
||||
let size = lm_image.size.expect("ImageSize should be present");
|
||||
assert_eq!(
|
||||
size.width, w as i32,
|
||||
"ImageSize.width should match the encoded PNG width after downscaling"
|
||||
);
|
||||
assert_eq!(
|
||||
size.height, h as i32,
|
||||
"ImageSize.height should match the encoded PNG height after downscaling"
|
||||
);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user