diff --git a/nbs/00_core.ipynb b/nbs/00_core.ipynb
index 8119bfc..9fa5939 100644
--- a/nbs/00_core.ipynb
+++ b/nbs/00_core.ipynb
@@ -1048,7 +1048,7 @@
" pid: str = 'current', # `current`, `all` or tmux pane_id (e.g. %0) for context\n",
" skip_system: bool = False, # Whether to skip system information in the AI's context\n",
" history_lines: int = None, # Number of history lines. Defaults to tmux scrollback history length\n",
- " mode: str = 'default', # Available ShellSage modes: ['default', 'sassy']\n",
+ " mode: str = None, # Available ShellSage modes: ['default', 'sassy']\n",
" model: str = None, # The LLM model that will be invoked on the LLM provider\n",
" vendor_name: str = None, # Vendor name for non auto-resolved models (e.g. 'codex', 'fireworks_ai', 'moonshot', 'deepseek', ...)\n",
" search: str = None, # Wheather to allow the LLM to search the internet\n",
@@ -1062,7 +1062,7 @@
" custom_instructions: str = None, # Extra instructions appended to the system prompt\n",
"):\n",
" safecmd = None\n",
- " opts = get_opts(history_lines=history_lines, model=model, search=search,\n",
+ " opts = get_opts(history_lines=history_lines, model=model, mode=mode, search=search,\n",
" base_url=base_url, api_key=api_key, code_theme=code_theme,\n",
" code_lexer=code_lexer, think=think, trust=trust, safecmd=safecmd,\n",
" vendor_name=vendor_name, log=None, custom_instructions=custom_instructions)\n",
@@ -1072,8 +1072,8 @@
" with Live(Spinner(\"dots\", text=\"Connecting...\"), auto_refresh=False) as live:\n",
" global _live, _md\n",
" _live = live\n",
- " if mode not in ['default', 'sassy']:\n",
- " raise Exception(f\"{mode} is not valid. Must be one of the following: ['default', 'sassy']\")\n",
+ " if opts.mode not in ['default', 'sassy']:\n",
+ " raise Exception(f\"{opts.mode} is not valid. Must be one of the following: ['default', 'sassy']\")\n",
" \n",
" _md = noop if raw else partial(Markdown, code_theme=opts.code_theme, inline_code_lexer=opts.code_lexer,\n",
" inline_code_theme=opts.code_theme)\n",
@@ -1098,14 +1098,14 @@
" \n",
" query = f'{ctxt}\\n\\n{query}\\n'\n",
"\n",
- " sage = get_sage(opts.model, mode, search=opts.search, use_safecmd=opts.safecmd, vendor_name=opts.vendor_name, custom_instructions=opts.custom_instructions)\n",
+ " sage = get_sage(opts.model, opts.mode, search=opts.search, use_safecmd=opts.safecmd, vendor_name=opts.vendor_name, custom_instructions=opts.custom_instructions)\n",
" async for res in get_res(sage, query, opts): live.update(_md(res), refresh=True)\n",
" \n",
" # Handle logging if the log flag is set\n",
" if opts.log:\n",
" db = mk_db()\n",
" db.logs.insert(Log(timestamp=datetime.now().isoformat(), query=query,\n",
- " response=res, model=opts.model, mode=mode))\n",
+ " response=res, model=opts.model, mode=opts.mode))\n",
" except KeyboardInterrupt: print(\"Interrupted.\")"
]
},
diff --git a/shell_sage/core.py b/shell_sage/core.py
index 1d3360e..944d47a 100644
--- a/shell_sage/core.py
+++ b/shell_sage/core.py
@@ -326,7 +326,7 @@ async def main(
pid: str = 'current', # `current`, `all` or tmux pane_id (e.g. %0) for context
skip_system: bool = False, # Whether to skip system information in the AI's context
history_lines: int = None, # Number of history lines. Defaults to tmux scrollback history length
- mode: str = 'default', # Available ShellSage modes: ['default', 'sassy']
+ mode: str = None, # Available ShellSage modes: ['default', 'sassy']
model: str = None, # The LLM model that will be invoked on the LLM provider
vendor_name: str = None, # Vendor name for non auto-resolved models (e.g. 'codex', 'fireworks_ai', 'moonshot', 'deepseek', ...)
search: str = None, # Wheather to allow the LLM to search the internet
@@ -340,7 +340,7 @@ async def main(
custom_instructions: str = None, # Extra instructions appended to the system prompt
):
safecmd = None
- opts = get_opts(history_lines=history_lines, model=model, search=search,
+ opts = get_opts(history_lines=history_lines, model=model, mode=mode, search=search,
base_url=base_url, api_key=api_key, code_theme=code_theme,
code_lexer=code_lexer, think=think, trust=trust, safecmd=safecmd,
vendor_name=vendor_name, log=None, custom_instructions=custom_instructions)
@@ -350,8 +350,8 @@ async def main(
with Live(Spinner("dots", text="Connecting..."), auto_refresh=False) as live:
global _live, _md
_live = live
- if mode not in ['default', 'sassy']:
- raise Exception(f"{mode} is not valid. Must be one of the following: ['default', 'sassy']")
+ if opts.mode not in ['default', 'sassy']:
+ raise Exception(f"{opts.mode} is not valid. Must be one of the following: ['default', 'sassy']")
_md = noop if raw else partial(Markdown, code_theme=opts.code_theme, inline_code_lexer=opts.code_lexer,
inline_code_theme=opts.code_theme)
@@ -376,14 +376,14 @@ async def main(
query = f'{ctxt}\n\n{query}\n'
- sage = get_sage(opts.model, mode, search=opts.search, use_safecmd=opts.safecmd, vendor_name=opts.vendor_name, custom_instructions=opts.custom_instructions)
+ sage = get_sage(opts.model, opts.mode, search=opts.search, use_safecmd=opts.safecmd, vendor_name=opts.vendor_name, custom_instructions=opts.custom_instructions)
async for res in get_res(sage, query, opts): live.update(_md(res), refresh=True)
# Handle logging if the log flag is set
if opts.log:
db = mk_db()
db.logs.insert(Log(timestamp=datetime.now().isoformat(), query=query,
- response=res, model=opts.model, mode=mode))
+ response=res, model=opts.model, mode=opts.mode))
except KeyboardInterrupt: print("Interrupted.")
# %% ../nbs/00_core.ipynb #f093b48b