Skip to content

Commit ffdf914

Browse files
author
Dylan Huang
committed
works
1 parent 9214eae commit ffdf914

2 files changed

Lines changed: 45 additions & 3 deletions

File tree

eval_protocol/dataset_logger/sqlite_dataset_logger_adapter.py

Lines changed: 32 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,9 @@
11
import os
2-
from abc import ABC, abstractmethod
32
from typing import TYPE_CHECKING, List, Optional
43

4+
from peewee import CharField, Model, SqliteDatabase
5+
from playhouse.sqlite_ext import JSONField
6+
57
from eval_protocol.dataset_logger.dataset_logger import DatasetLogger
68
from eval_protocol.dataset_logger.directory_utils import find_eval_protocol_dir
79

@@ -13,9 +15,36 @@ class SqliteDatasetLoggerAdapter(DatasetLogger):
1315
def __init__(self, db_path: Optional[str] = None):
1416
eval_protocol_dir = find_eval_protocol_dir()
1517
self.db_path = os.path.join(eval_protocol_dir, "logs.db")
18+
db = SqliteDatabase(self.db_path)
19+
20+
class BaseModel(Model):
21+
class Meta:
22+
database = db
23+
24+
class EvaluationRow(BaseModel):
25+
row_id = CharField(unique=True)
26+
data = JSONField()
27+
28+
self.EvaluationRow = EvaluationRow
29+
30+
db.connect()
31+
db.create_tables([EvaluationRow])
1632

1733
def log(self, row: "EvaluationRow") -> None:
18-
pass
34+
row_id = row.input_metadata.row_id
35+
data = row.model_dump(exclude_none=True, mode="json")
36+
# if row_id already exists, update the row
37+
if self.EvaluationRow.select().where(self.EvaluationRow.row_id == row_id).exists():
38+
self.EvaluationRow.update(data=data).where(self.EvaluationRow.row_id == row_id).execute()
39+
else:
40+
self.EvaluationRow.create(row_id=row_id, data=data)
1941

2042
def read(self, row_id: Optional[str] = None) -> List["EvaluationRow"]:
21-
return []
43+
from eval_protocol.models import EvaluationRow
44+
45+
if row_id is None:
46+
query = self.EvaluationRow.select().dicts()
47+
else:
48+
query = self.EvaluationRow.select().dicts().where(self.EvaluationRow.row_id == row_id)
49+
results = list(query)
50+
return [EvaluationRow(**result["data"]) for result in results]
Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,13 @@
1+
from eval_protocol.dataset_logger.sqlite_dataset_logger_adapter import SqliteDatasetLoggerAdapter
2+
from eval_protocol.models import EvaluationRow, InputMetadata, Message
3+
4+
5+
def test_log_and_read():
6+
logger = SqliteDatasetLoggerAdapter()
7+
messages = [Message(role="user", content="Hello")]
8+
input_metadata = InputMetadata(row_id="1")
9+
row = EvaluationRow(input_metadata=input_metadata, messages=messages)
10+
logger.log(row)
11+
saved = logger.read(row_id="1")[0]
12+
assert row.messages == saved.messages
13+
assert row.input_metadata == saved.input_metadata

0 commit comments

Comments
 (0)