Skip to content

Commit 895b5ca

Browse files
committed
Resolve linter errors
1 parent 755ecc4 commit 895b5ca

9 files changed

Lines changed: 126 additions & 108 deletions

File tree

paperbanana/agents/visualizer.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -107,6 +107,7 @@ async def _generate_plot(
107107
full_description = description
108108
if raw_data:
109109
import json
110+
110111
full_description += f"\n\n## Raw Data\n```json\n{json.dumps(raw_data, indent=2)}\n```"
111112

112113
# Load and format the plot visualizer prompt template

paperbanana/cli.py

Lines changed: 48 additions & 51 deletions
Original file line numberDiff line numberDiff line change
@@ -25,33 +25,23 @@
2525

2626
@app.command()
2727
def generate(
28-
input: str = typer.Option(
29-
..., "--input", "-i", help="Path to methodology text file"
30-
),
28+
input: str = typer.Option(..., "--input", "-i", help="Path to methodology text file"),
3129
caption: str = typer.Option(
3230
..., "--caption", "-c", help="Figure caption / communicative intent"
3331
),
34-
output: Optional[str] = typer.Option(
35-
None, "--output", "-o", help="Output image path"
36-
),
32+
output: Optional[str] = typer.Option(None, "--output", "-o", help="Output image path"),
3733
vlm_provider: Optional[str] = typer.Option(
3834
None, "--vlm-provider", help="VLM provider (gemini)"
3935
),
40-
vlm_model: Optional[str] = typer.Option(
41-
None, "--vlm-model", help="VLM model name"
42-
),
36+
vlm_model: Optional[str] = typer.Option(None, "--vlm-model", help="VLM model name"),
4337
image_provider: Optional[str] = typer.Option(
4438
None, "--image-provider", help="Image gen provider"
4539
),
46-
image_model: Optional[str] = typer.Option(
47-
None, "--image-model", help="Image gen model name"
48-
),
40+
image_model: Optional[str] = typer.Option(None, "--image-model", help="Image gen model name"),
4941
iterations: Optional[int] = typer.Option(
5042
None, "--iterations", "-n", help="Refinement iterations"
5143
),
52-
config: Optional[str] = typer.Option(
53-
None, "--config", help="Path to config YAML file"
54-
),
44+
config: Optional[str] = typer.Option(None, "--config", help="Path to config YAML file"),
5545
):
5646
"""Generate a methodology diagram from a text description."""
5747
# Load source text
@@ -81,6 +71,7 @@ def generate(
8171
settings = Settings.from_yaml(config, **overrides)
8272
else:
8373
from dotenv import load_dotenv
74+
8475
load_dotenv()
8576
settings = Settings(**overrides)
8677

@@ -91,13 +82,15 @@ def generate(
9182
diagram_type=DiagramType.METHODOLOGY,
9283
)
9384

94-
console.print(Panel.fit(
95-
f"[bold]PaperBanana[/bold] - Generating Methodology Diagram\n\n"
96-
f"VLM: {settings.vlm_provider} / {settings.vlm_model}\n"
97-
f"Image: {settings.image_provider} / {settings.image_model}\n"
98-
f"Iterations: {settings.refinement_iterations}",
99-
border_style="blue",
100-
))
85+
console.print(
86+
Panel.fit(
87+
f"[bold]PaperBanana[/bold] - Generating Methodology Diagram\n\n"
88+
f"VLM: {settings.vlm_provider} / {settings.vlm_model}\n"
89+
f"Image: {settings.image_provider} / {settings.image_model}\n"
90+
f"Iterations: {settings.refinement_iterations}",
91+
border_style="blue",
92+
)
93+
)
10194

10295
# Run pipeline
10396
from paperbanana.core.pipeline import PaperBananaPipeline
@@ -135,8 +128,10 @@ def plot(
135128

136129
# Load data
137130
import json as json_mod
131+
138132
if data_path.suffix == ".csv":
139133
import pandas as pd
134+
140135
df = pd.read_csv(data_path)
141136
raw_data = df.to_dict(orient="records")
142137
source_context = (
@@ -149,6 +144,7 @@ def plot(
149144
source_context = f"JSON data:\n{json_mod.dumps(raw_data, indent=2)[:2000]}"
150145

151146
from dotenv import load_dotenv
147+
152148
load_dotenv()
153149

154150
settings = Settings(
@@ -163,12 +159,14 @@ def plot(
163159
raw_data={"data": raw_data},
164160
)
165161

166-
console.print(Panel.fit(
167-
f"[bold]PaperBanana[/bold] - Generating Statistical Plot\n\n"
168-
f"Data: {data_path.name}\n"
169-
f"Intent: {intent}",
170-
border_style="green",
171-
))
162+
console.print(
163+
Panel.fit(
164+
f"[bold]PaperBanana[/bold] - Generating Statistical Plot\n\n"
165+
f"Data: {data_path.name}\n"
166+
f"Intent: {intent}",
167+
border_style="green",
168+
)
169+
)
172170

173171
from paperbanana.core.pipeline import PaperBananaPipeline
174172

@@ -183,16 +181,19 @@ async def _run():
183181
@app.command()
184182
def setup():
185183
"""Interactive setup wizard — get generating in 2 minutes with FREE APIs."""
186-
console.print(Panel.fit(
187-
"[bold]Welcome to PaperBanana Setup[/bold]\n\n"
188-
"We'll set up FREE API keys so you can start generating diagrams.",
189-
border_style="yellow",
190-
))
184+
console.print(
185+
Panel.fit(
186+
"[bold]Welcome to PaperBanana Setup[/bold]\n\n"
187+
"We'll set up FREE API keys so you can start generating diagrams.",
188+
border_style="yellow",
189+
)
190+
)
191191

192192
console.print("\n[bold]Step 1: Google Gemini API Key[/bold] (FREE, no credit card)")
193193
console.print("This powers the AI agents that plan and critique your diagrams.\n")
194194

195195
import webbrowser
196+
196197
open_browser = Prompt.ask(
197198
"Open browser to get a free Gemini API key?",
198199
choices=["y", "n"],
@@ -220,18 +221,10 @@ def setup():
220221

221222
@app.command()
222223
def evaluate(
223-
generated: str = typer.Option(
224-
..., "--generated", "-g", help="Path to generated image"
225-
),
226-
context: str = typer.Option(
227-
..., "--context", help="Path to source context text file"
228-
),
229-
caption: str = typer.Option(
230-
..., "--caption", "-c", help="Figure caption"
231-
),
232-
reference: str = typer.Option(
233-
..., "--reference", "-r", help="Path to human reference image"
234-
),
224+
generated: str = typer.Option(..., "--generated", "-g", help="Path to generated image"),
225+
context: str = typer.Option(..., "--context", help="Path to source context text file"),
226+
caption: str = typer.Option(..., "--caption", "-c", help="Figure caption"),
227+
reference: str = typer.Option(..., "--reference", "-r", help="Path to human reference image"),
235228
vlm_provider: str = typer.Option(
236229
"gemini", "--vlm-provider", help="VLM provider for evaluation"
237230
),
@@ -252,10 +245,12 @@ def evaluate(
252245
context_text = Path(context).read_text(encoding="utf-8")
253246

254247
from dotenv import load_dotenv
248+
255249
load_dotenv()
256250

257251
settings = Settings(vlm_provider=vlm_provider)
258252
from paperbanana.providers.registry import ProviderRegistry
253+
259254
vlm = ProviderRegistry.create_vlm(settings)
260255

261256
judge = VLMJudge(vlm)
@@ -276,12 +271,14 @@ async def _run():
276271
result = getattr(scores, dim)
277272
dim_lines.append(f"{dim.capitalize():14s} {result.winner}")
278273

279-
console.print(Panel.fit(
280-
"[bold]Evaluation Results (Comparative)[/bold]\n\n"
281-
+ "\n".join(dim_lines)
282-
+ f"\n[bold]{'Overall':14s} {scores.overall_winner}[/bold]",
283-
border_style="cyan",
284-
))
274+
console.print(
275+
Panel.fit(
276+
"[bold]Evaluation Results (Comparative)[/bold]\n\n"
277+
+ "\n".join(dim_lines)
278+
+ f"\n[bold]{'Overall':14s} {scores.overall_winner}[/bold]",
279+
border_style="cyan",
280+
)
281+
)
285282

286283
for dim in dims:
287284
result = getattr(scores, dim)

paperbanana/core/pipeline.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -293,6 +293,7 @@ async def generate(self, input: GenerationInput) -> GenerationOutput:
293293

294294
# Copy final image to output location
295295
import shutil
296+
296297
shutil.copy2(final_image, final_output_path)
297298

298299
# Build metadata

paperbanana/core/types.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -90,7 +90,8 @@ class DimensionResult(BaseModel):
9090

9191
winner: str = Field(description="Model | Human | Both are good | Both are bad")
9292
score: float = Field(
93-
ge=0.0, le=100.0,
93+
ge=0.0,
94+
le=100.0,
9495
description="100 (Model wins), 0 (Human wins), 50 (Tie)",
9596
)
9697
reasoning: str = Field(default="", description="Comparison reasoning")
@@ -109,11 +110,10 @@ class EvaluationScore(BaseModel):
109110
conciseness: DimensionResult
110111
readability: DimensionResult
111112
aesthetics: DimensionResult
112-
overall_winner: str = Field(
113-
description="Hierarchical aggregation result"
114-
)
113+
overall_winner: str = Field(description="Hierarchical aggregation result")
115114
overall_score: float = Field(
116-
ge=0.0, le=100.0,
115+
ge=0.0,
116+
le=100.0,
117117
description="100 (Model wins), 0 (Human wins), 50 (Tie)",
118118
)
119119

paperbanana/evaluation/judge.py

Lines changed: 3 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -94,9 +94,7 @@ async def evaluate(
9494
overall_score=overall_score,
9595
)
9696

97-
def _load_eval_prompt(
98-
self, dimension: str, source_context: str, caption: str
99-
) -> str:
97+
def _load_eval_prompt(self, dimension: str, source_context: str, caption: str) -> str:
10098
"""Load evaluation prompt for a specific dimension."""
10199
prompt_path = self.prompt_dir / "evaluation" / f"{dimension}.txt"
102100
if not prompt_path.exists():
@@ -122,9 +120,7 @@ def _parse_result(self, response: str, dimension: str) -> DimensionResult:
122120
winner = "Both are good"
123121

124122
score = WINNER_SCORE_MAP.get(winner, 50.0)
125-
return DimensionResult(
126-
winner=winner, score=score, reasoning=reasoning
127-
)
123+
return DimensionResult(winner=winner, score=score, reasoning=reasoning)
128124
except (json.JSONDecodeError, ValueError, TypeError) as e:
129125
logger.warning(
130126
"Failed to parse evaluation response",
@@ -137,9 +133,7 @@ def _parse_result(self, response: str, dimension: str) -> DimensionResult:
137133
reasoning="Could not parse evaluation response.",
138134
)
139135

140-
def _hierarchical_aggregate(
141-
self, results: dict[str, DimensionResult]
142-
) -> str:
136+
def _hierarchical_aggregate(self, results: dict[str, DimensionResult]) -> str:
143137
"""Apply hierarchical aggregation per paper Section 4.2.
144138
145139
Primary dimensions (Faithfulness + Readability) take precedence.

paperbanana/providers/registry.py

Lines changed: 3 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -27,18 +27,13 @@ def create_vlm(settings: Settings) -> VLMProvider:
2727
model=settings.vlm_model,
2828
)
2929
else:
30-
raise ValueError(
31-
f"Unknown VLM provider: {provider}. "
32-
f"Available: gemini"
33-
)
30+
raise ValueError(f"Unknown VLM provider: {provider}. Available: gemini")
3431

3532
@staticmethod
3633
def create_image_gen(settings: Settings) -> ImageGenProvider:
3734
"""Create an image generation provider based on settings."""
3835
provider = settings.image_provider.lower()
39-
logger.info(
40-
"Creating image gen provider", provider=provider, model=settings.image_model
41-
)
36+
logger.info("Creating image gen provider", provider=provider, model=settings.image_model)
4237

4338
if provider == "google_imagen":
4439
from paperbanana.providers.image_gen.google_imagen import GoogleImagenGen
@@ -48,7 +43,4 @@ def create_image_gen(settings: Settings) -> ImageGenProvider:
4843
model=settings.image_model,
4944
)
5045
else:
51-
raise ValueError(
52-
f"Unknown image provider: {provider}. "
53-
f"Available: google_imagen"
54-
)
46+
raise ValueError(f"Unknown image provider: {provider}. Available: google_imagen")

tests/test_agents/test_retriever.py

Lines changed: 17 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -19,8 +19,15 @@ class MockVLM:
1919
def __init__(self, response: str = ""):
2020
self._response = response
2121

22-
async def generate(self, prompt, images=None, system_prompt=None,
23-
temperature=1.0, max_tokens=4096, response_format=None):
22+
async def generate(
23+
self,
24+
prompt,
25+
images=None,
26+
system_prompt=None,
27+
temperature=1.0,
28+
max_tokens=4096,
29+
response_format=None,
30+
):
2431
return self._response
2532

2633
def is_available(self):
@@ -76,13 +83,15 @@ async def test_retriever_empty_candidates():
7683
@pytest.mark.asyncio
7784
async def test_retriever_parses_vlm_response():
7885
"""Test that retriever correctly parses VLM JSON response."""
79-
response = json.dumps({
80-
"selected_ids": ["ref_001", "ref_003"],
81-
"reasoning": {
82-
"ref_001": "Relevant because...",
83-
"ref_003": "Relevant because...",
86+
response = json.dumps(
87+
{
88+
"selected_ids": ["ref_001", "ref_003"],
89+
"reasoning": {
90+
"ref_001": "Relevant because...",
91+
"ref_003": "Relevant because...",
92+
},
8493
}
85-
})
94+
)
8695

8796
vlm = MockVLM(response=response)
8897
agent = RetrieverAgent(vlm)

0 commit comments

Comments
 (0)