Skip to content

Commit 69df748

Browse files
committed
Refactor item and lifecycle models to use Optional types and add validators for empty strings
- Updated the LifecycleUpdate and ItemType models to use Optional types for various fields, improving data handling. - Added a validator to convert empty strings to None for specific fields in both models, enhancing input validation. - Removed unnecessary checks in the benchmark API to streamline the upsert_benchmark function.
1 parent b32020e commit 69df748

4 files changed

Lines changed: 88 additions & 32 deletions

File tree

app/api/benchmark.py

Lines changed: 1 addition & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77

88
from models.base import User, CategoryBenchmark
99
from utils.auth import authenticate
10-
from utils.gear_lifecycle import get_all_benchmarks, DEFAULT_BENCHMARKS
10+
from utils.gear_lifecycle import get_all_benchmarks
1111

1212
logger = logging.getLogger(__name__)
1313

@@ -28,9 +28,6 @@ class BenchmarkUpdate(BaseModel):
2828

2929
@route.put("/{category_name}")
3030
def upsert_benchmark(category_name: str, payload: BenchmarkUpdate, user: User = Depends(authenticate)):
31-
if category_name not in DEFAULT_BENCHMARKS:
32-
raise HTTPException(400, f"Unknown category: {category_name}")
33-
3431
override = db.session.query(CategoryBenchmark).filter_by(
3532
user_id=user.id, category_name=category_name
3633
).first()

app/api/item.py

Lines changed: 20 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -3,8 +3,8 @@
33

44
from fastapi import APIRouter, Depends, HTTPException, File, UploadFile
55
from fastapi_sqlalchemy import db
6-
from pydantic import BaseModel
7-
from typing import List
6+
from pydantic import BaseModel, validator
7+
from typing import List, Optional
88
from io import StringIO
99
from sqlalchemy import or_, func
1010

@@ -38,15 +38,26 @@ class ItemType(BaseModel):
3838
product_url: str = None
3939
notes: str = None
4040

41-
acquired_date: str = None
42-
acquisition_type: str = None
43-
purchase_retailer: str = None
44-
condition: str = None
45-
status: str = None
46-
retired_date: str = None
47-
retired_reason: str = None
41+
acquired_date: Optional[str] = None
42+
acquisition_type: Optional[str] = None
43+
purchase_retailer: Optional[str] = None
44+
condition: Optional[str] = None
45+
status: Optional[str] = None
46+
retired_date: Optional[str] = None
47+
retired_reason: Optional[str] = None
4848
replaced_by_id: int = None
4949

50+
@validator(
51+
"acquired_date", "acquisition_type", "purchase_retailer",
52+
"condition", "status", "retired_date", "retired_reason",
53+
"product_url", "notes",
54+
pre=True, always=True,
55+
)
56+
def empty_str_to_none(cls, v):
57+
if isinstance(v, str) and v.strip() == "":
58+
return None
59+
return v
60+
5061

5162
def _find_catalog_product(session, brand_id: int, product_id: int, product_variant_id: int | None):
5263
q = session.query(CatalogProduct).filter(

app/api/item_lifecycle.py

Lines changed: 20 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,7 @@
33

44
from fastapi import APIRouter, Depends, HTTPException
55
from fastapi_sqlalchemy import db
6-
from pydantic import BaseModel
6+
from pydantic import BaseModel, validator
77
from typing import Optional
88

99
from models.base import User, Item, ItemLog
@@ -21,15 +21,25 @@
2121

2222

2323
class LifecycleUpdate(BaseModel):
24-
acquired_date: str = None
25-
acquisition_type: str = None
26-
purchase_retailer: str = None
27-
condition: str = None
28-
status: str = None
29-
retired_date: str = None
30-
retired_reason: str = None
24+
acquired_date: Optional[str] = None
25+
acquisition_type: Optional[str] = None
26+
purchase_retailer: Optional[str] = None
27+
condition: Optional[str] = None
28+
status: Optional[str] = None
29+
retired_date: Optional[str] = None
30+
retired_reason: Optional[str] = None
3131
replaced_by_id: int = None
3232

33+
@validator(
34+
"acquired_date", "acquisition_type", "purchase_retailer",
35+
"condition", "status", "retired_date", "retired_reason",
36+
pre=True, always=True,
37+
)
38+
def empty_str_to_none(cls, v):
39+
if isinstance(v, str) and v.strip() == "":
40+
return None
41+
return v
42+
3343

3444
@route.put("/{item_id}/lifecycle")
3545
def update_lifecycle(item_id: int, payload: LifecycleUpdate, user: User = Depends(authenticate)):
@@ -152,13 +162,15 @@ def get_replacement_score(item_id: int, user: User = Depends(authenticate)):
152162
)
153163

