Fixes found through integration test crate

This commit is contained in:
Scott Lyons 2024-08-10 00:26:07 -07:00
commit f17a214865
12 changed files with 815 additions and 63 deletions

View file

@ -2,12 +2,14 @@
use crate::{errors::*, models::*, tokens::*};
use base64::{engine::general_purpose::STANDARD, Engine as _};
use reqwest::header::CONTENT_TYPE;
use reqwest::{
header::{HeaderMap, ACCEPT, AUTHORIZATION},
multipart::{Form, Part},
Client, ClientBuilder, Method, RequestBuilder, Response,
};
use serde::{de::DeserializeOwned, Serialize};
use serde_json::Value;
use sha1::{Digest, Sha1};
use std::fmt::{Display, Formatter};
use std::path::Path;
@ -88,6 +90,12 @@ impl SzurubooruClient {
SzurubooruClient::new(host, auth, allow_insecure)
}
/// Create a new client with anonymous credentials
pub fn new_anonymous(host: &str, allow_insecure: bool) -> SzurubooruResult<Self> {
let auth = SzurubooruAuth::None;
SzurubooruClient::new(host, auth, allow_insecure)
}
fn new(host: &str, auth: SzurubooruAuth, allow_insecure: bool) -> SzurubooruResult<Self> {
let host = if host.ends_with("/") {
&host[0..host.len() - 1]
@ -103,6 +111,7 @@ impl SzurubooruClient {
let mut header_map = HeaderMap::new();
//header_map.append(AUTHORIZATION, token_header_value.parse().unwrap());
header_map.append(ACCEPT, "application/json".parse().unwrap());
header_map.append(CONTENT_TYPE, "application/json".parse().unwrap());
let client = ClientBuilder::new()
.danger_accept_invalid_certs(allow_insecure)
@ -288,7 +297,7 @@ impl<'a> SzurubooruRequest<'a> {
T: AsRef<str> + Display,
{
let mut req_url = self.client.base_url.clone();
req_url.set_path(path.as_ref());
req_url.set_path(&format!("/api{}", path.as_ref()));
if let Some(query_vec) = query {
let mut qpm = req_url.query_pairs_mut();
@ -315,7 +324,7 @@ impl<'a> SzurubooruRequest<'a> {
// This doesn't detect the required `mut` for some reason
#[allow(unused_mut)]
let mut req = self.client.client.request(method, req_url);
match &self.client.auth {
let req = match &self.client.auth {
SzurubooruAuth::TokenAuth(t) => {
let mut header_map = HeaderMap::new();
header_map.append(AUTHORIZATION, t.parse().unwrap());
@ -324,7 +333,8 @@ impl<'a> SzurubooruRequest<'a> {
}
SzurubooruAuth::BasicAuth(u, p) => req.basic_auth(u, Some(p)),
SzurubooruAuth::None => req,
}
};
req
}
#[tracing::instrument(skip(self), fields(base_url=self.client.base_url.to_string()))]
@ -351,6 +361,20 @@ impl<'a> SzurubooruRequest<'a> {
self.handle_request(request).await
}
async fn handle_response(&self, response: Response) -> SzurubooruResult<Response> {
if response.status().is_client_error() || response.status().is_server_error() {
let resp_json = response
.text()
.await
.map_err(SzurubooruClientError::RequestError)?;
let server_error = serde_json::from_str::<SzurubooruServerError>(&resp_json)
.map_err(|e| SzurubooruClientError::ResponseParsingError(e, resp_json))?;
Err(SzurubooruClientError::SzurubooruServerError(server_error))
} else {
Ok(response)
}
}
async fn handle_request<T: DeserializeOwned>(
&self,
request: RequestBuilder,
@ -361,10 +385,11 @@ impl<'a> SzurubooruRequest<'a> {
let response = self.client.client.execute(request).await;
let response = response
.map_err(SzurubooruClientError::RequestError)?
.error_for_status()
.map_err(SzurubooruClientError::RequestError)?;
let response = self
.handle_response(response.map_err(SzurubooruClientError::RequestError)?)
.await?;
//.error_for_status()
//.map_err(SzurubooruClientError::RequestError)?;
let response_text = response
.text()
@ -400,13 +425,13 @@ impl<'a> SzurubooruRequest<'a> {
pub async fn update_tag_category<T>(
&self,
name: T,
resource: &TagCategoryResource,
update_tag_cat: &CreateUpdateTagCategory,
) -> SzurubooruResult<TagCategoryResource>
where
T: AsRef<str> + Display,
{
let path = format!("/tag-category/{name}");
self.do_request(Method::PUT, &path, None, Some(resource))
self.do_request(Method::PUT, &path, None, Some(update_tag_cat))
.await
}
@ -427,8 +452,9 @@ impl<'a> SzurubooruRequest<'a> {
{
let path = format!("/tag-category/{name}");
let version_obj = ResourceVersion { version };
self.do_request(Method::DELETE, &path, None, Some(&version_obj))
self.do_request::<Value, _, _>(Method::DELETE, &path, None, Some(&version_obj))
.await
.map(|_| ())
}
/// Sets given tag category as default. All new tags created manually or automatically will
@ -447,9 +473,9 @@ impl<'a> SzurubooruRequest<'a> {
/// all possible query tokens, or use (QueryToken)[tokens::QueryToken] for a custom token
pub async fn list_tags(
&self,
query: &Vec<QueryToken>,
query: Option<&Vec<QueryToken>>,
) -> SzurubooruResult<PagedSearchResult<TagResource>> {
self.do_request(Method::GET, "/tags", Some(query), None::<&String>)
self.do_request(Method::GET, "/tags", query, None::<&String>)
.await
}
@ -502,8 +528,9 @@ impl<'a> SzurubooruRequest<'a> {
{
let path = format!("/tag/{name}");
let version_obj = ResourceVersion { version };
self.do_request(Method::DELETE, &path, None, Some(&version_obj))
self.do_request::<Value, _, _>(Method::DELETE, &path, None, Some(&version_obj))
.await
.map(|_| ())
}
/// Removes source tag and merges all of its usages, suggestions and implications to the
@ -697,11 +724,13 @@ impl<'a> SzurubooruRequest<'a> {
.build()
.map_err(SzurubooruClientError::RequestBuilderError)?;
self.client
let resp_res = self
.client
.client
.execute(request)
.await
.map_err(SzurubooruClientError::RequestError)
.map_err(|e| SzurubooruClientError::RequestError(e))?;
self.handle_response(resp_res).await
}
///Downloads the given post ID's image as a stream of bytes
@ -718,10 +747,11 @@ impl<'a> SzurubooruRequest<'a> {
///Downloads the given post ID's image as a (Bytes)[bytes::Bytes] struct
pub async fn get_post_content_bytes(&self, post_id: u32) -> SzurubooruResult<bytes::Bytes> {
let content_response = self.get_post_content(post_id).await?;
content_response
.bytes()
.await
.map_err(SzurubooruClientError::RequestError)
.map_err(|e| SzurubooruClientError::RequestError(e))
}
/// Retrieves posts that look like the input image

View file

@ -5,9 +5,10 @@
//! See [here](https://github.com/rr-/szurubooru/blob/master/doc/API.md#field-selecting) for
//! more information.
use chrono::NaiveDateTime;
use chrono::{DateTime, Utc};
use derive_builder::Builder;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use strum_macros::AsRefStr;
#[derive(Serialize, Deserialize, Debug, Clone)]
@ -98,9 +99,9 @@ pub struct TagResource {
/// the user by the web client on usage
pub suggestions: Option<Vec<MicroTagResource>>,
/// time the tag was created
pub creation_time: Option<NaiveDateTime>,
pub creation_time: Option<DateTime<Utc>>,
/// time the tag was edited
pub last_edit_time: Option<NaiveDateTime>,
pub last_edit_time: Option<DateTime<Utc>>,
/// the number of posts the tag was used in
pub usages: Option<u32>,
/// the tag description (instructions how to use, history etc.) The client should render
@ -128,21 +129,27 @@ pub struct TagResource {
#[builder(setter(strip_option))]
pub struct CreateUpdateTag {
#[serde(skip_serializing_if = "Option::is_none")]
#[builder(default)]
/// resource version. See [versioning](ResourceVersion)
pub version: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
#[builder(default)]
/// Tag names and aliases, must match `tag_name_regex` from the server's configuration
pub names: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
#[builder(default)]
/// Category that this tag belongs to. Must already exist
pub category: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
#[builder(default)]
/// The tag description in Markdown format
pub description: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
#[builder(default)]
/// Tags that should be implied when this tag is used
pub implications: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
#[builder(default)]
/// Tags that should be suggested when this tag is used
pub suggestions: Option<Vec<String>>,
}
@ -160,21 +167,31 @@ pub struct TagCategoryResource {
/// How many tags is the given category used with
pub usages: Option<u32>,
/// The order in which tags with this category are displayed, ascending
pub order: Option<String>,
pub order: Option<u32>,
/// Whether the tag category is the default one
pub default: Option<bool>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Builder)]
#[builder(setter(into))]
#[derive(Debug, Clone, Serialize, Deserialize, Default, Builder)]
#[builder(setter(strip_option))]
/// Used for creating or updating a Tag Category
pub struct CreateUpdateTagCategory {
/// Resource version. See [versioning](ResourceVersion)
#[serde(skip_serializing_if = "Option::is_none")]
#[builder(default)]
pub version: Option<u32>,
/// The name of the category to create
pub name: String,
#[serde(skip_serializing_if = "Option::is_none")]
#[builder(default)]
pub name: Option<String>,
/// The display color to use for the category
pub color: String,
#[serde(skip_serializing_if = "Option::is_none")]
#[builder(default)]
pub color: Option<String>,
/// The order in which tags with this category are displayed, ascending
pub order: String,
#[serde(skip_serializing_if = "Option::is_none")]
#[builder(default)]
pub order: Option<u32>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Builder)]
@ -261,9 +278,9 @@ pub struct PostResource {
/// The post identifier
pub id: Option<u32>,
/// Time the post was created
pub creation_time: Option<NaiveDateTime>,
pub creation_time: Option<DateTime<Utc>>,
/// Time the post was edited
pub last_edit_time: Option<NaiveDateTime>,
pub last_edit_time: Option<DateTime<Utc>>,
/// Whether the post is safe for work
pub safety: Option<PostSafety>,
#[serde(rename = "type")]
@ -311,7 +328,7 @@ pub struct PostResource {
/// How many posts are related to this post
pub relation_count: Option<u32>,
/// The last time the post was featured
pub last_feature_time: Option<NaiveDateTime>,
pub last_feature_time: Option<DateTime<Utc>>,
/// List of users who have favorited this post
pub favorited_by: Option<Vec<MicroUserResource>>,
/// Whether the post uses custom thumbnail
@ -326,7 +343,7 @@ pub struct PostResource {
}
#[derive(Debug, Clone, Serialize, Deserialize, Builder)]
#[builder(setter(into, strip_option))]
#[builder(setter(strip_option))]
#[serde(rename_all = "camelCase")]
/// A `struct` used to create or update a post. For updating purposes
/// the [version](CreateUpdatePost::version) field is required
@ -447,10 +464,10 @@ pub struct UserResource {
pub rank: Option<UserRank>,
#[serde(rename = "last-login-time")]
/// The last login time
pub last_login_time: Option<NaiveDateTime>,
pub last_login_time: Option<DateTime<Utc>>,
#[serde(rename = "creation-time")]
/// The user registration time
pub creation_time: Option<NaiveDateTime>,
pub creation_time: Option<DateTime<Utc>>,
/// How to render the user avatar
pub avatar_style: Option<UserAvatarStyle>,
/// The URL to the avatar
@ -474,7 +491,7 @@ pub struct UserResource {
pub favorite_post_count: Option<SzuruEither<u32, bool>>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Builder)]
#[derive(Debug, Clone, Serialize, Deserialize, Default, Builder)]
#[builder(setter(into, strip_option))]
#[serde(rename_all = "camelCase")]
/// `struct` used to create or update a user resource. The version field is only used when
@ -482,15 +499,20 @@ pub struct UserResource {
pub struct CreateUpdateUser {
/// Resource version. See [versioning](ResourceVersion)
#[serde(skip_serializing_if = "Option::is_none")]
#[builder(default)]
pub version: Option<u32>,
/// The username
#[builder(default)]
pub name: Option<String>,
/// The user's password
#[builder(default)]
pub password: Option<String>,
/// The user's desired rank, if not given will default to `default_rank` in the server's
/// configuration
#[builder(default)]
pub rank: Option<UserRank>,
/// The user avatar style, Gravatar or Manual
#[builder(default)]
pub avatar_style: Option<UserAvatarStyle>,
}
@ -517,15 +539,15 @@ pub struct UserAuthTokenResource {
/// Whether the token is still valid for authentication
pub enabled: Option<bool>,
/// Time when the token expires
pub expiration_time: Option<NaiveDateTime>,
pub expiration_time: Option<DateTime<Utc>>,
/// Resource version. See [versioning](ResourceVersion)
pub version: Option<u32>,
/// time the user token was created
pub creation_time: Option<NaiveDateTime>,
pub creation_time: Option<DateTime<Utc>>,
/// time the user token was edited
pub last_edit_time: Option<NaiveDateTime>,
pub last_edit_time: Option<DateTime<Utc>>,
/// the last time this token was used
pub last_usage_time: Option<NaiveDateTime>,
pub last_usage_time: Option<DateTime<Utc>>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Builder, Default)]
@ -545,7 +567,7 @@ pub struct CreateUpdateUserAuthToken {
pub note: Option<String>,
/// Time when the token expires
#[serde(skip_serializing_if = "Option::is_none")]
pub expiration_time: Option<NaiveDateTime>,
pub expiration_time: Option<DateTime<Utc>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
@ -578,8 +600,14 @@ pub struct GlobalInfoConfig {
pub tag_category_name_regex: String,
/// Default user rank upon signup
pub default_user_rank: String,
/// Whether safety is enabled
pub enable_safety: bool,
/// Contact email for this server
pub contact_email: Option<String>,
/// Is sending email enabled for this server
pub can_send_mails: bool,
/// Available privileges enabled for this server
pub privileges: Vec<String>,
pub privileges: HashMap<String, String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
@ -593,11 +621,11 @@ pub struct GlobalInfo {
/// The current featured post
pub featured_post: Option<u32>,
/// The time the current featured post was featured
pub featuring_time: Option<NaiveDateTime>,
pub featuring_time: Option<DateTime<Utc>>,
/// The user who uploaded the featured post
pub featuring_user: Option<u32>,
/// The current server time
pub server_time: NaiveDateTime,
pub server_time: DateTime<Utc>,
/// The configuration for this server
pub config: GlobalInfoConfig,
}
@ -660,9 +688,9 @@ pub struct PoolResource {
/// An ordered list of posts. Posts are ordered by insertion by default
pub posts: Option<Vec<MicroPostResource>>,
/// Time the pool was created
pub creation_time: Option<NaiveDateTime>,
pub creation_time: Option<DateTime<Utc>>,
/// Time the pool was edited
pub last_edit_time: Option<NaiveDateTime>,
pub last_edit_time: Option<DateTime<Utc>>,
/// The total number of posts the pool has
pub post_count: Option<u32>,
/// The pool description (instructions how to use, history etc). The client should render
@ -758,9 +786,9 @@ pub struct CommentResource {
/// The text of the comment
pub text: Option<String>,
/// When was the comment posted
pub creation_time: Option<NaiveDateTime>,
pub creation_time: Option<DateTime<Utc>>,
/// When was the last time this comment was edited
pub last_edit_time: Option<NaiveDateTime>,
pub last_edit_time: Option<DateTime<Utc>>,
/// The sum of the -1/0/+1 scores by other users
pub score: Option<i32>,
/// The user's own score for this comment
@ -890,7 +918,7 @@ pub struct SnapshotResource {
/// The data associated with this resource change
pub data: Option<SnapshotData>,
/// When this resource change occurred
pub time: Option<NaiveDateTime>,
pub time: Option<DateTime<Utc>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
@ -923,3 +951,78 @@ pub struct AroundPostResult {
/// The next post, if it exists
next: Option<u32>,
}
#[cfg(test)]
mod tests {
use crate::models::{GlobalInfo, GlobalInfoConfig, TagCategoryResource};
use chrono::Datelike;
#[test]
fn test_parse_global_info() {
let cfg_str = r#"{
"name": "integrationland",
"userNameRegex": "^[a-zA-Z0-9_-]{1,32}$",
"passwordRegex": "^.{5,}$",
"tagNameRegex": "^\\S+$",
"tagCategoryNameRegex": "^[^\\s%+#/]+$",
"defaultUserRank": "regular",
"enableSafety": true,
"contactEmail": null,
"canSendMails": false,
"privileges": {
"users:create:self": "anonymous",
"users:create:any": "administrator",
"comments:edit:own": "regular",
"comments:list": "regular",
"comments:view": "regular",
"comments:score": "regular",
"snapshots:list": "power",
"uploads:create": "regular",
"uploads:useDownloader": "power"
}
}"#;
let global_config =
serde_json::from_str::<GlobalInfoConfig>(cfg_str).expect("Unable to parse cfg_str");
assert_eq!(global_config.can_send_mails, false);
let info_str = r#"{"postCount": 0,
"diskUsage": 0,
"serverTime": "2024-08-09T21:41:24.123623Z",
"config": {
"name": "integrationland",
"userNameRegex": "^[a-zA-Z0-9_-]{1,32}$",
"passwordRegex": "^.{5,}$",
"tagNameRegex": "^\\S+$",
"tagCategoryNameRegex": "^[^\\s%+#/]+$",
"defaultUserRank": "regular",
"enableSafety": true,
"contactEmail": null,
"canSendMails": false,
"privileges": {
"users:create:self": "anonymous"
}
},
"featuredPost": null,
"featuringUser": null,
"featuringTime": null
}"#;
let global_info =
serde_json::from_str::<GlobalInfo>(info_str).expect("Unable to parse info_str");
assert_eq!(global_info.server_time.year(), 2024);
}
#[test]
fn test_parse_tag_category_resource() {
let input_str = r#" {
"name": "default",
"version": 1,
"color": "default",
"usages": 0,
"default": true,
"order": 1
}"#;
let tag_cat = serde_json::from_str::<TagCategoryResource>(input_str)
.expect("Unable to parse tag category string");
assert_eq!(tag_cat.name, Some("default".to_string()));
}
}