Refactoring, more tests, lint fixes, optimise

* Refactoring, more tests, lint fixes, optimise
* Lint CI fixes
* Changing version, readme help fixes
* Validation rule, validation tests, minor improvements
This commit is contained in:
Farkhod Sadykov
2023-04-02 02:52:55 +02:00
committed by GitHub
parent a0eb1ad720
commit 763fef15d3
17 changed files with 514 additions and 237 deletions
+8 -9
View File
@@ -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
+20 -16
View File
@@ -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
+7 -8
View File
@@ -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 ✅'
+1 -1
View File
@@ -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",
+2 -1
View File
@@ -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
+79 -73
View File
@@ -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)
+3 -77
View File
@@ -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)
+6 -7
View File
@@ -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]],
+2 -2
View File
@@ -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:
View File
+181
View File
@@ -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)
+30
View File
@@ -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,
)
+39
View File
@@ -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
+4 -13
View File
@@ -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)
+16 -17
View File
@@ -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)
+110 -6
View File
@@ -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
+6 -7
View File
@@ -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
)