Underground-Digital/Workflow-Engine
0
1from collections.abc import Generator2 3import google.generativeai.types.generation_types as generation_config_types4import pytest5from _pytest.monkeypatch import MonkeyPatch6from google.ai import generativelanguage as glm7from google.ai.generativelanguage_v1beta.types import content as gag_content8from google.generativeai import GenerativeModel9from google.generativeai.client import _ClientManager, configure10from google.generativeai.types import GenerateContentResponse, content_types, safety_types11from google.generativeai.types.generation_types import BaseGenerateContentResponse12 13current_api_key = ""14 15 16class MockGoogleResponseClass:17 _done = False18 19 def __iter__(self):20 full_response_text = "it's google!"21 22 for i in range(0, len(full_response_text) + 1, 1):23 if i == len(full_response_text):24 self._done = True25 yield GenerateContentResponse(26 done=True, iterator=None, result=glm.GenerateContentResponse({}), chunks=[]27 )28 else:29 yield GenerateContentResponse(30 done=False, iterator=None, result=glm.GenerateContentResponse({}), chunks=[]31 )32 33 34class MockGoogleResponseCandidateClass:35 finish_reason = "stop"36 37 @property38 def content(self) -> gag_content.Content:39 return gag_content.Content(parts=[gag_content.Part(text="it's google!")])40 41 42class MockGoogleClass:43 @staticmethod44 def generate_content_sync() -> GenerateContentResponse:45 return GenerateContentResponse(done=True, iterator=None, result=glm.GenerateContentResponse({}), chunks=[])46 47 @staticmethod48 def generate_content_stream() -> Generator[GenerateContentResponse, None, None]:49 return MockGoogleResponseClass()50 51 def generate_content(52 self: GenerativeModel,53 contents: content_types.ContentsType,54 *,55 generation_config: generation_config_types.GenerationConfigType | None = None,56 safety_settings: safety_types.SafetySettingOptions | None = None,57 stream: bool = False,58 **kwargs,59 ) -> GenerateContentResponse:60 global current_api_key61 62 if len(current_api_key) < 16:63 raise Exception("Invalid API key")64 65 if stream:66 return MockGoogleClass.generate_content_stream()67 68 return MockGoogleClass.generate_content_sync()69 70 @property71 def generative_response_text(self) -> str:72 return "it's google!"73 74 @property75 def generative_response_candidates(self) -> list[MockGoogleResponseCandidateClass]:76 return [MockGoogleResponseCandidateClass()]77 78 def make_client(self: _ClientManager, name: str):79 global current_api_key80 81 if name.endswith("_async"):82 name = name.split("_")[0]83 cls = getattr(glm, name.title() + "ServiceAsyncClient")84 else:85 cls = getattr(glm, name.title() + "ServiceClient")86 87 # Attempt to configure using defaults.88 if not self.client_config:89 configure()90 91 client_options = self.client_config.get("client_options", None)92 if client_options:93 current_api_key = client_options.api_key94 95 def nop(self, *args, **kwargs):96 pass97 98 original_init = cls.__init__99 cls.__init__ = nop100 client: glm.GenerativeServiceClient = cls(**self.client_config)101 cls.__init__ = original_init102 103 if not self.default_metadata:104 return client105 106 107@pytest.fixture108def setup_google_mock(request, monkeypatch: MonkeyPatch):109 monkeypatch.setattr(BaseGenerateContentResponse, "text", MockGoogleClass.generative_response_text)110 monkeypatch.setattr(BaseGenerateContentResponse, "candidates", MockGoogleClass.generative_response_candidates)111 monkeypatch.setattr(GenerativeModel, "generate_content", MockGoogleClass.generate_content)112 monkeypatch.setattr(_ClientManager, "make_client", MockGoogleClass.make_client)113 114 yield115 116 monkeypatch.undo()117 