use crate::Transport;
use reqwest::IntoUrl;
use unitycatalog_common::models::temporary_credentials::v1::TemporaryCredential;
use unitycatalog_common::{
model_versions::v1::GetModelVersionRequest,
models::temporary_credentials::v1::{
GenerateTemporaryModelVersionCredentialsRequest, GenerateTemporaryPathCredentialsRequest,
GenerateTemporaryTableCredentialsRequest, GenerateTemporaryVolumeCredentialsRequest,
generate_temporary_model_version_credentials_request::Operation as MvOperation,
generate_temporary_path_credentials_request::Operation as PthOperation,
generate_temporary_table_credentials_request::Operation as TblOperation,
generate_temporary_volume_credentials_request::Operation as VolOperation,
},
tables::v1::GetTableRequest,
volumes::v1::GetVolumeRequest,
};
use url::Url;
use uuid::Uuid;
use crate::Result;
use crate::codegen::tables::TableServiceClient;
pub(super) use crate::codegen::temporary_credentials::TemporaryCredentialClient as TemporaryCredentialClientBase;
use crate::codegen::volumes::client::VolumeServiceClient;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum TableReference {
Id(Uuid),
Name(String),
}
impl From<String> for TableReference {
fn from(name: String) -> Self {
TableReference::Name(name)
}
}
impl From<&str> for TableReference {
fn from(name: &str) -> Self {
TableReference::Name(name.to_string())
}
}
impl From<Uuid> for TableReference {
fn from(id: Uuid) -> Self {
TableReference::Id(id)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum VolumeReference {
Id(Uuid),
Name(String),
}
impl From<String> for VolumeReference {
fn from(name: String) -> Self {
VolumeReference::Name(name)
}
}
impl From<&str> for VolumeReference {
fn from(name: &str) -> Self {
VolumeReference::Name(name.to_string())
}
}
impl From<Uuid> for VolumeReference {
fn from(id: Uuid) -> Self {
VolumeReference::Id(id)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TableOperation {
Read,
ReadWrite,
}
impl From<TableOperation> for i32 {
fn from(operation: TableOperation) -> Self {
match operation {
TableOperation::Read => TblOperation::Read as i32,
TableOperation::ReadWrite => TblOperation::ReadWrite as i32,
}
}
}
impl From<TableOperation> for TblOperation {
fn from(operation: TableOperation) -> Self {
match operation {
TableOperation::Read => TblOperation::Read,
TableOperation::ReadWrite => TblOperation::ReadWrite,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PathOperation {
Read,
ReadWrite,
CreateTable,
}
impl From<PathOperation> for i32 {
fn from(operation: PathOperation) -> Self {
match operation {
PathOperation::Read => PthOperation::PathRead as i32,
PathOperation::ReadWrite => PthOperation::PathReadWrite as i32,
PathOperation::CreateTable => PthOperation::PathCreateTable as i32,
}
}
}
impl From<PathOperation> for PthOperation {
fn from(operation: PathOperation) -> Self {
match operation {
PathOperation::Read => PthOperation::PathRead,
PathOperation::ReadWrite => PthOperation::PathReadWrite,
PathOperation::CreateTable => PthOperation::PathCreateTable,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum VolumeOperation {
Read,
ReadWrite,
}
impl From<VolumeOperation> for i32 {
fn from(operation: VolumeOperation) -> Self {
match operation {
VolumeOperation::Read => VolOperation::ReadVolume as i32,
VolumeOperation::ReadWrite => VolOperation::WriteVolume as i32,
}
}
}
impl From<VolumeOperation> for VolOperation {
fn from(operation: VolumeOperation) -> Self {
match operation {
VolumeOperation::Read => VolOperation::ReadVolume,
VolumeOperation::ReadWrite => VolOperation::WriteVolume,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ModelVersionOperation {
Read,
ReadWrite,
}
impl From<ModelVersionOperation> for i32 {
fn from(operation: ModelVersionOperation) -> Self {
match operation {
ModelVersionOperation::Read => MvOperation::ReadModelVersion as i32,
ModelVersionOperation::ReadWrite => MvOperation::ReadWriteModelVersion as i32,
}
}
}
impl From<ModelVersionOperation> for MvOperation {
fn from(operation: ModelVersionOperation) -> Self {
match operation {
ModelVersionOperation::Read => MvOperation::ReadModelVersion,
ModelVersionOperation::ReadWrite => MvOperation::ReadWriteModelVersion,
}
}
}
#[derive(Clone)]
pub struct TemporaryCredentialClient {
client: TemporaryCredentialClientBase,
}
impl TemporaryCredentialClient {
pub fn new_with_url(client: Transport, mut base_url: Url) -> Self {
if !base_url.path().ends_with('/') {
base_url.set_path(&format!("{}/", base_url.path()));
}
Self {
client: TemporaryCredentialClientBase::new(client, base_url),
}
}
pub fn new(client: TemporaryCredentialClientBase) -> Self {
Self { client }
}
async fn post_credential<R: serde::Serialize>(
&self,
path: &str,
request: &R,
) -> Result<TemporaryCredential> {
let url = self.client.base_url.join(path)?;
let response = self
.client
.client
.post(url)
.json(request)
.send_raw()
.await
.map_err(crate::Error::from_api_send)?;
let response = crate::error::check_api_response(response).await?;
let bytes = response.bytes().await?;
let mut value: serde_json::Value = serde_json::from_slice(&bytes)?;
if let Some(obj) = value.as_object_mut() {
const ONEOF_KEYS: [&str; 10] = [
"aws_temp_credentials",
"awsTempCredentials",
"azure_user_delegation_sas",
"azureUserDelegationSas",
"azure_aad",
"azureAad",
"gcp_oauth_token",
"gcpOauthToken",
"r2_temp_credentials",
"r2TempCredentials",
];
obj.retain(|key, v| !(v.is_null() && ONEOF_KEYS.contains(&key.as_str())));
}
Ok(serde_json::from_value(value)?)
}
pub async fn temporary_table_credential(
&self,
table: impl Into<TableReference>,
operation: TableOperation,
) -> Result<(TemporaryCredential, Uuid)> {
let (table_id, storage_location) = match table.into() {
TableReference::Id(id) => (id.as_hyphenated().to_string(), None),
TableReference::Name(name) => {
let table_client = TableServiceClient::new(
self.client.client.clone(),
self.client.base_url.clone(),
);
let table_info = table_client
.get_table(&GetTableRequest {
full_name: name,
include_browse: Some(false),
include_delta_metadata: Some(false),
include_manifest_capabilities: Some(false),
..Default::default()
})
.await?;
(
table_info.table_id.clone().unwrap_or_default(),
table_info.storage_location.clone(),
)
}
};
let uuid =
Uuid::parse_str(&table_id).map_err(unitycatalog_common::Error::InvalidIdentifier)?;
let mut credential = self
.post_credential(
"temporary-table-credentials",
&GenerateTemporaryTableCredentialsRequest {
table_id,
operation: TblOperation::from(operation).into(),
..Default::default()
},
)
.await?;
backfill_credential_url(&mut credential, storage_location);
Ok((credential, uuid))
}
pub async fn temporary_path_credential(
&self,
path: impl IntoUrl,
operation: PathOperation,
dry_run: impl Into<Option<bool>>,
) -> Result<(TemporaryCredential, Url)> {
let url = path.into_url()?;
Ok((
self.post_credential(
"temporary-path-credentials",
&GenerateTemporaryPathCredentialsRequest {
url: url.to_string(),
operation: PthOperation::from(operation).into(),
dry_run: dry_run.into(),
..Default::default()
},
)
.await?,
url,
))
}
pub async fn temporary_volume_credential(
&self,
volume: impl Into<VolumeReference>,
operation: VolumeOperation,
) -> Result<(TemporaryCredential, Uuid)> {
let (volume_id, storage_location) = match volume.into() {
VolumeReference::Id(id) => (id.as_hyphenated().to_string(), None),
VolumeReference::Name(name) => {
let volume_client = VolumeServiceClient::new(
self.client.client.clone(),
self.client.base_url.clone(),
);
let info = volume_client
.get_volume(&GetVolumeRequest {
name,
include_browse: Some(false),
..Default::default()
})
.await?;
(info.volume_id, Some(info.storage_location))
}
};
let uuid =
Uuid::parse_str(&volume_id).map_err(unitycatalog_common::Error::InvalidIdentifier)?;
let mut credential = self
.post_credential(
"temporary-volume-credentials",
&GenerateTemporaryVolumeCredentialsRequest {
volume_id,
operation: VolOperation::from(operation).into(),
..Default::default()
},
)
.await?;
backfill_credential_url(&mut credential, storage_location);
Ok((credential, uuid))
}
pub async fn temporary_model_version_credential(
&self,
full_name: impl Into<String>,
version: i64,
operation: ModelVersionOperation,
) -> Result<TemporaryCredential> {
let full_name = full_name.into();
let [catalog_name, schema_name, model_name] =
<[String; 3]>::try_from(full_name.split('.').map(str::to_string).collect::<Vec<_>>())
.map_err(|_| {
unitycatalog_common::Error::invalid_argument(
"full_name must be a three-level catalog.schema.model name",
)
})?;
let mv_client = crate::codegen::model_versions::ModelVersionServiceClient::new(
self.client.client.clone(),
self.client.base_url.clone(),
);
let storage_location = mv_client
.get_model_version(&GetModelVersionRequest {
full_name: full_name.clone(),
version,
..Default::default()
})
.await
.ok()
.and_then(|mv| mv.storage_location);
let mut credential = self
.post_credential(
"temporary-model-version-credentials",
&GenerateTemporaryModelVersionCredentialsRequest {
catalog_name,
schema_name,
model_name,
version,
operation: MvOperation::from(operation).into(),
..Default::default()
},
)
.await?;
backfill_credential_url(&mut credential, storage_location);
Ok(credential)
}
}
fn backfill_credential_url(credential: &mut TemporaryCredential, storage_location: Option<String>) {
if credential.url.is_empty()
&& let Some(location) = storage_location
&& !location.is_empty()
{
credential.url = location;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn malformed_table_id_is_error_not_panic() {
let result =
Uuid::parse_str("not-a-uuid").map_err(unitycatalog_common::Error::InvalidIdentifier);
let err: crate::Error = result.unwrap_err().into();
assert!(matches!(err, crate::Error::Common { .. }));
}
}