Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
77 changes: 69 additions & 8 deletions graybox/ai/ai_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,9 @@
Service layer for External APIs (LLMs, Embeddings, Rerankers).
"""

import base64
import os

os.environ.setdefault("LITELLM_LOCAL_MODEL_COST_MAP", "true")

from typing import Optional, List, Any
Expand All @@ -24,6 +26,7 @@
responses,
aresponses,
)
from litellm.utils import supports_pdf_input
import litellm
import requests
import logging
Expand Down Expand Up @@ -114,7 +117,57 @@ def get_embedding_params(self, **kwargs):
}
return params

def _build_messages(self, system_prompt: str, prompt: str) -> list[dict]:
def _build_messages(
self, system_prompt: str, prompt: str, image: str, file_input: str
) -> list[dict]:

if image and file_input:
raise ValueError(
"Cannot provide both `image` and `file_input` simultaneously."
)

if file_input:
if not supports_pdf_input(self.config.llm.model_name, None):
raise ValueError("Model does not support PDF input")

return [
{
"role": "system",
"content": system_prompt or "You are a helpful assistant.",
},
{"type": "text", "text": prompt},
{
"type": "file",
"file": {
"file_id": file_input,
"format": "application/pdf",
},
},
]

if image:
if not image.startswith(("http://", "https://")):
try:
with open(image, "rb") as image_file:
image = base64.b64encode(image_file.read())
image = f"data:image/png;base64,{image.decode('utf-8')}"
except Exception:
raise ValueError(f"Invalid Image path.")

return [
{
"role": "system",
"content": system_prompt or "You are a helpful assistant.",
},
{
"role": "user",
"content": [
{"type": "text", "text": prompt},
{"type": "image_url", "image_url": {"url": image}},
],
},
]

return [
{
"role": "system",
Expand Down Expand Up @@ -155,8 +208,10 @@ def _extract_response(self, response, api_type: str) -> dict:
)
def llm_call(self, system_prompt: str = None, prompt: str = None, **kwargs) -> dict:
stream = kwargs.pop("stream", False)
image = kwargs.pop("image", None)
file_input = kwargs.pop("file_input", None)
messages = self._build_messages(system_prompt, prompt, image, file_input)
params = self.get_llm_params(**kwargs)
messages = self._build_messages(system_prompt, prompt)

try:
if params["api_type"] == "responses":
Expand Down Expand Up @@ -201,8 +256,10 @@ def llm_call(self, system_prompt: str = None, prompt: str = None, **kwargs) -> d
def llm_call_stream(
self, system_prompt: str = None, prompt: str = None, **kwargs
) -> dict:
image = kwargs.pop("image", None)
file_input = kwargs.pop("file_input", None)
messages = self._build_messages(system_prompt, prompt, image, file_input)
params = self.get_llm_params(**kwargs)
messages = self._build_messages(system_prompt, prompt)

try:
if params["api_type"] == "responses":
Expand Down Expand Up @@ -251,13 +308,15 @@ def llm_call_stream(
def llm_call_batch(
self, system_prompt: str = None, prompts: list[str] = None, **kwargs
) -> dict:
images = kwargs.pop("images", [None] * len(prompts or []))
file_inputs = kwargs.pop("file_inputs", [None] * len(prompts or []))
batched_messages = [
self._build_messages(system_prompt, prompt, image, file_input)
for prompt, image, file_input in zip(prompts, images, file_inputs)
]
params = self.get_llm_params(**kwargs)
if params["api_type"] != "chat_completion":
raise ValueError("Batch calls only support chat_completion")

batched_messages = [
self._build_messages(system_prompt, p) for p in (prompts or [])
]
try:
responses_list = batch_completion(
model=params["model"],
Expand Down Expand Up @@ -298,8 +357,10 @@ def llm_call_batch(
async def llm_call_async(
self, system_prompt: str = None, prompt: str = None, **kwargs
) -> dict:
image = kwargs.pop("image", None)
file_input = kwargs.pop("file_input", None)
messages = self._build_messages(system_prompt, prompt, image, file_input)
params = self.get_llm_params(**kwargs)
messages = self._build_messages(system_prompt, prompt)

try:
if params["api_type"] == "responses":
Expand Down
10 changes: 10 additions & 0 deletions graybox/config.example.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,16 @@ llm:
base_url: http://localhost:11434 #leave this to empty string if not applicable
temperature: 0.1

# ---------------------------------------------------------------------------
# Ask / Chat answer style — OPTIONAL.
# This controls presentation only (tone, verbosity, structure, formatting,
# audience, personality, explanation style, and requested language). It does
# NOT override Gray Box's grounding, citation, uncertainty, or refusal rules.
# Environment variable GRAYBOX_ANSWER_STYLE_PROMPT overrides this value.
# ---------------------------------------------------------------------------
prompts:
answer_style: ""

# ---------------------------------------------------------------------------
# Retrieval tuning — the defaults below are sane, leave them unless you've
# noticed search feels too strict/loose in `ask`/`search`.
Expand Down
12 changes: 12 additions & 0 deletions graybox/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,9 @@
"temperature": 0.0,
},
"auto_refresh_summaries": True,
"prompts": {
"answer_style": "",
},
"retrieval": {
"top_k": 5,
# Both thresholds below are normalized to [0, 1] on the SAME scale:
Expand Down Expand Up @@ -138,13 +141,19 @@ class EmbeddingsConfig:
kwargs: dict = field(default_factory=dict)


@dataclass
class PromptsConfig:
answer_style: str = ""


@dataclass
class Config:
root: Path
workspace_manager: WorkspaceManager
llm: LLMConfig
retrieval: RetrievalConfig
embeddings: EmbeddingsConfig
prompts: PromptsConfig = field(default_factory=PromptsConfig)
auto_refresh_summaries: bool = True
raw: dict = field(default_factory=dict)
config_path: Path | None = None
Expand Down Expand Up @@ -221,6 +230,7 @@ def for_workspace(self, workspace: Workspace | str) -> "Config":
llm=self.llm,
retrieval=self.retrieval,
embeddings=self.embeddings,
prompts=self.prompts,
raw=copy.deepcopy(self.raw),
config_path=self.config_path,
)
Expand Down Expand Up @@ -288,6 +298,7 @@ def load_config(path: str | None = None) -> Config:
"GRAYBOX_DEFAULT_WORKSPACE": ("default_workspace",),
"GRAYBOX_LLM_MODEL": ("llm", "model_name"),
"GRAYBOX_LLM_BASE_URL": ("llm", "base_url"),
"GRAYBOX_ANSWER_STYLE_PROMPT": ("prompts", "answer_style"),
"GRAYBOX_LLM_API_KEY": ("llm", "api_key"),
"GRAYBOX_TEMPERATURE": ("llm", "temperature"),
"GRAYBOX_TOP_K": ("retrieval", "top_k"),
Expand Down Expand Up @@ -339,6 +350,7 @@ def load_config(path: str | None = None) -> Config:
llm=LLMConfig(**cfg["llm"]),
retrieval=RetrievalConfig(**cfg["retrieval"]),
embeddings=EmbeddingsConfig(**cfg.get("embeddings", {})),
prompts=PromptsConfig(**cfg.get("prompts", {})),
auto_refresh_summaries=cfg.get("auto_refresh_summaries", True),
raw=cfg,
config_path=candidate,
Expand Down
40 changes: 32 additions & 8 deletions graybox/prompts.py
Original file line number Diff line number Diff line change
Expand Up @@ -289,14 +289,38 @@

Question: {question}

Answer the question using only the context above. Cite the relevant source tag(s) inline
with markers like [] or [1][2] etc., right after each claim they support and then write
sources for each marker below the answer under header "Citations" like [1] project/atlas [2] people/aaryan
number of sources to cite can be 1 or more depending on how many sources were actually used and relevant for query.
If the context contains some relevant evidence but not a complete answer, answer
with the supported details and state what is not explicitly provided. Use the
exact refusal below only when the context contains no relevant information:
"I don't have enough information in the knowledge base to answer that."
Answer the question using ONLY the context above.

CITATION RULES:
- Every factual claim must be supported by the supplied context.
- Cite each factual claim with a numeric citation marker such as [1], [2], or [1][2].
- Use ONLY numeric citation markers. Never put source paths, source tags, filenames, or labels directly inside the answer.
- Place the citation marker immediately after the claim it supports.
- Use the smallest number of citations necessary to support the claim.
- Do not cite sources that were not actually used.
- Do not invent citation numbers.
- Do not use empty citation markers such as [].
- Do not write citations in any other format.

After the answer, provide the source mapping under exactly this header:

Citations

Use this exact format:

[1] source/path
[2] source/path

Only include citation numbers that actually appear in the answer.
Keep this section compact. Do not add explanations or commentary to it.

If the context contains some relevant evidence but not a complete answer,
answer with the supported details and state what is not explicitly provided.
Use the exact refusal below only when the context contains no relevant information: "I don't have enough information in the knowledge base to answer that."

====================

Return ONLY the answer and, when applicable, the compact "Citations" section.
"""

