From 1376c9559570bc4ce02934fb7a1b3510a95c8f6e Mon Sep 17 00:00:00 2001 From: TimmekHW Date: Mon, 7 Jul 2025 21:09:16 +0300 Subject: [PATCH] Updated /suggest-tags to the current state --- .../_response/ai/generate_image.py | 2 +- .../sdk/ai/generate_image/__init__.py | 1 + .../sdk/ai/generate_image/suggest_tags.py | 124 ++++++++++-------- tests/test_suggest_tags.py | 31 +++++ 4 files changed, 99 insertions(+), 59 deletions(-) create mode 100644 tests/test_suggest_tags.py diff --git a/src/novelai_python/_response/ai/generate_image.py b/src/novelai_python/_response/ai/generate_image.py index 20f21b9..b8d582f 100755 --- a/src/novelai_python/_response/ai/generate_image.py +++ b/src/novelai_python/_response/ai/generate_image.py @@ -1,7 +1,7 @@ # -*- coding: utf-8 -*- # @Time : 2024/1/26 上午11:16 # @Author : sudoskys -# @File : text2image.py +# @File : generate_image.py from typing import Tuple, List diff --git a/src/novelai_python/sdk/ai/generate_image/__init__.py b/src/novelai_python/sdk/ai/generate_image/__init__.py index b9a18d7..8eff6fb 100755 --- a/src/novelai_python/sdk/ai/generate_image/__init__.py +++ b/src/novelai_python/sdk/ai/generate_image/__init__.py @@ -26,6 +26,7 @@ get_model_group, ModelGroups, get_supported_params, get_modifiers from .params import Params, get_default_params from .schema import Character, V4Prompt, V4NegativePrompt, PositionMap +from .suggest_tags import SuggestTags from ...schema import ApiBaseModel from ...._exceptions import APIError, AuthError, ConcurrentGenerationError, SessionHttpError, DataSerializationError from ...._response.ai.generate_image import ImageGenerateResp, RequestParams diff --git a/src/novelai_python/sdk/ai/generate_image/suggest_tags.py b/src/novelai_python/sdk/ai/generate_image/suggest_tags.py index ad8cf4d..7e2ab20 100755 --- a/src/novelai_python/sdk/ai/generate_image/suggest_tags.py +++ b/src/novelai_python/sdk/ai/generate_image/suggest_tags.py @@ -1,29 +1,28 @@ # -*- coding: utf-8 -*- -# @Time : 2024/2/13 下午8:09 -# @Author : sudoskys -# @File : suggest-tags.py +# @Time : 2025/7/7 下午20:11 +# @Author : TimmekHW +# @File : suggest_tags.py -from typing import Optional, Union -from urllib.parse import urlparse +from typing import Optional, Union, Literal import curl_cffi import httpx from curl_cffi.requests import AsyncSession from loguru import logger -from pydantic import PrivateAttr +from pydantic import PrivateAttr, Field from novelai_python.sdk.ai._enum import Model from ...schema import ApiBaseModel from ...._exceptions import APIError, AuthError, SessionHttpError from ...._response.ai.generate_image import SuggestTagsResp from ....credential import CredentialBase -from ....utils import try_jsonfy class SuggestTags(ApiBaseModel): - _endpoint: str = PrivateAttr("https://api.novelai.net") - model: Model = Model.NAI_DIFFUSION_3 - prompt: str = "landscape" + _endpoint: str = PrivateAttr("https://image.novelai.net") + model: Union[Model, str] = Field(default=Model.NAI_DIFFUSION_4_5_FULL, description="The image model") + prompt: str = Field(..., description="The incomplete tag query") + lang: Literal["en", "jp"] = Field(default="en", description="The language of the tag query") @property def endpoint(self): @@ -38,57 +37,66 @@ def base_url(self): return f"{self.endpoint.strip('/')}/ai/generate-image/suggest-tags" async def request(self, - session: Union[AsyncSession, CredentialBase], + session: Optional[Union[AsyncSession, CredentialBase]] = None, *, override_headers: Optional[dict] = None ) -> SuggestTagsResp: """ - Request to get user subscription information - :param override_headers: - :param session: - :return: + Request to get tag suggestions for image generation + + :param session: Session object for making requests + :param override_headers: Optional headers to override defaults + :return: SuggestTagsResp containing tag suggestions + :raises AuthError: If authentication fails + :raises APIError: If API returns an error + :raises SessionHttpError: If session/network error occurs """ - # Data Build + # Prepare request data request_data = self.model_dump(mode="json", exclude_none=True) - if isinstance(session, AsyncSession): - pass - elif isinstance(session, CredentialBase): - pass - # Header - if override_headers: - session.headers.clear() - session.headers.update(override_headers) - try: - assert hasattr(session, "get"), "session must have get method." - response = await session.get( - url=self.base_url + "?" + "&".join([f"{k}={v}" for k, v in request_data.items()]) - ) - if ( - "application/json" not in response.headers.get('Content-Type') - or response.status_code != 200 - ): - error_message = await self.handle_error_response(response=response, request_data=request_data) - status_code = error_message.get("statusCode", response.status_code) - message = error_message.get("message", "Unknown error") - if status_code in [400, 401, 402]: - # 400 : validation error - # 401 : unauthorized - # 402 : payment required - # 409 : conflict - raise AuthError(message, request=request_data, code=status_code, response=error_message) - if status_code in [500]: - # An unknown error occured. - raise APIError(message, request=request_data, code=status_code, response=error_message) - raise APIError(message, request=request_data, code=status_code, response=error_message) - return SuggestTagsResp.model_validate(response.json()) - except curl_cffi.requests.errors.RequestsError as exc: - logger.exception(exc) - raise SessionHttpError("An AsyncSession RequestsError occurred, maybe SSL error. Try again later!") - except httpx.HTTPError as exc: - logger.exception(exc) - raise SessionHttpError("An HTTPError occurred, maybe SSL error. Try again later!") - except APIError as e: - raise e - except Exception as e: - logger.opt(exception=e).exception("An Unexpected error occurred") - raise e + + # Create session if not provided + if session is None: + from curl_cffi.requests import AsyncSession + async with AsyncSession() as sess: + return await self._make_request(sess, request_data, override_headers) + else: + async with session if isinstance(session, AsyncSession) else await session.get_session() as sess: + return await self._make_request(sess, request_data, override_headers) + + async def _make_request(self, sess, request_data, override_headers): + if override_headers: + sess.headers.clear() + sess.headers.update(override_headers) + + # Build query string + query_params = "&".join([f"{k}={v}" for k, v in request_data.items()]) + + try: + self.ensure_session_has_get_method(sess) + response = await sess.get(f"{self.base_url}?{query_params}") + + if response.status_code != 200: + error_message = await self.handle_error_response(response, request_data) + status_code = error_message.get("statusCode", response.status_code) + message = error_message.get("message", "Unknown error") + + if status_code in [400, 401, 402]: + raise AuthError(message, request=request_data, code=status_code, response=error_message) + elif status_code == 500: + raise APIError(message, request=request_data, code=status_code, response=error_message) + else: + raise APIError(message, request=request_data, code=status_code, response=error_message) + + return SuggestTagsResp.model_validate(response.json()) + + except curl_cffi.requests.errors.RequestsError as exc: + logger.exception(exc) + raise SessionHttpError("A RequestsError occurred (e.g., SSL error). Try again later.") + except httpx.HTTPError as exc: + logger.exception(exc) + raise SessionHttpError("An HTTP error occurred. Try again later.") + except APIError as e: + raise e + except Exception as e: + logger.opt(exception=e).exception("Unexpected error occurred during the request.") + raise Exception("An unexpected error occurred.") from e \ No newline at end of file diff --git a/tests/test_suggest_tags.py b/tests/test_suggest_tags.py new file mode 100644 index 0000000..6906806 --- /dev/null +++ b/tests/test_suggest_tags.py @@ -0,0 +1,31 @@ +# -*- coding: utf-8 -*- +# @Time : 2025/01/07 +# @File : test_suggest_tags.py + +import asyncio +from novelai_python.sdk.ai.generate_image import SuggestTags +from novelai_python.sdk.ai._enum import Model + + +def test_suggest_tags_default_model(): + async def run(): + suggest = SuggestTags(prompt="senko") + result = await suggest.request() + assert suggest.model == Model.NAI_DIFFUSION_4_5_FULL + assert suggest.lang == "en" + assert len(result.tags) > 0 + assert all(hasattr(tag, 'tag') for tag in result.tags) + + asyncio.run(run()) + + +def test_suggest_tags_with_nai_diffusion_3(): + async def run(): + suggest = SuggestTags( + prompt="senko", + model=Model.NAI_DIFFUSION_3 + ) + result = await suggest.request() + assert len(result.tags) > 0 + + asyncio.run(run()) \ No newline at end of file