diff --git a/.github/workflows/lint_test.yml b/.github/workflows/lint_test.yml index 1df49dd..56c115c 100644 --- a/.github/workflows/lint_test.yml +++ b/.github/workflows/lint_test.yml @@ -1,4 +1,4 @@ -name: Run linters and unittests +name: Lint and Test on: push: @@ -30,19 +30,18 @@ jobs: - name: Analysing the code with pylint run: | pylint \ - *.py \ + ./sgpt \ + ./tests \ + --disable=fixme \ + --disable=cyclic-import \ + --disable=useless-import-alias \ --disable=missing-function-docstring \ - --disable=too-many-arguments \ --disable=missing-module-docstring \ - --disable=import-error \ --disable=missing-class-docstring \ - --disable=too-many-instance-attributes \ --disable=too-many-function-args \ - --disable=unspecified-encoding \ - --max-line-length=120 + - name: Analysing the code with black - run: | - black --target-version py310 -l 120 *.py + run: black ./sgpt ./tests --check --target-version py310 - name: Run unittests run: | export OPENAI_API_KEY=test_api_key diff --git a/README.md b/README.md index bdca1e5..1355ef3 100644 --- a/README.md +++ b/README.md @@ -7,7 +7,7 @@ A command-line productivity tool powered by OpenAI's ChatGPT (GPT-3.5). As devel ## Installation ```shell -pip install shell-gpt +pip install shell-gpt==0.8.2 ``` You'll need an OpenAI API key, you can generate one [here](https://beta.openai.com/account/api-keys). @@ -211,21 +211,25 @@ REQUEST_TIMEOUT=60 ### Full list of arguments ```shell -╭─ Arguments ───────────────────────────────────────────────────────────────────────────────────────────────╮ -│ prompt [PROMPT] The prompt to generate completions for. │ -╰───────────────────────────────────────────────────────────────────────────────────────────────────────────╯ -╭─ Options ─────────────────────────────────────────────────────────────────────────────────────────────────╮ -│ --temperature FLOAT RANGE [0.0<=x<=1.0] Randomness of generated output. [default: 1.0] │ -│ --top-probability FLOAT RANGE [0.1<=x<=1.0] Limits highest probable tokens (words). [default: 1.0] │ -│ --chat TEXT Follow conversation with id (chat mode). [default: None] │ -│ --show-chat TEXT Show all messages from provided chat id. [default: None] │ -│ --list-chat List all existing chat ids. [default: no-list-chat] │ -│ --shell Generate and execute shell command. │ -│ --code Provide code as output. [default: no-code] │ -│ --editor Open $EDITOR to provide a prompt. [default: no-editor] │ -│ --cache Cache completion results. [default: cache] │ -│ --help Show this message and exit. │ -╰───────────────────────────────────────────────────────────────────────────────────────────────────────────╯ +╭─ Arguments ─────────────────────────────────────────────────────────────────────────────────────────────╮ +│ prompt [PROMPT] The prompt to generate completions for. │ +╰─────────────────────────────────────────────────────────────────────────────────────────────────────────╯ +╭─ Options ───────────────────────────────────────────────────────────────────────────────────────────────╮ +│ --temperature FLOAT RANGE [0.0<=x<=1.0] Randomness of generated output. [default: 0.1] │ +│ --top-probability FLOAT RANGE [0.1<=x<=1.0] Limits highest probable tokens (words). [default: 1.0] │ +│ --editor Open $EDITOR to provide a prompt. [default: no-editor] │ +│ --cache Cache completion results. [default: cache] │ +│ --help Show this message and exit. │ +╰─────────────────────────────────────────────────────────────────────────────────────────────────────────╯ +╭─ Chat Options ──────────────────────────────────────────────────────────────────────────────────────────╮ +│ --chat TEXT Follow conversation with id (chat mode). [default: None] │ +│ --show-chat TEXT Show all messages from provided chat id. [default: None] │ +│ --list-chat List all existing chat ids. [default: no-list-chat] │ +╰─────────────────────────────────────────────────────────────────────────────────────────────────────────╯ +╭─ Assistance Options ────────────────────────────────────────────────────────────────────────────────────╮ +│ --shell -s Generate and execute shell commands. │ +│ --code Generate only code. [default: no-code] │ +╰─────────────────────────────────────────────────────────────────────────────────────────────────────────╯ ``` ## Docker diff --git a/pre-push.sh b/pre-push.sh index cf8c93d..ad882fb 100755 --- a/pre-push.sh +++ b/pre-push.sh @@ -21,26 +21,25 @@ if [ -f "$PWD/venv/bin/activate" ]; then source $PWD/venv/bin/activate fi -if black --exclude venv --check --target-version py310 -l 120 *.py +if black ./sgpt ./tests --check --target-version py310 then echo 'Black passed ✅' else echo 'Black failed ❌' - echo 'RUN: black --exclude venv --target-version py310 -l 120 *.py' + echo 'RUN: black ./sgpt ./tests --target-version py310' exit 1 fi if pylint \ - *.py \ + ./sgpt \ + ./tests \ + --disable=fixme \ + --disable=cyclic-import \ + --disable=useless-import-alias \ --disable=missing-function-docstring \ - --disable=too-many-arguments \ --disable=missing-module-docstring \ - --disable=import-error \ --disable=missing-class-docstring \ - --disable=too-many-instance-attributes \ --disable=too-many-function-args \ - --disable=unspecified-encoding \ - --max-line-length=120 \ --ignore=venv then echo 'Pylint passed ✅' diff --git a/setup.py b/setup.py index 827b606..dfca28d 100644 --- a/setup.py +++ b/setup.py @@ -3,7 +3,7 @@ from setuptools import setup, find_packages # pylint: disable=consider-using-with setup( name="shell_gpt", - version="0.8.1", + version="0.8.2", packages=find_packages(), install_requires=[ "typer~=0.7.0", diff --git a/sgpt/__init__.py b/sgpt/__init__.py index 9cd8690..7fc6674 100644 --- a/sgpt/__init__.py +++ b/sgpt/__init__.py @@ -1,7 +1,8 @@ from . import config as config from .cache import Cache as Cache -from .cache import ChatCache as ChatCache from .client import OpenAIClient as OpenAIClient +from .handlers.chat_handler import ChatHandler as ChatHandler +from .handlers.default_handler import DefaultHandler as DefaultHandler from . import utils as utils from .app import main as main from .app import entry_point as cli diff --git a/sgpt/app.py b/sgpt/app.py index c204bb8..71b9c2d 100644 --- a/sgpt/app.py +++ b/sgpt/app.py @@ -12,95 +12,101 @@ API Key is stored locally for easy use in future runs. import os -from typing import Mapping, List import typer # Click is part of typer. -from click import MissingParameter, BadParameter -from sgpt import config, make_prompt, OpenAIClient -from sgpt.utils import ( - echo_chat_ids, - echo_chat_messages, - get_edited_prompt, -) +from click import MissingParameter, BadArgumentUsage +from sgpt import config, OpenAIClient +from sgpt import ChatHandler, DefaultHandler +from sgpt.utils import get_edited_prompt -def get_completion( - messages: List[Mapping[str, str]], - temperature: float, - top_p: float, - caching: bool, - chat: str, -): - api_host = config.get("OPENAI_API_HOST") - api_key = config.get("OPENAI_API_KEY") - client = OpenAIClient(api_host, api_key) - return client.get_completion( - messages=messages, - model="gpt-3.5-turbo", - temperature=temperature, - top_probability=top_p, - caching=caching, - chat_id=chat, - ) - - -def main( - prompt: str = typer.Argument(None, show_default=False, help="The prompt to generate completions for."), - temperature: float = typer.Option(1.0, min=0.0, max=1.0, help="Randomness of generated output."), - top_probability: float = typer.Option(1.0, min=0.1, max=1.0, help="Limits highest probable tokens (words)."), - chat: str = typer.Option(None, help="Follow conversation with id (chat mode)."), - show_chat: str = typer.Option(None, help="Show all messages from provided chat id."), - list_chat: bool = typer.Option(False, help="List all existing chat ids."), - shell: bool = typer.Option(False, "--shell", "-s", help="Generate and execute shell command."), - code: bool = typer.Option(False, help="Provide code as output."), - editor: bool = typer.Option(False, help="Open $EDITOR to provide a prompt."), - cache: bool = typer.Option(True, help="Cache completion results."), +def main( # pylint: disable=too-many-arguments + prompt: str = typer.Argument( + None, + show_default=False, + help="The prompt to generate completions for.", + ), + temperature: float = typer.Option( + 0.1, + min=0.0, + max=1.0, + help="Randomness of generated output.", + ), + top_probability: float = typer.Option( + 1.0, + min=0.1, + max=1.0, + help="Limits highest probable tokens (words).", + ), + chat: str = typer.Option( + None, + help="Follow conversation with id (chat mode).", + rich_help_panel="Chat Options", + ), + show_chat: str = typer.Option( # pylint: disable=W0613 + None, + help="Show all messages from provided chat id.", + callback=ChatHandler.show_messages, + rich_help_panel="Chat Options", + ), + list_chat: bool = typer.Option( # pylint: disable=W0613 + False, + help="List all existing chat ids.", + callback=ChatHandler.list_ids, + rich_help_panel="Chat Options", + ), + shell: bool = typer.Option( + False, + "--shell", + "-s", + help="Generate and execute shell commands.", + rich_help_panel="Assistance Options", + ), + code: bool = typer.Option( + False, + help="Generate only code.", + rich_help_panel="Assistance Options", + ), + editor: bool = typer.Option( + False, + help="Open $EDITOR to provide a prompt.", + ), + cache: bool = typer.Option( + True, + help="Cache completion results.", + ), ) -> None: - if list_chat: - echo_chat_ids() - return - if show_chat: - echo_chat_messages(show_chat) - return - if not prompt and not editor: raise MissingParameter(param_hint="PROMPT", param_type="string") + if shell and code: + raise BadArgumentUsage("--shell and --code options cannot be used together.") + if editor: prompt = get_edited_prompt() - if chat and OpenAIClient.chat_cache.exists(chat): - chat_history = OpenAIClient.chat_cache.get_messages(chat) - is_shell_chat = chat_history[0].endswith("###\nCommand:") - is_code_chat = chat_history[0].endswith("###\nCode:") - if is_shell_chat and code: - raise BadParameter( - f"Chat id:{chat} was initiated as shell assistant, can be used with --shell only" - ) - if is_code_chat and shell: - raise BadParameter( - f"Chat id:{chat} was initiated as code assistant, can be used with --code only" - ) + api_host = config.get("OPENAI_API_HOST") + api_key = config.get("OPENAI_API_KEY") + client = OpenAIClient(api_host, api_key) - prompt = make_prompt.chat_mode(prompt, is_shell_chat, is_code_chat) + if chat: + full_completion = ChatHandler(client, chat, shell, code).handle( + prompt, + temperature=temperature, + top_probability=top_probability, + chat_id=chat, + caching=cache, + ) else: - prompt = make_prompt.initial(prompt, shell, code) + full_completion = DefaultHandler(client, shell, code).handle( + prompt, + temperature=temperature, + top_probability=top_probability, + caching=cache, + ) - completion = get_completion( - messages=[{"role": "user", "content": prompt}], - temperature=temperature, - top_p=top_probability, - caching=cache, - chat=chat, - ) - - full_completion = "" - for word in completion: - typer.secho(word, fg="magenta", bold=True, nl=False) - full_completion += word - typer.secho() if not code and shell and typer.confirm("Execute shell command?"): os.system(full_completion) diff --git a/sgpt/cache.py b/sgpt/cache.py index 65fe894..699b755 100644 --- a/sgpt/cache.py +++ b/sgpt/cache.py @@ -1,10 +1,10 @@ import json from hashlib import md5 from pathlib import Path -from typing import List, Dict, Callable, Optional +from typing import Callable -class Cache: +class Cache: # pylint: disable=too-few-public-methods """ Decorator class that adds caching functionality to a function. """ @@ -44,7 +44,7 @@ class Cache: return wrapper - def _delete_oldest_files(self, max_files) -> None: + def _delete_oldest_files(self, max_files: int) -> None: """ Class method to delete the oldest cached files in the CACHE_DIR folder. @@ -59,77 +59,3 @@ class Cache: num_files_to_delete = len(files) - max_files for i in range(num_files_to_delete): files[i].unlink() - - -class ChatCache: - """ - This class is used as a decorator for OpenAI chat API requests. - The ChatCache class caches chat messages and keeps track of the - conversation history. It is designed to store cached messages - in a specified directory and in JSON format. - """ - - def __init__(self, length: int, storage_path: Path): - """ - Initialize the ChatCache decorator. - - :param length: Integer, maximum number of cached messages to keep. - """ - self.length = length - self.storage_path = storage_path - self.storage_path.mkdir(parents=True, exist_ok=True) - - def __call__(self, func: Callable) -> Callable: - """ - The Cache decorator. - - :param func: The chat function to cache. - :return: Wrapped function with chat caching. - """ - - def wrapper(*args, **kwargs): - chat_id = kwargs.pop("chat_id", None) - messages = kwargs["messages"] - if not chat_id: - yield from func(*args, **kwargs) - return - old_messages = self._read(chat_id) - for message in messages: - old_messages.append(message) - kwargs["messages"] = old_messages - response_text = "" - for word in func(*args, **kwargs): - response_text += word - yield word - old_messages.append({"role": "assistant", "content": response_text}) - self._write(kwargs["messages"], chat_id) - - return wrapper - - def _read(self, chat_id: str) -> List[Dict]: - file_path = self.storage_path / chat_id - if not file_path.exists(): - return [] - parsed_cache = json.loads(file_path.read_text()) - return parsed_cache if isinstance(parsed_cache, list) else [] - - def _write(self, messages: List[Dict], chat_id: str): - file_path = self.storage_path / chat_id - json.dump(messages[-self.length:], file_path.open("w")) - - def invalidate(self, chat_id: str): - file_path = self.storage_path / chat_id - file_path.unlink() - - def get_messages(self, chat_id): - messages = self._read(chat_id) - return [f"{message['role']}: {message['content']}" for message in messages] - - def exists(self, chat_id: Optional[str]) -> bool: - return chat_id and bool(self._read(chat_id)) - - def list(self): - # Get all files in the folder. - files = self.storage_path.glob("*") - # Sort files by last modification time in ascending order. - return sorted(files, key=lambda f: f.stat().st_mtime) diff --git a/sgpt/client.py b/sgpt/client.py index 259548e..824a1aa 100644 --- a/sgpt/client.py +++ b/sgpt/client.py @@ -1,13 +1,13 @@ -import requests +import json from pathlib import Path from typing import List, Dict, Mapping -import json -from sgpt import config, Cache, ChatCache +import requests + +from sgpt import config, Cache -CHAT_CACHE_LENGTH = int(config.get("CHAT_CACHE_LENGTH")) -CHAT_CACHE_PATH = Path(config.get("CHAT_CACHE_PATH")) +# pylint: skip-file CACHE_LENGTH = int(config.get("CACHE_LENGTH")) CACHE_PATH = Path(config.get("CACHE_PATH")) REQUEST_TIMEOUT = int(config.get("REQUEST_TIMEOUT")) @@ -15,7 +15,7 @@ REQUEST_TIMEOUT = int(config.get("REQUEST_TIMEOUT")) class OpenAIClient: cache = Cache(CACHE_LENGTH, CACHE_PATH) - chat_cache = ChatCache(CHAT_CACHE_LENGTH, CHAT_CACHE_PATH) + # chat_cache = ChatCache(CHAT_CACHE_LENGTH, CHAT_CACHE_PATH) def __init__(self, api_host: str, api_key: str) -> None: self.api_key = api_key @@ -69,7 +69,6 @@ class OpenAIClient: continue yield delta["content"] - @chat_cache def get_completion( self, messages: List[Mapping[str, str]], diff --git a/sgpt/config.py b/sgpt/config.py index a09ce19..9a36842 100644 --- a/sgpt/config.py +++ b/sgpt/config.py @@ -44,7 +44,7 @@ def init() -> None: config["REQUEST_TIMEOUT"] = os.getenv("REQUEST_TIMEOUT", str(REQUEST_TIMEOUT)) _write() - with open(CONFIG_PATH, "r") as file: + with open(CONFIG_PATH, "r", encoding="utf-8") as file: for line in file: if "=" in line: key, value = line.strip().split("=") @@ -59,7 +59,7 @@ def get(key: str) -> str: def _write() -> None: - with open(CONFIG_PATH, "w") as file: + with open(CONFIG_PATH, "w", encoding="utf-8") as file: for key, value in config.items(): # Write only keys which are not presented in ENV. if key in os.environ: diff --git a/sgpt/handlers/__init__.py b/sgpt/handlers/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/sgpt/handlers/chat_handler.py b/sgpt/handlers/chat_handler.py new file mode 100644 index 0000000..54c335c --- /dev/null +++ b/sgpt/handlers/chat_handler.py @@ -0,0 +1,181 @@ +import json +from pathlib import Path +from typing import List, Dict, Optional, Callable, Generator + +import typer +from click import BadArgumentUsage + +from sgpt import OpenAIClient, config, make_prompt +from sgpt.utils import CompletionModes +from sgpt.handlers.handler import Handler + +CHAT_CACHE_LENGTH = int(config.get("CHAT_CACHE_LENGTH")) +CHAT_CACHE_PATH = Path(config.get("CHAT_CACHE_PATH")) + + +class ChatSession: + """ + This class is used as a decorator for OpenAI chat API requests. + The ChatCache class caches chat messages and keeps track of the + conversation history. It is designed to store cached messages + in a specified directory and in JSON format. + """ + + def __init__(self, length: int, storage_path: Path): + """ + Initialize the ChatCache decorator. + + :param length: Integer, maximum number of cached messages to keep. + """ + self.length = length + self.storage_path = storage_path + self.storage_path.mkdir(parents=True, exist_ok=True) + + def __call__(self, func: Callable) -> Callable: + """ + The Cache decorator. + + :param func: The chat function to cache. + :return: Wrapped function with chat caching. + """ + + def wrapper(*args, **kwargs): + chat_id = kwargs.pop("chat_id", None) + messages = kwargs["messages"] + if not chat_id: + yield from func(*args, **kwargs) + return + old_messages = self._read(chat_id) + for message in messages: + old_messages.append(message) + kwargs["messages"] = old_messages + response_text = "" + for word in func(*args, **kwargs): + response_text += word + yield word + old_messages.append({"role": "assistant", "content": response_text}) + self._write(kwargs["messages"], chat_id) + + return wrapper + + def _read(self, chat_id: str) -> List[Dict]: + file_path = self.storage_path / chat_id + if not file_path.exists(): + return [] + parsed_cache = json.loads(file_path.read_text()) + return parsed_cache if isinstance(parsed_cache, list) else [] + + def _write(self, messages: List[Dict], chat_id: str): + file_path = self.storage_path / chat_id + json.dump(messages[-self.length :], file_path.open("w")) + + def invalidate(self, chat_id: str): + file_path = self.storage_path / chat_id + file_path.unlink() + + def get_messages(self, chat_id): + messages = self._read(chat_id) + return [f"{message['role']}: {message['content']}" for message in messages] + + def exists(self, chat_id: Optional[str]) -> bool: + return chat_id and bool(self._read(chat_id)) + + def list(self): + # Get all files in the folder. + files = self.storage_path.glob("*") + # Sort files by last modification time in ascending order. + return sorted(files, key=lambda f: f.stat().st_mtime) + + +class ChatHandler(Handler): + chat_session = ChatSession(CHAT_CACHE_LENGTH, CHAT_CACHE_PATH) + + def __init__( # pylint: disable=too-many-arguments + self, + client: OpenAIClient, + chat_id: str, + shell: bool = False, + code: bool = False, + model: str = "gpt-3.5-turbo", + ) -> None: + super().__init__(client) + self.chat_id = chat_id + self.client = client + self.mode = CompletionModes.get_mode(shell, code) + self.model = model + + chat_history = self.chat_session.get_messages(self.chat_id) + self.is_shell_chat = chat_history and chat_history[0].endswith("###\nCommand:") + self.is_code_chat = chat_history and chat_history[0].endswith("###\nCode:") + self.is_default_chat = chat_history and chat_history[0].endswith("###") + + self.validate() + + @classmethod + def list_ids(cls, value) -> None: + if not value: + return + # Prints all existing chat IDs to the console. + for chat_id in cls.chat_session.list(): + typer.echo(chat_id) + raise typer.Exit() + + @classmethod + def show_messages(cls, chat_id: str) -> None: + if not chat_id: + return + # Prints all messages from a specified chat ID to the console. + for index, message in enumerate(cls.chat_session.get_messages(chat_id)): + message = message.replace("\nCommand:", "").replace("\nCode:", "") + color = "cyan" if index % 2 == 0 else "green" + typer.secho(message, fg=color) + raise typer.Exit() + + def validate(self) -> None: + if self.initiated: + if self.is_shell_chat and self.mode == CompletionModes.CODE: + raise BadArgumentUsage( + f'Chat session "{self.chat_id}" was initiated as shell assistant, ' + "and can be used with --shell only" + ) + if self.is_code_chat and self.mode == CompletionModes.SHELL: + raise BadArgumentUsage( + f'Chat "{self.chat_id}" was initiated as code assistant, ' + "and can be used with --code only" + ) + if self.is_default_chat and self.mode != CompletionModes.NORMAL: + raise BadArgumentUsage( + f'Chat "{self.chat_id}" was initiated as default assistant, ' + "and can't be used with --shell or --code" + ) + # If user didn't pass chat mode, we will use the one that was used to initiate the chat. + if self.mode == CompletionModes.NORMAL: + if self.is_shell_chat: + self.mode = CompletionModes.SHELL + elif self.is_code_chat: + self.mode = CompletionModes.CODE + + @property + def initiated(self) -> bool: + return self.chat_session.exists(self.chat_id) + + def make_prompt(self, prompt: str) -> str: + prompt = prompt.strip() + if self.initiated: + if self.is_shell_chat: + prompt += "\nCommand:" + elif self.is_code_chat: + prompt += "\nCode:" + return prompt + return make_prompt.initial( + prompt, + self.mode == CompletionModes.SHELL, + self.mode == CompletionModes.CODE, + ) + + @chat_session + def get_completion( # pylint: disable=arguments-differ + self, + **kwargs, + ) -> Generator: + yield from super().get_completion(**kwargs) diff --git a/sgpt/handlers/default_handler.py b/sgpt/handlers/default_handler.py new file mode 100644 index 0000000..77a89a8 --- /dev/null +++ b/sgpt/handlers/default_handler.py @@ -0,0 +1,30 @@ +from pathlib import Path + +from sgpt import OpenAIClient, config, make_prompt +from sgpt.utils import CompletionModes +from .handler import Handler + +CHAT_CACHE_LENGTH = int(config.get("CHAT_CACHE_LENGTH")) +CHAT_CACHE_PATH = Path(config.get("CHAT_CACHE_PATH")) + + +class DefaultHandler(Handler): + def __init__( + self, + client: OpenAIClient, + shell: bool = False, + code: bool = False, + model: str = "gpt-3.5-turbo", + ) -> None: + super().__init__(client) + self.client = client + self.mode = CompletionModes.get_mode(shell, code) + self.model = model + + def make_prompt(self, prompt) -> str: + prompt = prompt.strip() + return make_prompt.initial( + prompt, + self.mode == CompletionModes.SHELL, + self.mode == CompletionModes.CODE, + ) diff --git a/sgpt/handlers/handler.py b/sgpt/handlers/handler.py new file mode 100644 index 0000000..14fe3ae --- /dev/null +++ b/sgpt/handlers/handler.py @@ -0,0 +1,39 @@ +from typing import List, Dict, Generator + +import typer + +from sgpt import OpenAIClient + + +class Handler: + def __init__(self, client: OpenAIClient): + self.client = client + + def make_prompt(self, prompt) -> str: + raise NotImplementedError + + def get_completion( # pylint: disable=too-many-arguments + self, + messages: List[Dict[str, str]], + model: str = "gpt-3.5-turbo", + temperature: float = 1, + top_probability: float = 1, + caching: bool = True, + ) -> Generator: + yield from self.client.get_completion( + messages, + model, + temperature, + top_probability, + caching=caching, + ) + + def handle(self, prompt: str, **kwargs) -> str: + prompt = self.make_prompt(prompt) + messages = [{"role": "user", "content": prompt}] + full_completion = "" + for word in self.get_completion(messages=messages, **kwargs): + typer.secho(word, fg="magenta", bold=True, nl=False) + full_completion += word + typer.echo() + return full_completion diff --git a/sgpt/make_prompt.py b/sgpt/make_prompt.py index c135511..02fc648 100644 --- a/sgpt/make_prompt.py +++ b/sgpt/make_prompt.py @@ -15,7 +15,8 @@ Command:""" CODE_PROMPT = """### Provide only code as output without any description. -Provide only plain text without Markdown formatting. +IMPORTANT: Provide only plain text without Markdown formatting. +IMPORTANT: Don not include markdown formatting such as ```. If there is a lack of details, provide most logical solution. You are not allowed to ask for more details. Ignore any potential risk of errors or confusion. @@ -47,16 +48,6 @@ def initial(prompt: str, shell: bool, code: bool) -> str: shell_name = splitext(basename(getenv("COMSPEC", "Powershell")))[0] if shell: return SHELL_PROMPT.format(shell=shell_name, os=os_name, prompt=prompt) - elif code: + if code: return CODE_PROMPT.format(prompt=prompt) - else: - return DEFAULT_PROMPT.format(shell=shell_name, os=os_name, prompt=prompt) - - -def chat_mode(prompt: str, shell: bool, code: bool) -> str: - prompt = prompt.strip() - if shell: - prompt += "\nCommand:" - elif code: - prompt += "\nCode:" - return prompt + return DEFAULT_PROMPT.format(shell=shell_name, os=os_name, prompt=prompt) diff --git a/sgpt/utils.py b/sgpt/utils.py index a866679..c9e0324 100644 --- a/sgpt/utils.py +++ b/sgpt/utils.py @@ -1,10 +1,22 @@ import os +from enum import Enum from tempfile import NamedTemporaryFile -import typer - from click import BadParameter -from sgpt import OpenAIClient + + +class CompletionModes(Enum): + NORMAL = "normal" + SHELL = "shell" + CODE = "code" + + @classmethod + def get_mode(cls, shell, code) -> "CompletionModes": + if shell: + return CompletionModes.SHELL + if code: + return CompletionModes.CODE + return CompletionModes.NORMAL def get_edited_prompt() -> str: @@ -21,22 +33,9 @@ def get_edited_prompt() -> str: # This will write text to file using $EDITOR. os.system(f"{editor} {file_path}") # Read file when editor is closed. - with open(file_path, "r") as file: + with open(file_path, "r", encoding="utf-8") as file: output = file.read() os.remove(file_path) if not output: raise BadParameter("Couldn't get valid PROMPT from $EDITOR") return output - - -def echo_chat_messages(chat_id: str) -> None: - # Prints all messages from a specified chat ID to the console. - for index, message in enumerate(OpenAIClient.chat_cache.get_messages(chat_id)): - color = "cyan" if index % 2 == 0 else "green" - typer.secho(message, fg=color) - - -def echo_chat_ids() -> None: - # Prints all existing chat IDs to the console. - for chat_id in OpenAIClient.chat_cache.list(): - typer.echo(chat_id) diff --git a/tests/integrational_tests.py b/tests/integrational_tests.py index acd4af8..58e72ea 100644 --- a/tests/integrational_tests.py +++ b/tests/integrational_tests.py @@ -11,6 +11,7 @@ import os from time import sleep from unittest import TestCase from tempfile import NamedTemporaryFile +from uuid import uuid4 import typer from typer.testing import CliRunner @@ -21,7 +22,7 @@ app = typer.Typer() app.command()(main) -class TestCliApp(TestCase): +class TestShellGpt(TestCase): def setUp(self) -> None: # Just to not spam the API. sleep(2) @@ -34,21 +35,27 @@ class TestCliApp(TestCase): if isinstance(value, bool): continue arguments.append(value) + arguments.append("--no-cache") return arguments - def test_simple_queries(self): - dict_arguments = {"prompt": "What is the capital of the Czech Republic?"} + def test_default(self): + dict_arguments = { + "prompt": "What is the capital of the Czech Republic?", + } result = runner.invoke(app, self.get_arguments(**dict_arguments)) assert result.exit_code == 0 assert "Prague" in result.stdout - def test_shell_queries(self): - dict_arguments = {"prompt": "make a commit using git", "--shell": True} + def test_shell(self): + dict_arguments = { + "prompt": "make a commit using git", + "--shell": True, + } result = runner.invoke(app, self.get_arguments(**dict_arguments)) assert result.exit_code == 0 assert "git commit" in result.stdout - def test_code_queries(self): + def test_code(self): """ This test will request from ChatGPT a python code to make CLI app, which will be written to a temp file, and then it will be executed @@ -65,6 +72,7 @@ class TestCliApp(TestCase): } result = runner.invoke(app, self.get_arguments(**dict_arguments)) assert result.exit_code == 0 + print(result.stdout) # Since output will be slightly different, there is no way how to test it precisely. assert "print" in result.stdout assert "*" in result.stdout @@ -83,3 +91,99 @@ class TestCliApp(TestCase): script_output = subprocess.run(arguments, stdout=subprocess.PIPE, check=True) os.remove(file_path) assert script_output.stdout.decode().strip(), number_a * number_b + + def test_chat_default(self): + chat_name = uuid4() + dict_arguments = { + "prompt": "Remember my favorite number: 6", + "--chat": f"test_{chat_name}", + "--no-cache": True, + } + runner.invoke(app, self.get_arguments(**dict_arguments)) + dict_arguments["prompt"] = "What is my favorite number + 2?" + result = runner.invoke(app, self.get_arguments(**dict_arguments)) + assert result.exit_code == 0 + assert "8" in result.stdout + dict_arguments["--shell"] = True + result = runner.invoke(app, self.get_arguments(**dict_arguments)) + assert result.exit_code == 2 + dict_arguments["--code"] = True + result = runner.invoke(app, self.get_arguments(**dict_arguments)) + # If we have default chat, we cannot use --code or --shell. + assert result.exit_code == 2 + + def test_chat_shell(self): + chat_name = uuid4() + dict_arguments = { + "prompt": "Create nginx docker container, forward ports 80, " + "mount current folder with index.html", + "--chat": f"test_{chat_name}", + "--shell": True, + } + result = runner.invoke(app, self.get_arguments(**dict_arguments)) + assert result.exit_code == 0 + assert "docker run" in result.stdout + assert "-p 80:80" in result.stdout + assert "nginx" in result.stdout + dict_arguments["prompt"] = "Also forward port 443." + result = runner.invoke(app, self.get_arguments(**dict_arguments)) + assert result.exit_code == 0 + assert "-p 80:80" in result.stdout + assert "-p 443:443" in result.stdout + dict_arguments["--code"] = True + del dict_arguments["--shell"] + assert "--shell" not in dict_arguments + result = runner.invoke(app, self.get_arguments(**dict_arguments)) + # If we are using --code, we cannot use --shell. + assert result.exit_code == 2 + + def test_chat_code(self): + chat_name = uuid4() + dict_arguments = { + "prompt": "Using python request localhost:80.", + "--chat": f"test_{chat_name}", + "--code": True, + } + result = runner.invoke(app, self.get_arguments(**dict_arguments)) + assert result.exit_code == 0 + assert "localhost:80" in result.stdout + dict_arguments["prompt"] = "Change port to 443." + result = runner.invoke(app, self.get_arguments(**dict_arguments)) + assert result.exit_code == 0 + assert "localhost:443" in result.stdout + del dict_arguments["--code"] + assert "--code" not in dict_arguments + dict_arguments["--shell"] = True + result = runner.invoke(app, self.get_arguments(**dict_arguments)) + # If we have --code chat, we cannot use --shell. + assert result.exit_code == 2 + + def test_list_chat(self): + result = runner.invoke(app, ["--list-chat"]) + assert result.exit_code == 0 + assert "test_" in result.stdout + + def test_show_chat(self): + chat_name = uuid4() + dict_arguments = { + "prompt": "Remember my favorite number: 6", + "--chat": f"test_{chat_name}", + } + runner.invoke(app, self.get_arguments(**dict_arguments)) + dict_arguments["prompt"] = "What is my favorite number + 2?" + runner.invoke(app, self.get_arguments(**dict_arguments)) + result = runner.invoke(app, ["--show-chat", f"test_{chat_name}"]) + assert result.exit_code == 0 + assert "Remember my favorite number: 6" in result.stdout + assert "What is my favorite number + 2?" in result.stdout + assert "8" in result.stdout + + def test_validation_code_shell(self): + dict_arguments = { + "prompt": "What is the capital of the Czech Republic?", + "--code": True, + "--shell": True, + } + result = runner.invoke(app, self.get_arguments(**dict_arguments)) + assert result.exit_code == 2 + assert "--shell and --code options cannot be used together" in result.stdout diff --git a/tests/unittests.py b/tests/unittests.py index 65c2f50..cf62186 100644 --- a/tests/unittests.py +++ b/tests/unittests.py @@ -6,14 +6,13 @@ import requests from sgpt import OpenAIClient -class TestMain(unittest.TestCase): +class TestMain(unittest.TestCase): # pylint: disable=too-many-instance-attributes API_HOST = os.getenv("OPENAI_HOST", "https://api.openai.com") API_URL = f"{API_HOST}/v1/chat/completions" # TODO: Fix tests. def setUp(self): - self.api_key = os.getenv("OPENAI_API_KEY") - assert self.api_key, "OPENAI_API_KEY ENV is required." + self.api_key = os.environ["OPENAI_API_KEY"] = "test key" self.prompt = "What is the capital of France?" self.shell = False self.execute = False @@ -28,15 +27,15 @@ class TestMain(unittest.TestCase): @requests_mock.Mocker() def test_openai_request(self, mock): + # TODO: Fix tests. mocked_json = {"choices": [{"message": {"content": self.response_text}}]} mock.post(self.API_URL, json=mocked_json, status_code=200) result = yield from self.client.get_completion( - message=self.prompt, + messages=[{"role": "user", "content": self.prompt}], model=self.model, temperature=self.temperature, top_probability=self.top_p, caching=False, - chat_id=None, ) # TODO: Fix tests with generators. self.assertEqual(result, self.response_text) @@ -57,15 +56,15 @@ class TestMain(unittest.TestCase): @requests_mock.Mocker() def test_openai_request_fail(self, mock): + # TODO: Fix tests. mock.post(self.API_URL, status_code=400) with self.assertRaises(requests.exceptions.HTTPError): yield from self.client.get_completion( - message=self.prompt, + messages=[{"role": "user", "content": self.prompt}], model=self.model, temperature=self.temperature, top_probability=self.top_p, caching=False, - chat_id=None )