DIGEST_SYSTEM = """You are a workplace journal writer. Given a set of raw notes and the wiki pages
Expand Down
30 changes: 28 additions & 2 deletions graybox/retrieval.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,6 +92,8 @@ def _build_history_block(
) -> str:
lines = []
for turn in history:
if _is_refusal(turn.answer) or NO_EVIDENCE_MSG in turn.answer:
continue
lines.append(f"User: {turn.question}")
lines.append(f"Assistant: {turn.answer}")

Expand Down Expand Up @@ -243,10 +245,34 @@ def _build_context(hits: list[Hit], all_workspaces: bool = False) -> str:


def _system_prompt(cfg: Config) -> str:
"""Build the authoritative system prompt used by Ask/Chat answers.

RETRIEVAL_SYSTEM is always the foundation. Workspace context is added
next, and the optional answer-style configuration is appended last as a
presentation-only preference. The style block is explicitly delimited
and cannot change Gray Box's grounding contract.
"""
parts = [RETRIEVAL_SYSTEM]

ctx = workspace_context_block(cfg)
if ctx:
return f"{RETRIEVAL_SYSTEM}\n\nWorkspace context:\n{ctx}"
return RETRIEVAL_SYSTEM
parts.append(f"Workspace context:\n{ctx}")

style = (getattr(getattr(cfg, "prompts", None), "answer_style", "") or "").strip()
if style:
parts.append(
"User-defined answer preferences (presentation only):\n"
"<answer_style_preferences>\n"
f"{style}\n"
"</answer_style_preferences>\n\n"
"These preferences control only tone, verbosity, structure, formatting, "
"audience, personality, explanation style, and requested language. "
"They do not override Gray Box's grounding, citation, timestamp, "
"uncertainty, evidence, or refusal rules. If a preference conflicts "
"with those rules, the Gray Box retrieval rules take precedence."
)

return "\n\n".join(parts)


def _get_workspace_meta(hit: Any) -> tuple[str | None, str | None]:
Expand Down
4 changes: 2 additions & 2 deletions tests/test_ai_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,13 +82,13 @@ def test_non_retryable_exception_returns_error_dict(self, cfg):

def test_messages_built_with_system_and_user(self, cfg):
svc = AIService(cfg)
messages = svc._build_messages("system text", "user text")
messages = svc._build_messages("system text", "user text", None, None)
assert messages[0] == {"role": "system", "content": "system text"}
assert messages[1] == {"role": "user", "content": "user text"}

def test_default_system_prompt_when_none(self, cfg):
svc = AIService(cfg)
messages = svc._build_messages(None, "user text")
messages = svc._build_messages(None, "user text", None, None)
assert messages[0]["content"] == "You are a helpful assistant."

def test_streaming_call_returns_stream_wrapper(self, cfg):
Expand Down
Loading