mirror of
https://github.com/TheR1D/shell_gpt.git
synced 2026-07-03 14:10:18 +02:00
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:
@@ -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
|
||||
|
||||
@@ -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
@@ -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 ✅'
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user