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:
Mohamad Khani
2026-07-14 01:52:12 +03:30
commit b9819977a5
3984 changed files with 1487015 additions and 0 deletions

View 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"] }

View File

@@ -0,0 +1 @@
../../LICENSE-GPL

View 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"),
}
}
}

View 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
}
}

View 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))
}
}
}

View File

@@ -0,0 +1,3 @@
pub mod cloud_model;
pub use cloud_model::*;

View 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."
)
}
}

View 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);
}
}

View 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"
);
}
}