Compare commits
10 Commits
juk
...
c5c4a6628f
| Author | SHA1 | Date | |
|---|---|---|---|
| c5c4a6628f | |||
| f8ed0e3636 | |||
| caf5244d52 | |||
| ca3a53e68b | |||
| 820d938060 | |||
| 8bb2a002a6 | |||
| 01de75bef3 | |||
| 1bb553b223 | |||
| bb8aa2f817 | |||
| 0d31b88567 |
@@ -5,6 +5,17 @@ def openai_api_key(api_key: str) -> None:
|
||||
openai.api_key = api_key
|
||||
|
||||
|
||||
def display_models() -> None:
|
||||
not_ready = []
|
||||
for engine in sorted(openai.Engine.list()['data'], key=lambda x: x['id']):
|
||||
if engine['ready']:
|
||||
print(engine['id'])
|
||||
else:
|
||||
not_ready.append(engine['id'])
|
||||
if len(not_ready) > 0:
|
||||
print('\nNot ready: ' + ', '.join(not_ready))
|
||||
|
||||
|
||||
def ai(chat: list[dict[str, str]],
|
||||
config: dict,
|
||||
number: int
|
||||
|
||||
+12
-4
@@ -9,7 +9,7 @@ import argparse
|
||||
import pathlib
|
||||
from .utils import terminal_width, process_tags, display_chat, display_source_code, display_tags_frequency
|
||||
from .storage import save_answers, create_chat, get_tags, get_tags_unique, read_file, dump_data
|
||||
from .api_client import ai, openai_api_key
|
||||
from .api_client import ai, openai_api_key, display_models
|
||||
from itertools import zip_longest
|
||||
|
||||
|
||||
@@ -61,9 +61,10 @@ def process_and_display_chat(args: argparse.Namespace,
|
||||
display_chat(chat, dump, args.only_source_code)
|
||||
return chat, full_question, tags
|
||||
|
||||
|
||||
def process_and_display_tags(args: argparse.Namespace,
|
||||
config: dict,
|
||||
dump: bool=False
|
||||
dump: bool = False
|
||||
) -> None:
|
||||
display_tags_frequency(get_tags(config, None), dump)
|
||||
|
||||
@@ -96,6 +97,7 @@ def create_parser() -> argparse.ArgumentParser:
|
||||
group.add_argument('-D', '--chat-dump', help="Print chat history as Python structure", action='store_true')
|
||||
group.add_argument('-d', '--chat', help="Print chat history as readable text", action='store_true')
|
||||
group.add_argument('-l', '--list-tags', help="List all tags and their frequency", action='store_true')
|
||||
group.add_argument('-L', '--list-models', help="List all available models", action='store_true')
|
||||
parser.add_argument('-c', '--config', help='Config file name.', default=default_config)
|
||||
parser.add_argument('-m', '--max-tokens', help='Max tokens to use', type=int)
|
||||
parser.add_argument('-T', '--temperature', help='Temperature to use', type=float)
|
||||
@@ -104,8 +106,12 @@ def create_parser() -> argparse.ArgumentParser:
|
||||
parser.add_argument('-s', '--source', nargs='*', help='Source add content of a file to the query')
|
||||
parser.add_argument('-S', '--only-source-code', help='Print only source code', action='store_true')
|
||||
parser.add_argument('-w', '--with-tags', help="Print chat history with tags.", action='store_true')
|
||||
parser.add_argument('-W', '--with-file', help="Print chat history with filename.", action='store_true')
|
||||
parser.add_argument('-a', '--match-all-tags', help="All given tags must match when selecting chat history entries.", action='store_true')
|
||||
parser.add_argument('-W', '--with-file',
|
||||
help="Print chat history with filename.",
|
||||
action='store_true')
|
||||
parser.add_argument('-a', '--match-all-tags',
|
||||
help="All given tags must match when selecting chat history entries.",
|
||||
action='store_true')
|
||||
tags_arg = parser.add_argument('-t', '--tags', nargs='*', help='List of tag names', metavar='TAGS')
|
||||
tags_arg.completer = tags_completer # type: ignore
|
||||
extags_arg = parser.add_argument('-e', '--extags', nargs='*', help='List of tag names to exclude', metavar='EXTAGS')
|
||||
@@ -144,6 +150,8 @@ def main() -> int:
|
||||
process_and_display_chat(args, config)
|
||||
elif args.list_tags:
|
||||
process_and_display_tags(args, config)
|
||||
elif args.list_models:
|
||||
display_models()
|
||||
|
||||
return 0
|
||||
|
||||
|
||||
@@ -5,23 +5,26 @@ from .utils import terminal_width, append_message, message_to_chat
|
||||
from typing import List, Dict, Any, Optional
|
||||
|
||||
|
||||
def read_file(fname: str, tags_only: bool = False) -> Dict[str, Any]:
|
||||
def read_file(fname: pathlib.Path, tags_only: bool = False) -> Dict[str, Any]:
|
||||
with open(fname, "r") as fd:
|
||||
tagline = fd.readline().strip().split(':', maxsplit=1)[1].strip()
|
||||
# also support tags separated by ',' (old format)
|
||||
separator = ',' if ',' in tagline else ' '
|
||||
tags = [t.strip() for t in tagline.split(separator)]
|
||||
if tags_only:
|
||||
return {"tags": [x.strip() for x in fd.readline().strip().split(':')[1].strip().split(',')]}
|
||||
return {"tags": tags}
|
||||
text = fd.read().strip().split('\n')
|
||||
tags = [x.strip() for x in text.pop(0).split(':')[1].strip().split(',')]
|
||||
question_idx = text.index("=== QUESTION ===") + 1
|
||||
answer_idx = text.index("==== ANSWER ====")
|
||||
question = "\n".join(text[question_idx:answer_idx]).strip()
|
||||
answer = "\n".join(text[answer_idx + 1:]).strip()
|
||||
return {"question": question, "answer": answer, "tags": tags,
|
||||
"file": pathlib.Path(fname).name}
|
||||
"file": fname.name}
|
||||
|
||||
|
||||
def dump_data(data: Dict[str, Any]) -> str:
|
||||
with io.StringIO() as fd:
|
||||
fd.write(f'TAGS: {", ".join(data["tags"])}\n')
|
||||
fd.write(f'TAGS: {" ".join(data["tags"])}\n')
|
||||
fd.write(f'=== QUESTION ===\n{data["question"]}\n')
|
||||
fd.write(f'==== ANSWER ====\n{data["answer"]}\n')
|
||||
return fd.getvalue()
|
||||
@@ -29,7 +32,7 @@ def dump_data(data: Dict[str, Any]) -> str:
|
||||
|
||||
def write_file(fname: str, data: Dict[str, Any]) -> None:
|
||||
with open(fname, "w") as fd:
|
||||
fd.write(f'TAGS: {", ".join(data["tags"])}\n')
|
||||
fd.write(f'TAGS: {" ".join(data["tags"])}\n')
|
||||
fd.write(f'=== QUESTION ===\n{data["question"]}\n')
|
||||
fd.write(f'==== ANSWER ====\n{data["answer"]}\n')
|
||||
|
||||
@@ -74,6 +77,7 @@ def create_chat(question: Optional[str],
|
||||
if file.suffix == '.yaml':
|
||||
with open(file, 'r') as f:
|
||||
data = yaml.load(f, Loader=yaml.FullLoader)
|
||||
data['file'] = file.name
|
||||
elif file.suffix == '.txt':
|
||||
data = read_file(file)
|
||||
else:
|
||||
@@ -111,5 +115,6 @@ def get_tags(config: Dict[str, Any], prefix: Optional[str]) -> List[str]:
|
||||
result.append(tag)
|
||||
return result
|
||||
|
||||
|
||||
def get_tags_unique(config: Dict[str, Any], prefix: Optional[str]) -> List[str]:
|
||||
return list(set(get_tags(config, prefix)))
|
||||
|
||||
@@ -15,11 +15,11 @@ def process_tags(tags: list[str], extags: list[str], otags: list[str]) -> None:
|
||||
printed_messages = []
|
||||
|
||||
if tags:
|
||||
printed_messages.append(f"Tags: {', '.join(tags)}")
|
||||
printed_messages.append(f"Tags: {' '.join(tags)}")
|
||||
if extags:
|
||||
printed_messages.append(f"Excluding tags: {', '.join(extags)}")
|
||||
printed_messages.append(f"Excluding tags: {' '.join(extags)}")
|
||||
if otags:
|
||||
printed_messages.append(f"Output tags: {', '.join(otags)}")
|
||||
printed_messages.append(f"Output tags: {' '.join(otags)}")
|
||||
|
||||
if printed_messages:
|
||||
print("\n".join(printed_messages))
|
||||
@@ -41,7 +41,7 @@ def message_to_chat(message: Dict[str, str],
|
||||
append_message(chat, 'user', message['question'])
|
||||
append_message(chat, 'assistant', message['answer'])
|
||||
if with_tags:
|
||||
tags = ", ".join(message['tags'])
|
||||
tags = " ".join(message['tags'])
|
||||
append_message(chat, 'tags', tags)
|
||||
if with_file:
|
||||
append_message(chat, 'file', message['file'])
|
||||
@@ -74,9 +74,10 @@ def display_chat(chat, dump=False, source_code=False) -> None:
|
||||
else:
|
||||
print(f"{message['role'].upper()}: {message['content']}")
|
||||
|
||||
|
||||
def display_tags_frequency(tags: List[str], dump=False) -> None:
|
||||
if dump:
|
||||
pp(tags)
|
||||
return
|
||||
for tag in set(tags):
|
||||
print(f"-{tag} : {tags.count(tag)}")
|
||||
print(f"- {tag}: {tags.count(tag)}")
|
||||
|
||||
+17
-7
@@ -21,9 +21,11 @@ class TestCreateChat(unittest.TestCase):
|
||||
self.tags = ['test_tag']
|
||||
|
||||
@patch('os.listdir')
|
||||
@patch('pathlib.Path.iterdir')
|
||||
@patch('builtins.open')
|
||||
def test_create_chat_with_tags(self, open_mock, listdir_mock):
|
||||
def test_create_chat_with_tags(self, open_mock, iterdir_mock, listdir_mock):
|
||||
listdir_mock.return_value = ['testfile.txt']
|
||||
iterdir_mock.return_value = [pathlib.Path(x) for x in listdir_mock.return_value]
|
||||
open_mock.return_value.__enter__.return_value = io.StringIO(dump_data(
|
||||
{'question': 'test_content', 'answer': 'some answer',
|
||||
'tags': ['test_tag']}))
|
||||
@@ -41,9 +43,11 @@ class TestCreateChat(unittest.TestCase):
|
||||
{'role': 'user', 'content': self.question})
|
||||
|
||||
@patch('os.listdir')
|
||||
@patch('pathlib.Path.iterdir')
|
||||
@patch('builtins.open')
|
||||
def test_create_chat_with_other_tags(self, open_mock, listdir_mock):
|
||||
def test_create_chat_with_other_tags(self, open_mock, iterdir_mock, listdir_mock):
|
||||
listdir_mock.return_value = ['testfile.txt']
|
||||
iterdir_mock.return_value = [pathlib.Path(x) for x in listdir_mock.return_value]
|
||||
open_mock.return_value.__enter__.return_value = io.StringIO(dump_data(
|
||||
{'question': 'test_content', 'answer': 'some answer',
|
||||
'tags': ['other_tag']}))
|
||||
@@ -57,9 +61,11 @@ class TestCreateChat(unittest.TestCase):
|
||||
{'role': 'user', 'content': self.question})
|
||||
|
||||
@patch('os.listdir')
|
||||
@patch('pathlib.Path.iterdir')
|
||||
@patch('builtins.open')
|
||||
def test_create_chat_without_tags(self, open_mock, listdir_mock):
|
||||
def test_create_chat_without_tags(self, open_mock, iterdir_mock, listdir_mock):
|
||||
listdir_mock.return_value = ['testfile.txt', 'testfile2.txt']
|
||||
iterdir_mock.return_value = [pathlib.Path(x) for x in listdir_mock.return_value]
|
||||
open_mock.side_effect = (
|
||||
io.StringIO(dump_data({'question': 'test_content',
|
||||
'answer': 'some answer',
|
||||
@@ -95,7 +101,10 @@ class TestHandleQuestion(unittest.TestCase):
|
||||
question=[self.question],
|
||||
source=None,
|
||||
only_source_code=False,
|
||||
number=3
|
||||
number=3,
|
||||
match_all_tags=False,
|
||||
with_tags=False,
|
||||
with_file=False,
|
||||
)
|
||||
self.config = {
|
||||
'db': 'test_files',
|
||||
@@ -119,7 +128,8 @@ class TestHandleQuestion(unittest.TestCase):
|
||||
mock_create_chat.assert_called_once_with(self.question,
|
||||
self.args.tags,
|
||||
self.args.extags,
|
||||
self.config)
|
||||
self.config,
|
||||
False, False, False)
|
||||
mock_pp.assert_called_once_with("test_chat")
|
||||
mock_ai.assert_called_with("test_chat",
|
||||
self.config,
|
||||
@@ -203,7 +213,7 @@ class TestCreateParser(unittest.TestCase):
|
||||
mock_add_mutually_exclusive_group.assert_called_once_with(required=True)
|
||||
mock_group.add_argument.assert_any_call('-p', '--print', help='File to print')
|
||||
mock_group.add_argument.assert_any_call('-q', '--question', nargs='*', help='Question to ask')
|
||||
mock_group.add_argument.assert_any_call('-D', '--chat-dump', help="Print chat as Python structure", action='store_true')
|
||||
mock_group.add_argument.assert_any_call('-d', '--chat', help="Print chat as readable text", action='store_true')
|
||||
mock_group.add_argument.assert_any_call('-D', '--chat-dump', help="Print chat history as Python structure", action='store_true')
|
||||
mock_group.add_argument.assert_any_call('-d', '--chat', help="Print chat history as readable text", action='store_true')
|
||||
self.assertTrue('.config.yaml' in parser.get_default('config'))
|
||||
self.assertEqual(parser.get_default('number'), 1)
|
||||
|
||||
Reference in New Issue
Block a user