diff --git a/szurubooru-client/Cargo.toml b/szurubooru-client/Cargo.toml index 8a4565b..a462d75 100644 --- a/szurubooru-client/Cargo.toml +++ b/szurubooru-client/Cargo.toml @@ -17,7 +17,7 @@ derive_builder = "0.20.0" futures-util = "0.3.30" hex = "0.4.3" openssl = { version = "0.10.66", features = ["vendored"] } -pyo3 = { version="0.22.0", optional=true, features=["chrono-tz", "chrono", "extension-module", "serde", "experimental-async"] } +pyo3 = { version="0.22.0", optional=true, features=["chrono-tz", "chrono", "serde", "experimental-async"] } reqwest = { version = "0.12.5", features = ["json", "multipart", "stream"] } serde = { version = "1.0.204", features = ["derive"] } serde-pyobject = { version = "0.4.0", optional = true } @@ -36,9 +36,9 @@ mockito = "1.4.0" tokio = { version = "1.39.2", features = ["full"] } [features] -python = ["dep:pyo3", "dep:tokio", "dep:serde-pyobject"] -serde-pyobject = ["dep:serde-pyobject"] +python = ["dep:pyo3", "dep:tokio", "dep:serde-pyobject", "pyo3/extension-module"] +extension-module = ["pyo3/extension-module"] [lib] name = "szurubooru_client" -crate-type = ["cdylib"] +crate-type = ["cdylib", "lib"] diff --git a/szurubooru-client/pyproject.toml b/szurubooru-client/pyproject.toml index f0eacf9..ef61234 100644 --- a/szurubooru-client/pyproject.toml +++ b/szurubooru-client/pyproject.toml @@ -3,7 +3,7 @@ requires = ["maturin>=1.7,<2.0"] build-backend = "maturin" [project] -name = "szurubooru-client" +name = "szurubooru_client" requires-python = ">=3.8" classifiers = [ "Programming Language :: Rust", @@ -12,4 +12,4 @@ classifiers = [ ] dynamic = ["version"] [tool.maturin] -features = ["pyo3/extension-module"] +features = ["pyo3/extension-module", "python"] diff --git a/szurubooru-client/src/client.rs b/szurubooru-client/src/client.rs index eb3aec0..342c0eb 100644 --- a/szurubooru-client/src/client.rs +++ b/szurubooru-client/src/client.rs @@ -420,12 +420,14 @@ impl<'a> SzurubooruRequest<'a> { async fn handle_response(&self, response: Response) -> SzurubooruResult { if response.status().is_client_error() || response.status().is_server_error() { + let status = response.status(); let resp_json = response .text() .await .map_err(SzurubooruClientError::RequestError)?; + let server_error = serde_json::from_str::(&resp_json) - .map_err(|e| SzurubooruClientError::ResponseParsingError(e, resp_json))?; + .map_err(|e| SzurubooruClientError::ResponseError(status, resp_json))?; Err(SzurubooruClientError::SzurubooruServerError(server_error)) } else { Ok(response) @@ -436,7 +438,7 @@ impl<'a> SzurubooruRequest<'a> { &self, request: RequestBuilder, ) -> SzurubooruResult { - let request = request + let mut request = request .build() .map_err(SzurubooruClientError::RequestBuilderError)?; @@ -450,6 +452,7 @@ impl<'a> SzurubooruRequest<'a> { .text() .await .map_err(SzurubooruClientError::RequestError)?; + serde_json::from_str::>(&response_text) .map_err(|e| SzurubooruClientError::ResponseParsingError(e, response_text))? .into_result() @@ -599,7 +602,7 @@ impl<'a> SzurubooruRequest<'a> { /// Removes source tag and merges all of its usages, suggestions and implications to the /// target tag. Other tag properties such as category and aliases do not get transferred /// and are discarded. - pub async fn merge_tag(&self, merge_opts: &MergeTags) -> SzurubooruResult { + pub async fn merge_tags(&self, merge_opts: &MergeTags) -> SzurubooruResult { self.do_request(Method::POST, "/api/tag-merge", None, Some(merge_opts)) .await } @@ -639,6 +642,11 @@ impl<'a> SzurubooruRequest<'a> { method: Method, cupost: &CreateUpdatePost, ) -> SzurubooruResult { + if method == Method::POST && cupost.safety.is_none() { + return Err(SzurubooruClientError::ValidationError( + "Safety must be set".to_string(), + )); + } self.do_request(method, path, None, Some(cupost)).await } @@ -1009,7 +1017,12 @@ impl<'a> SzurubooruRequest<'a> { path: impl AsRef, ) -> SzurubooruResult<()> { let mut stream = self.get_image_bytestream(post_id).await?; - let mut file = File::open(path.as_ref()).map_err(SzurubooruClientError::IOError)?; + let mut file = File::options() + .write(true) + .truncate(true) + .create(true) + .open(path.as_ref()) + .map_err(SzurubooruClientError::IOError)?; self.write_content_to_file(&mut file, &mut stream).await } @@ -1067,7 +1080,7 @@ impl<'a> SzurubooruRequest<'a> { // Need to add a reverse search for bytes /// Searches for an exact match of a file based on the SHA1 checksum - pub async fn posts_for_file( + pub async fn post_for_file( &self, mut file: &mut File, ) -> SzurubooruResult> { @@ -1081,21 +1094,17 @@ impl<'a> SzurubooruRequest<'a> { .list_posts(Some(&vec![qt])) .await .map(|psr| self.propagate_urls(psr))?; - Ok(if psr.total > 1 { - Some(psr.results.swap_remove(0)) - } else { - None - }) + Ok(psr.results.first().map(|pr| pr.clone())) } /// Searches for an exact match of a file path based on the SHA1 checksum - pub async fn posts_for_file_path( + pub async fn post_for_file_path( &self, file_path: impl AsRef, ) -> SzurubooruResult> { let mut file = File::open(file_path).map_err(SzurubooruClientError::IOError)?; - self.posts_for_file(&mut file).await + self.post_for_file(&mut file).await } /// Retrieves information about an existing post. @@ -1137,6 +1146,11 @@ impl<'a> SzurubooruRequest<'a> { /// Updates score of authenticated user for given post. Valid scores are -1, 0 and 1. pub async fn rate_post(&self, post_id: u32, score: i8) -> SzurubooruResult { + if score < -1 || score > 1 { + return Err(SzurubooruClientError::ValidationError( + "Score must be -1, 0 or 1".to_string(), + )); + } let rating_obj = RateResource { score }; let path = format!("/api/post/{post_id}/score"); self.do_request(Method::PUT, &path, None, Some(&rating_obj)) @@ -1380,6 +1394,11 @@ impl<'a> SzurubooruRequest<'a> { comment_id: u32, score: i8, ) -> SzurubooruResult { + if score < -1 || score > 1 { + return Err(SzurubooruClientError::ValidationError( + "Score must be -1, 0 or 1".to_string(), + )); + } let path = format!("/api/comment/{comment_id}/score"); let rating = RateResource { score }; self.do_request(Method::PUT, &path, None, Some(&rating)) diff --git a/szurubooru-client/src/errors.rs b/szurubooru-client/src/errors.rs index 5de5323..32fc673 100644 --- a/szurubooru-client/src/errors.rs +++ b/szurubooru-client/src/errors.rs @@ -5,8 +5,11 @@ use crate::models::SzuruEither; use base64::EncodeSliceError; use derive_builder::UninitializedFieldError; #[cfg(feature = "python")] -use pyo3::{exceptions::PyRuntimeError, prelude::*}; +use pyo3::{create_exception, exceptions::PyException, prelude::*}; + +use reqwest::StatusCode; use serde::{Deserialize, Serialize}; +use strum_macros::AsRefStr; use thiserror::Error; use url::ParseError as UParseError; @@ -17,7 +20,7 @@ pub trait IntoClientResult { fn into_result(self) -> SzurubooruResult; } -#[derive(Debug, Error)] +#[derive(Debug, Error, AsRefStr)] /// Type that represents the various error states that can occur when interacting with /// Szurubooru pub enum SzurubooruClientError { @@ -38,6 +41,9 @@ pub enum SzurubooruClientError { /// Error occurred pas part of the request to the server #[error("Request error {0}")] RequestError(#[source] reqwest::Error), + /// Error response with a text response from the server + #[error("Response error {0}: Server reply: {1}")] + ResponseError(StatusCode, String), /// Error parsing the JSON response from the server #[error("Response Parsing error: {0}: {1}")] ResponseParsingError( @@ -51,8 +57,8 @@ pub enum SzurubooruClientError { #[error("JSON Serialization error: {0}")] JSONSerializationError(#[source] serde_json::Error), /// Error when validation fails for one of the Builder types - #[error("Builder validation error: {0}")] - BuilderValidationError(String), + #[error("Vlidation error: {0}")] + ValidationError(String), /// Error occurred when reading a file #[error("IO Error: {0}")] IOError(#[source] std::io::Error), @@ -69,14 +75,17 @@ impl From for SzurubooruClientError { impl From for SzurubooruClientError { fn from(value: UninitializedFieldError) -> Self { - SzurubooruClientError::BuilderValidationError(value.to_string()) + SzurubooruClientError::ValidationError(value.to_string()) } } +#[cfg(feature = "python")] +create_exception!(szurubooru_client, SzuruPyClientError, PyException); + #[cfg(feature = "python")] impl std::convert::From for PyErr { fn from(value: SzurubooruClientError) -> Self { - PyRuntimeError::new_err(value.to_string()) + SzuruPyClientError::new_err((value.as_ref().to_string(), value.to_string())) } } diff --git a/szurubooru-client/src/lib.rs b/szurubooru-client/src/lib.rs index 3cdc990..18660f6 100644 --- a/szurubooru-client/src/lib.rs +++ b/szurubooru-client/src/lib.rs @@ -22,7 +22,7 @@ //! ``` //! //! For all other methods for making the requests, see the documentation. - +#![feature(cfg_eval)] #![warn(missing_docs)] #![warn(rustdoc::missing_crate_level_docs)] @@ -34,9 +34,6 @@ pub use client::SzurubooruRequest; pub mod errors; pub use errors::SzurubooruResult; pub mod models; - -#[cfg(feature = "python")] -pub mod pyclient; pub mod tokens; #[cfg(feature = "python")] @@ -49,27 +46,26 @@ use pyo3::prelude::*; #[cfg_attr(feature = "python", pymodule)] /// A Python wrapper around [SzurubooruClient] mod szurubooru_client { - use pyo3::prelude::*; #[pymodule_export] pub use crate::{ + errors::SzuruPyClientError, models::{ AroundPostResult, CommentResource, GlobalInfo, ImageSearchResult, - ImageSearchSimilarPost, MergePool, MergePost, MergeTags, MicroPoolResource, - MicroPostResource, MicroTagResource, MicroUserResource, NoteResource, - PoolCategoryResource, PoolResource, PostResource, PostSafety, PostType, - SnapshotCreationDeletionData, SnapshotData, SnapshotModificationData, - SnapshotOperationType, SnapshotResource, SnapshotResourceType, TagCategoryResource, - TagResource, TagSibling, TemporaryFileUpload, UserAuthTokenResource, UserAvatarStyle, - UserRank, UserResource, + ImageSearchSimilarPost, MicroPoolResource, MicroPostResource, MicroTagResource, + MicroUserResource, NoteResource, PoolCategoryResource, PoolResource, PostResource, + PostSafety, PostType, SnapshotCreationDeletionData, SnapshotData, + SnapshotModificationData, SnapshotOperationType, SnapshotResource, + SnapshotResourceType, TagCategoryResource, TagResource, TagSibling, + UserAuthTokenResource, UserAvatarStyle, UserRank, UserResource, }, py::asynchronous::PythonAsyncClient, py::synchronous::PythonSyncClient, tokens::{ anonymous_token, named_token, sort_token, special_token, CommentNamedToken, CommentSortToken, PoolNamedToken, PoolSortToken, PostNamedToken, PostSortToken, - PostSpecialToken, SnapshotNamedToken, TagNamedToken, TagSortToken, UserNamedToken, - UserSortToken, + PostSpecialToken, QueryToken, SnapshotNamedToken, TagNamedToken, TagSortToken, + UserNamedToken, UserSortToken, }, }; } diff --git a/szurubooru-client/src/models.rs b/szurubooru-client/src/models.rs index 9c21597..528b079 100644 --- a/szurubooru-client/src/models.rs +++ b/szurubooru-client/src/models.rs @@ -266,14 +266,12 @@ pub struct CreateUpdateTagCategory { #[derive(Debug, Clone, Serialize, Deserialize, Builder)] #[builder(setter(strip_option), build_fn(error = "SzurubooruClientError"))] -#[cfg_attr(all(feature = "python"), pyclass(get_all))] #[serde(rename_all = "camelCase")] /// Removes source tag and merges all of its usages, suggestions and implications to the target tag. /// Other tag properties such as category and aliases do not get transferred and are discarded. pub struct MergeTags { /// Version of the tag to remove #[serde(rename = "removeVersion")] - #[cfg(feature = "python")] pub remove_tag_version: u32, /// The name of the tag to remove #[serde(rename = "remove")] @@ -514,6 +512,7 @@ pub struct CreateUpdatePost { #[serde(skip_serializing_if = "Option::is_none")] pub tags: Option>, /// Required field, represents the SFW/NSFW state of a post + #[builder(default)] #[serde(skip_serializing_if = "Option::is_none")] pub safety: Option, /// The origin of the post's content @@ -541,27 +540,21 @@ pub struct CreateUpdatePost { /// [upload_temporary_file](crate::SzurubooruRequest::upload_temporary_file) #[builder(default)] pub content_token: Option, + /// Upload the post anonymously + #[builder(default)] + #[serde(skip_serializing_if = "Option::is_none")] + pub anonymous: Option, } #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] -#[cfg_attr(all(feature = "python"), pyclass(get_all))] /// A token representing a temporary file upload pub struct TemporaryFileUpload { /// Temporary upload token pub token: String, } -#[cfg(feature = "python")] -#[cfg_attr(all(feature = "python"), pymethods)] -impl TemporaryFileUpload { - fn __repr__(&self) -> String { - format!("{:?}", self) - } -} - #[derive(Debug, Clone, Serialize, Deserialize, Builder)] -#[cfg_attr(all(feature = "python"), pyclass(get_all))] #[builder(build_fn(error = "SzurubooruClientError"))] #[serde(rename_all = "camelCase")] /// Removes source post and merges all of its tags, relations, scores, favorites and comments to @@ -647,61 +640,110 @@ pub enum UserAvatarStyle { #[serde(rename_all = "camelCase")] /// A single user pub struct UserResource { - #[cfg(feature = "python")] - #[pyo3(get)] /// Resource version. See [versioning](ResourceVersion) - pub version: Option, #[cfg(feature = "python")] #[pyo3(get)] + pub version: Option, + + /// Resource version. See [versioning](ResourceVersion) + #[cfg(not(feature = "python"))] + pub version: Option, + /// The user's username + #[cfg(feature = "python")] + #[pyo3(get)] pub name: Option, + + /// The user's username + #[cfg(not(feature = "python"))] + pub name: Option, + /// The user email. It is available only if the request is authenticated by the same user, /// or the authenticated user can change the email. If it's unavailable, the server returns /// `false`. If the user hasn't specified an email, the server returns [None](Option::None) pub email: Option>, + + /// The user rank, which effectively affects their privileges #[cfg(feature = "python")] #[pyo3(get)] - /// The user rank, which effectively affects their privileges pub rank: Option, + + /// The user rank, which effectively affects their privileges + #[cfg(not(feature = "python"))] + pub rank: Option, + + /// The last login time #[cfg(feature = "python")] #[pyo3(get)] #[serde(rename = "last-login-time")] - /// The last login time pub last_login_time: Option>, - #[cfg(feature = "python")] - #[pyo3(get)] - #[serde(rename = "creation-time")] - /// The user registration time - pub creation_time: Option>, - #[cfg(feature = "python")] - #[pyo3(get)] - /// How to render the user avatar - pub avatar_style: Option, - #[cfg(feature = "python")] - #[pyo3(get)] - /// The URL to the avatar - pub avatar_url: Option, + + /// The last login time #[cfg(not(feature = "python"))] - /// The URL to the avatar - pub avatar_url: Option, + #[serde(rename = "last-login-time")] + pub last_login_time: Option>, + + /// The user registration time + #[serde(rename = "creation-time")] #[cfg(feature = "python")] #[pyo3(get)] + pub creation_time: Option>, + + /// The user registration time + #[serde(rename = "creation-time")] + #[cfg(not(feature = "python"))] + pub creation_time: Option>, + + /// How to render the user avatar + #[cfg(feature = "python")] + #[pyo3(get)] + pub avatar_style: Option, + + /// How to render the user avatar + #[cfg(not(feature = "python"))] + pub avatar_style: Option, + + /// The URL to the avatar + #[cfg(feature = "python")] + #[pyo3(get)] + pub avatar_url: Option, + + /// The URL to the avatar + #[cfg(not(feature = "python"))] + pub avatar_url: Option, + /// Number of comments + #[cfg(feature = "python")] + #[pyo3(get)] #[serde(rename = "comment-count")] pub comment_count: Option, + + /// Number of comments + #[cfg(not(feature = "python"))] + #[serde(rename = "comment-count")] + pub comment_count: Option, + + /// Number of uploaded posts #[cfg(feature = "python")] #[pyo3(get)] - /// Number of uploaded posts #[serde(rename = "uploaded-post-count")] pub uploaded_post_count: Option, + + /// Number of uploaded posts + #[cfg(not(feature = "python"))] + #[serde(rename = "uploaded-post-count")] + pub uploaded_post_count: Option, + /// Number of liked posts. It is available only if the request is authenticated by the same /// user. If it's unavailable, the server returns `false` #[serde(rename = "liked-post-count")] pub liked_post_count: Option>, + /// Number of disliked posts. It is available only if the request is authenticated by the same /// user. If it's unavailable, the server returns `false`. #[serde(rename = "disliked-post-count")] pub disliked_post_count: Option>, + /// Number of favorited posts #[serde(rename = "favorite-post-count")] pub favorite_post_count: Option>, @@ -1108,7 +1150,6 @@ pub struct CreateUpdatePool { #[derive(Debug, Clone, Serialize, Deserialize, Builder, Default)] #[builder(build_fn(error = "SzurubooruClientError"))] -#[cfg_attr(all(feature = "python"), pyclass(get_all))] #[serde(rename_all = "camelCase")] /// This type is used to specify which pools should be merged. Uses the builder pattern like so: /// @@ -1116,7 +1157,7 @@ pub struct CreateUpdatePool { /// use szurubooru_client::models::MergePoolBuilder; /// // Merge pool ID 1 at version 1 to pool ID 3 at version 5 /// let merge_pool = MergePoolBuilder::default() -/// .remove_version(1) +/// .remove_pool_version(1) /// .remove(1) /// .merge_to_version(5) /// .merge_to(3) @@ -1125,7 +1166,7 @@ pub struct CreateUpdatePool { /// ``` pub struct MergePool { /// Version of the pool to remove. Must match the current Pool version - #[serde(rename = "removePool")] + #[serde(rename = "removeVersion")] pub remove_pool_version: u32, /// Pool ID to remove #[serde(rename = "remove")] @@ -1304,11 +1345,17 @@ impl WithBaseURL for SnapshotCreationDeletionData { #[serde(rename_all = "camelCase")] /// Data for a modified resource pub struct SnapshotModificationData { - /// The type of snapshot - #[serde(rename = "type")] #[cfg(feature = "python")] + #[serde(rename = "type")] #[pyo3(get)] + /// The type of snapshot pub snapshot_type: String, + + #[cfg(not(feature = "python"))] + #[serde(rename = "type")] + /// The type of snapshot + pub snapshot_type: String, + /// The JSON value for the modified resource. A dictionary diff that depends on the resource /// kind. /// diff --git a/szurubooru-client/src/py/asynchronous.rs b/szurubooru-client/src/py/asynchronous.rs index 0322f0d..f350fc5 100644 --- a/szurubooru-client/src/py/asynchronous.rs +++ b/szurubooru-client/src/py/asynchronous.rs @@ -16,6 +16,8 @@ pub struct PythonAsyncClient { impl PythonAsyncClient { #[new] #[pyo3(signature = (host, username=None, token=None, password=None, allow_insecure=None))] + /// + /// pub fn new( host: String, username: Option, @@ -57,6 +59,30 @@ impl PythonAsyncClient { .map_err(Into::into) } + #[pyo3(signature = (name, color=None, order=None, fields=None))] + pub async fn create_tag_category( + &self, + name: String, + color: Option, + order: Option, + fields: Option>, + ) -> PyResult { + let mut cutagcat = CreateUpdateTagCategoryBuilder::default(); + cutagcat.name(name); + if let Some(color) = color { + cutagcat.color(color); + } + if let Some(order) = order { + cutagcat.order(order); + } + let cutagcat = cutagcat.build()?; + self.client + .with_optional_fields(fields) + .create_tag_category(&cutagcat) + .await + .map_err(Into::into) + } + #[pyo3(signature = (name, version, color=None, order=None, fields=None))] pub async fn update_tag_category( &self, @@ -134,7 +160,8 @@ impl PythonAsyncClient { #[pyo3(signature = (names, category=None, description=None, implications=None, suggestions=None, fields=None))] pub async fn create_tag( &self, - names: Vec, + //names: Vec, + names: Py, category: Option, description: Option, implications: Option>, @@ -142,7 +169,19 @@ impl PythonAsyncClient { fields: Option>, ) -> PyResult { let mut cubuild = CreateUpdateTagBuilder::default(); - cubuild.names(names); + Python::with_gil(|py| { + if let Ok(name) = names.extract::(py) { + Ok(cubuild.names(vec![name])) + } else { + let list_res = names.extract::>(py); + if let Ok(names) = list_res { + Ok(cubuild.names(names)) + } else { + Err(list_res.err().unwrap()) + } + } + })?; + //cubuild.names(names); if let Some(cat) = category { cubuild.category(cat); } @@ -163,10 +202,11 @@ impl PythonAsyncClient { .map_err(Into::into) } - #[pyo3(signature = (name, names, category=None, description=None, implications=None, suggestions=None, fields=None))] + #[pyo3(signature = (name, version, names=None, category=None, description=None, implications=None, suggestions=None, fields=None))] pub async fn update_tag( &self, name: String, + version: u32, names: Option>, category: Option, description: Option, @@ -175,6 +215,7 @@ impl PythonAsyncClient { fields: Option>, ) -> PyResult { let mut cubuild = CreateUpdateTagBuilder::default(); + cubuild.version(version); if let Some(names) = names { cubuild.names(names); } @@ -220,7 +261,7 @@ impl PythonAsyncClient { } #[pyo3(signature = (remove_tag, remove_tag_version, merge_to_tag, merge_to_version, fields=None))] - pub async fn merge_tag( + pub async fn merge_tags( &self, remove_tag: String, remove_tag_version: u32, @@ -236,7 +277,7 @@ impl PythonAsyncClient { .build()?; self.client .with_optional_fields(fields) - .merge_tag(&mtags) + .merge_tags(&mtags) .await .map_err(Into::into) } @@ -269,7 +310,7 @@ impl PythonAsyncClient { } #[pyo3(signature = (url=None, token=None, file_path=None, thumbnail_path=None, tags=None, safety=None, source=None, - relations=None, notes=None, flags=None, fields=None))] + relations=None, notes=None, flags=None, anonymous=None, fields=None))] pub async fn create_post( &self, url: Option, @@ -282,6 +323,7 @@ impl PythonAsyncClient { relations: Option>, notes: Option>, flags: Option>, + anonymous: Option, fields: Option>, ) -> PyResult { let mut cupost = CreateUpdatePostBuilder::default(); @@ -303,6 +345,9 @@ impl PythonAsyncClient { if let Some(flags) = flags { cupost.flags(flags); } + if let Some(anonymous) = anonymous { + cupost.anonymous(anonymous); + } if let Some(token) = token { cupost.content_token(token); @@ -447,7 +492,7 @@ impl PythonAsyncClient { .map_err(Into::into) } - pub async fn reverse_search_image(&self, image_path: PathBuf) -> PyResult { + pub async fn reverse_image_search(&self, image_path: PathBuf) -> PyResult { self.client .request() .reverse_search_file_path(image_path) @@ -458,7 +503,7 @@ impl PythonAsyncClient { pub async fn post_for_image(&self, image_path: PathBuf) -> PyResult> { self.client .request() - .posts_for_file_path(image_path) + .post_for_file_path(image_path) .await .map_err(Into::into) } @@ -493,13 +538,14 @@ impl PythonAsyncClient { } #[pyo3(signature = (remove_post, remove_post_version, merge_to_post, - merge_to_version, fields=None))] + merge_to_version, replace_post_content=false, fields=None))] pub async fn merge_post( &self, remove_post: u32, remove_post_version: u32, merge_to_post: u32, merge_to_version: u32, + replace_post_content: bool, fields: Option>, ) -> PyResult { let mpost = MergePostBuilder::default() @@ -507,6 +553,7 @@ impl PythonAsyncClient { .remove_post(remove_post) .merge_to_version(merge_to_version) .merge_to_post(merge_to_post) + .replace_post_content(replace_post_content) .build()?; self.client .with_optional_fields(fields) @@ -697,14 +744,26 @@ impl PythonAsyncClient { #[pyo3(signature = (names, category=None, description=None, posts=None, fields=None))] pub async fn create_pool<'py>( &self, - names: Vec, + names: Py, category: Option, description: Option, posts: Option>, fields: Option>, ) -> PyResult { let mut cupool = CreateUpdatePoolBuilder::default(); - cupool.names(names); + Python::with_gil(|py| { + if let Ok(name) = names.extract::(py) { + Ok(cupool.names(vec![name])) + } else { + let list_res = names.extract::>(py); + if let Ok(names) = list_res { + Ok(cupool.names(names)) + } else { + Err(list_res.err().unwrap()) + } + } + })?; + //cupool.names(names); if let Some(cat) = category { cupool.category(cat); } @@ -885,15 +944,11 @@ impl PythonAsyncClient { rating: i8, fields: Option>, ) -> PyResult { - if rating < -1 || rating > 1 { - Err(PyValueError::new_err("Rating must be -1, 0, or 1")) - } else { - self.client - .with_optional_fields(fields) - .rate_comment(comment_id, rating) - .await - .map_err(Into::into) - } + self.client + .with_optional_fields(fields) + .rate_comment(comment_id, rating) + .await + .map_err(Into::into) } #[pyo3(signature = (query=None, fields=None, limit=None, offset=None))] @@ -1024,11 +1079,12 @@ impl PythonAsyncClient { .map(|ur| ur.results) } - #[pyo3(signature = (user_name, note=None, expiration_time=None, fields=None))] + #[pyo3(signature = (user_name, note=None, enabled=None, expiration_time=None, fields=None))] pub async fn create_user_token( &self, user_name: String, note: Option, + enabled: Option, expiration_time: Option>, fields: Option>, ) -> PyResult { @@ -1039,6 +1095,9 @@ impl PythonAsyncClient { if let Some(etime) = expiration_time { cutoken.expiration_time(etime); } + if let Some(enabled) = enabled { + cutoken.enabled(enabled); + } let cutoken = cutoken.build()?; self.client .with_optional_fields(fields) @@ -1137,11 +1196,12 @@ impl PythonAsyncClient { .map_err(Into::into) } - pub async fn upload_temporary_file(&self, file_path: PathBuf) -> PyResult { + pub async fn upload_temporary_file(&self, file_path: PathBuf) -> PyResult { self.client .request() .upload_temporary_file_from_path(file_path) .await .map_err(Into::into) + .map(|t| t.token) } } diff --git a/szurubooru-client/src/py/mod.rs b/szurubooru-client/src/py/mod.rs index dfdc2b4..e77a80d 100644 --- a/szurubooru-client/src/py/mod.rs +++ b/szurubooru-client/src/py/mod.rs @@ -1,6 +1,8 @@ use crate::models::PagedSearchResult; +use pyo3::exceptions::PyException; use pyo3::prelude::*; use pyo3::types::PyList; +use pyo3::types::PyListMethods; pub mod asynchronous; pub mod synchronous; @@ -20,6 +22,12 @@ impl PyPagedSearchResult { fn __repr__(&self) -> String { format!("{:?}", self) } + + /*fn __len__(&self) -> PyResult { + Python::with_gil(|py| { + Ok(self.results.bind_borrowed(py).len()) + }) + }*/ } impl> From> for PyPagedSearchResult { diff --git a/szurubooru-client/src/py/synchronous.rs b/szurubooru-client/src/py/synchronous.rs index 8ebca8d..aee0641 100644 --- a/szurubooru-client/src/py/synchronous.rs +++ b/szurubooru-client/src/py/synchronous.rs @@ -10,6 +10,20 @@ use std::path::{Path, PathBuf}; use tokio::runtime::{Builder, Runtime}; #[pyclass(name = "SzurubooruSyncClient")] +/// Constructor for the SzurubooruSyncClient +/// This client is completely synchronous. For the `asyncio` compatible version, +/// see [szurubooru_client.PythonAsyncClient](SzurubooruAsyncClient) +/// +/// ## Arguments +/// * `host`: Base host URL for the Szurubooru instance. Should be the protocol, hostname and any port +/// E.g `http://localhost:9801` +/// * `username`: The username used to authenticate against the Szurubooru instance. Leave blank for +/// anonymous authentication +/// * `password`: The password to use for `Basic` authentication. Token authentication should +/// be preferred +/// * `token`: The token to use for `Bearer` authentication. +/// * `allow_insecure`: Disable cert validation. Disables SSL authentication +/// pub struct PythonSyncClient { client: PythonAsyncClient, runtime: Runtime, @@ -19,6 +33,7 @@ pub struct PythonSyncClient { impl PythonSyncClient { #[new] #[pyo3(signature = (host, username=None, token=None, password=None, allow_insecure=None))] + /// This method is for creating new instances of the SzurubooruSyncClient pub fn new( host: String, username: Option, @@ -32,6 +47,7 @@ impl PythonSyncClient { } #[pyo3(signature = (fields=None))] + /// List the available tag categories pub fn list_tag_categories( &self, fields: Option>, @@ -40,6 +56,18 @@ impl PythonSyncClient { .block_on(self.client.list_tag_categories(fields)) } + #[pyo3(signature = (name, color=None, order=None, fields=None))] + pub fn create_tag_category( + &self, + name: String, + color: Option, + order: Option, + fields: Option>, + ) -> PyResult { + self.runtime + .block_on(self.client.create_tag_category(name, color, order, fields)) + } + #[pyo3(signature = (name, version, color=None, order=None, fields=None))] pub fn update_tag_category( &self, @@ -90,7 +118,8 @@ impl PythonSyncClient { #[pyo3(signature = (names, category=None, description=None, implications=None, suggestions=None, fields=None))] pub fn create_tag( &self, - names: Vec, + //names: Vec, + names: Py, category: Option, description: Option, implications: Option>, @@ -107,10 +136,12 @@ impl PythonSyncClient { )) } - #[pyo3(signature = (name, names, category=None, description=None, implications=None, suggestions=None, fields=None))] + #[pyo3(signature = (name, version, names=None, category=None, description=None, + implications=None, suggestions=None, fields=None))] pub fn update_tag( &self, name: String, + version: u32, names: Option>, category: Option, description: Option, @@ -120,6 +151,7 @@ impl PythonSyncClient { ) -> PyResult { self.runtime.block_on(self.client.update_tag( name, + version, names, category, description, @@ -139,7 +171,7 @@ impl PythonSyncClient { } #[pyo3(signature = (remove_tag, remove_tag_version, merge_to_tag, merge_to_version, fields=None))] - pub fn merge_tag( + pub fn merge_tags( &self, remove_tag: String, remove_tag_version: u32, @@ -147,7 +179,7 @@ impl PythonSyncClient { merge_to_version: u32, fields: Option>, ) -> PyResult { - self.runtime.block_on(self.client.merge_tag( + self.runtime.block_on(self.client.merge_tags( remove_tag, remove_tag_version, merge_to_tag, @@ -173,7 +205,7 @@ impl PythonSyncClient { } #[pyo3(signature = (url=None, token=None, file_path=None, thumbnail_path=None, tags=None, safety=None, source=None, - relations=None, notes=None, flags=None, fields=None))] + relations=None, notes=None, flags=None, anonymous=None, fields=None))] pub fn create_post( &self, url: Option, @@ -186,6 +218,7 @@ impl PythonSyncClient { relations: Option>, notes: Option>, flags: Option>, + anonymous: Option, fields: Option>, ) -> PyResult { self.runtime.block_on(self.client.create_post( @@ -199,6 +232,7 @@ impl PythonSyncClient { relations, notes, flags, + anonymous, fields, )) } @@ -258,9 +292,9 @@ impl PythonSyncClient { .block_on(self.client.download_thumbnail_to_path(post_id, file_path)) } - pub fn reverse_search_image(&self, image_path: PathBuf) -> PyResult { + pub fn reverse_image_search(&self, image_path: PathBuf) -> PyResult { self.runtime - .block_on(self.client.reverse_search_image(image_path)) + .block_on(self.client.reverse_image_search(image_path)) } pub fn post_for_image(&self, image_path: PathBuf) -> PyResult> { @@ -283,13 +317,14 @@ impl PythonSyncClient { } #[pyo3(signature = (remove_post, remove_post_version, merge_to_post, - merge_to_version, fields=None))] + merge_to_version, replace_post_content=false, fields=None))] pub fn merge_post( &self, remove_post: u32, remove_post_version: u32, merge_to_post: u32, merge_to_version: u32, + replace_post_content: bool, fields: Option>, ) -> PyResult { self.runtime.block_on(self.client.merge_post( @@ -297,6 +332,7 @@ impl PythonSyncClient { remove_post_version, merge_to_post, merge_to_version, + replace_post_content, fields, )) } @@ -422,7 +458,7 @@ impl PythonSyncClient { #[pyo3(signature = (names, category=None, description=None, posts=None, fields=None))] pub fn create_pool( &self, - names: Vec, + names: Py, category: Option, description: Option, posts: Option>, @@ -622,17 +658,19 @@ impl PythonSyncClient { .block_on(self.client.list_user_tokens(user_name, fields)) } - #[pyo3(signature = (user_name, note=None, expiration_time=None, fields=None))] + #[pyo3(signature = (user_name, note=None, enabled=None, expiration_time=None, fields=None))] pub fn create_user_token( &self, user_name: String, note: Option, + enabled: Option, expiration_time: Option>, fields: Option>, ) -> PyResult { self.runtime.block_on(self.client.create_user_token( user_name, note, + enabled, expiration_time, fields, )) @@ -702,7 +740,7 @@ impl PythonSyncClient { self.runtime.block_on(self.client.global_info()) } - pub fn upload_temporary_file(&self, file_path: PathBuf) -> PyResult { + pub fn upload_temporary_file(&self, file_path: PathBuf) -> PyResult { self.runtime .block_on(self.client.upload_temporary_file(file_path)) } diff --git a/szurubooru-client/src/pyclient.rs b/szurubooru-client/src/pyclient.rs deleted file mode 100644 index 8b13789..0000000 --- a/szurubooru-client/src/pyclient.rs +++ /dev/null @@ -1 +0,0 @@ - diff --git a/szurubooru-client/src/tokens.rs b/szurubooru-client/src/tokens.rs index 921b365..c14a49b 100644 --- a/szurubooru-client/src/tokens.rs +++ b/szurubooru-client/src/tokens.rs @@ -157,7 +157,7 @@ pub fn sort_token(key: &Bound<'_, PyAny>) -> PyResult { #[cfg(feature = "python")] #[cfg_attr(all(feature = "python"), pyfunction)] -pub fn anonymous_token(key: &Bound<'_, PyString>) -> PyResult { +pub fn anonymous_token(key: &Bound<'_, PyAny>) -> PyResult { QueryToken::anonymous_py(key) } @@ -183,7 +183,11 @@ impl QueryToken { #[pyo3(name = "token")] #[staticmethod] pub fn token_py(key: &Bound<'_, PyAny>, value: &Bound<'_, PyAny>) -> PyResult { - let value = value.extract::()?; + let value = if let Ok(value) = value.extract::() { + value.to_string() + } else { + value.extract::()? + }; if let Ok(tnt) = key.extract::() { Ok(QueryToken::token(tnt, value)) @@ -226,7 +230,7 @@ impl QueryToken { #[pyo3(name = "anonymous")] #[staticmethod] - pub fn anonymous_py(key: &Bound<'_, PyString>) -> PyResult { + pub fn anonymous_py(key: &Bound<'_, PyAny>) -> PyResult { let key = key.extract::()?; Ok(QueryToken::anonymous(key)) }