154164
benchmark = get_benchmark(db.session, user.id, category_name)
165+
is_default_fallback = benchmark.pop("is_default_fallback", False)
155166
score = replacement_score(item.acquired_date, item.condition, benchmark)
156167

157168
return {
158169
"item_id": item.id,
159170
"score": score,
160171
"category": category_name,
161172
"benchmark": benchmark,
173+
"is_default_fallback": is_default_fallback,
162174
"acquired_date": str(item.acquired_date) if item.acquired_date else None,
163175
"condition": item.condition,
164176
}

app/utils/gear_lifecycle.py

Lines changed: 47 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,8 @@
11
import datetime
22

3-
from models.base import CategoryBenchmark
3+
from sqlalchemy import or_
4+
5+
from models.base import Category, CategoryBenchmark
46

57
DEFAULT_BENCHMARKS = {
68
"Shelter": {"lifespan_years": 5, "expected_nights": 300},
@@ -31,38 +33,72 @@
3133
}
3234

3335

36+
def _apply_override(merged: dict, override, fields=_BENCHMARK_FIELDS) -> dict:
37+
for field in fields:
38+
val = getattr(override, field, None)
39+
if val is not None:
40+
merged[field] = float(val) if field != "distance_unit" else val
41+
return merged
42+
43+
3444
def get_benchmark(session, user_id: int, category_name: str) -> dict:
45+
is_known = category_name in DEFAULT_BENCHMARKS
3546
defaults = DEFAULT_BENCHMARKS.get(category_name, DEFAULT_BENCHMARKS["Miscellaneous"])
47+
3648
override = session.query(CategoryBenchmark).filter_by(
3749
user_id=user_id, category_name=category_name
3850
).first()
39-
if not override:
40-
return dict(defaults)
51+
4152
merged = dict(defaults)
42-
for field in _BENCHMARK_FIELDS:
43-
val = getattr(override, field, None)
44-
if val is not None:
45-
merged[field] = float(val) if field != "distance_unit" else val
53+
if override:
54+
_apply_override(merged, override)
55+
merged["is_default_fallback"] = False
56+
else:
57+
merged["is_default_fallback"] = not is_known
4658
return merged
4759

4860

61+
def _user_category_names(session, user_id: int) -> set[str]:
62+
rows = session.query(Category.name).filter(
63+
or_(Category.user_id == user_id, Category.user_id.is_(None))
64+
).all()
65+
return {r.name for r in rows if r.name}
66+
67+
4968
def get_all_benchmarks(session, user_id: int) -> dict[str, dict]:
5069
overrides = session.query(CategoryBenchmark).filter_by(user_id=user_id).all()
5170
override_map = {o.category_name: o for o in overrides}
5271

72+
user_cats = _user_category_names(session, user_id)
73+
5374
result = {}
75+
5476
for cat_name, defaults in DEFAULT_BENCHMARKS.items():
5577
merged = dict(defaults)
78+
merged["is_default_fallback"] = False
79+
override = override_map.get(cat_name)
80+
if override:
81+
_apply_override(merged, override)
82+
merged["has_override"] = True
83+
else:
84+
merged["has_override"] = False
85+
result[cat_name] = merged
86+
87+
misc_defaults = DEFAULT_BENCHMARKS["Miscellaneous"]
88+
for cat_name in user_cats:
89+
if cat_name in result:
90+
continue
91+
merged = dict(misc_defaults)
92+
merged["is_default_fallback"] = True
5693
override = override_map.get(cat_name)
5794
if override:
58-
for field in _BENCHMARK_FIELDS:
59-
val = getattr(override, field, None)
60-
if val is not None:
61-
merged[field] = float(val) if field != "distance_unit" else val
95+
_apply_override(merged, override)
6296
merged["has_override"] = True
97+
merged["is_default_fallback"] = False
6398
else:
6499
merged["has_override"] = False
65100
result[cat_name] = merged
101+
66102
return result
67103

68104

0 commit comments

Comments
 (0)