Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
107 changes: 41 additions & 66 deletions crates/collab/src/api/extensions.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,16 +4,15 @@ use anyhow::Context as _;
use aws_sdk_s3::presigning::PresigningConfig;
use axum::{
Extension, Json, Router,
extract::{Path, Query},
extract::{Path, Query, RawQuery},
http::StatusCode,
response::Redirect,
routing::get,
};
use cloud_api_types::{ExtensionApiManifest, ExtensionProvides, GetExtensionsResponse};
use collections::{BTreeSet, HashMap};
use cloud_api_types::{ExtensionApiManifest, GetExtensionsResponse};
use collections::HashMap;
use semver::Version as SemanticVersion;
use serde::Deserialize;
use std::str::FromStr;
use std::{sync::Arc, time::Duration};
use time::PrimitiveDateTime;
use util::{ResultExt, maybe};
Expand All @@ -33,74 +32,50 @@ pub fn router() -> Router {
)
}

#[derive(Debug, Deserialize)]
struct GetExtensionsParams {
filter: Option<String>,
/// A comma-delimited list of features that the extension must provide.
///
/// For example:
/// - `themes`
/// - `themes,icon-themes`
/// - `languages,language-servers`
#[serde(default)]
provides: Option<String>,
#[serde(default)]
max_schema_version: i32,
}
const UPSTREAM_EXTENSIONS_URL: &str = "https://cloud.zed.dev/extensions";

async fn get_extensions(
Extension(app): Extension<Arc<AppState>>,
Query(params): Query<GetExtensionsParams>,
) -> Result<Json<GetExtensionsResponse>> {
let provides_filter = params.provides.map(|provides| {
provides
.split(',')
.map(|value| value.trim())
.filter_map(|value| ExtensionProvides::from_str(value).ok())
.collect::<BTreeSet<_>>()
});
async fn get_extensions(RawQuery(query): RawQuery) -> Result<Json<GetExtensionsResponse>> {
let upstream_url = match query {
Some(query) => format!("{UPSTREAM_EXTENSIONS_URL}?{query}"),
None => UPSTREAM_EXTENSIONS_URL.to_string(),
};

let mut extensions = app
.db
.get_extensions(
params.filter.as_deref(),
provides_filter.as_ref(),
params.max_schema_version,
1_000,
let response = reqwest::get(&upstream_url).await.map_err(|error| {
tracing::error!(
?error,
"failed to proxy request to upstream extensions service"
);
Error::http(
StatusCode::BAD_GATEWAY,
"upstream extensions service unavailable".into(),
)
.await?;

if let Some(filter) = params.filter.as_deref() {
let extension_id = filter.to_lowercase();
let mut exact_match = None;
extensions.retain(|extension| {
if extension.id.as_ref() == extension_id {
exact_match = Some(extension.clone());
false
} else {
true
}
});
if exact_match.is_none() {
exact_match = app
.db
.get_extensions_by_ids(&[&extension_id], None)
.await?
.first()
.cloned();
}

if let Some(exact_match) = exact_match {
extensions.insert(0, exact_match);
}
};
})?;

if let Some(query) = params.filter.as_deref() {
let count = extensions.len();
tracing::info!(query, count, "extension_search")
let status = response.status();
if !status.is_success() {
let upstream_status =
StatusCode::from_u16(status.as_u16()).unwrap_or(StatusCode::BAD_GATEWAY);
let body = response.text().await.unwrap_or_default();
tracing::error!(
status = status.as_u16(),
body,
"upstream extensions service returned an error"
);
return Err(Error::http(upstream_status, body));
}

Ok(Json(GetExtensionsResponse { data: extensions }))
let body: GetExtensionsResponse = response.json().await.map_err(|error| {
tracing::error!(
?error,
"failed to parse response from upstream extensions service"
);
Error::http(
StatusCode::BAD_GATEWAY,
"failed to parse upstream response".into(),
)
})?;

Ok(Json(body))
}

#[derive(Debug, Deserialize)]
Expand Down
81 changes: 0 additions & 81 deletions crates/collab/src/db/queries/extensions.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,36 +8,6 @@ use util::ResultExt;
use super::*;

impl Database {
pub async fn get_extensions(
&self,
filter: Option<&str>,
provides_filter: Option<&BTreeSet<ExtensionProvides>>,
max_schema_version: i32,
limit: usize,
) -> Result<Vec<ExtensionMetadata>> {
self.transaction(|tx| async move {
let mut condition = Condition::all()
.add(
extension::Column::LatestVersion
.into_expr()
.eq(extension_version::Column::Version.into_expr()),
)
.add(extension_version::Column::SchemaVersion.lte(max_schema_version));
if let Some(filter) = filter {
let fuzzy_name_filter = Self::fuzzy_like_string(filter);
condition = condition.add(Expr::cust_with_expr("name ILIKE $1", fuzzy_name_filter));
}

if let Some(provides_filter) = provides_filter {
condition = apply_provides_filter(condition, provides_filter);
}

self.get_extensions_where(condition, Some(limit as u64), &tx)
.await
})
.await
}

pub async fn get_extensions_by_ids(
&self,
ids: &[&str],
Expand Down Expand Up @@ -396,57 +366,6 @@ impl Database {
}
}

fn apply_provides_filter(
mut condition: Condition,
provides_filter: &BTreeSet<ExtensionProvides>,
) -> Condition {
if provides_filter.contains(&ExtensionProvides::Themes) {
condition = condition.add(extension_version::Column::ProvidesThemes.eq(true));
}

if provides_filter.contains(&ExtensionProvides::IconThemes) {
condition = condition.add(extension_version::Column::ProvidesIconThemes.eq(true));
}

if provides_filter.contains(&ExtensionProvides::Languages) {
condition = condition.add(extension_version::Column::ProvidesLanguages.eq(true));
}

if provides_filter.contains(&ExtensionProvides::Grammars) {
condition = condition.add(extension_version::Column::ProvidesGrammars.eq(true));
}

if provides_filter.contains(&ExtensionProvides::LanguageServers) {
condition = condition.add(extension_version::Column::ProvidesLanguageServers.eq(true));
}

if provides_filter.contains(&ExtensionProvides::ContextServers) {
condition = condition.add(extension_version::Column::ProvidesContextServers.eq(true));
}

if provides_filter.contains(&ExtensionProvides::AgentServers) {
condition = condition.add(extension_version::Column::ProvidesAgentServers.eq(true));
}

if provides_filter.contains(&ExtensionProvides::SlashCommands) {
condition = condition.add(extension_version::Column::ProvidesSlashCommands.eq(true));
}

if provides_filter.contains(&ExtensionProvides::IndexedDocsProviders) {
condition = condition.add(extension_version::Column::ProvidesIndexedDocsProviders.eq(true));
}

if provides_filter.contains(&ExtensionProvides::Snippets) {
condition = condition.add(extension_version::Column::ProvidesSnippets.eq(true));
}

if provides_filter.contains(&ExtensionProvides::DebugAdapters) {
condition = condition.add(extension_version::Column::ProvidesDebugAdapters.eq(true));
}

condition
}

fn metadata_from_extension_and_version(
extension: extension::Model,
version: extension_version::Model,
Expand Down
Loading