refactor domain into a module
This commit is contained in:
parent
0a868fceb5
commit
63a568490e
8 changed files with 161 additions and 148 deletions
0
open_notebook/domain/__init__.py
Normal file
0
open_notebook/domain/__init__.py
Normal file
147
open_notebook/domain/base.py
Normal file
147
open_notebook/domain/base.py
Normal file
|
|
@ -0,0 +1,147 @@
|
||||||
|
from datetime import datetime
|
||||||
|
from typing import Any, ClassVar, Dict, List, Optional, Type, TypeVar
|
||||||
|
|
||||||
|
from loguru import logger
|
||||||
|
from pydantic import BaseModel, ValidationError, field_validator
|
||||||
|
|
||||||
|
from open_notebook.exceptions import (
|
||||||
|
DatabaseOperationError,
|
||||||
|
InvalidInputError,
|
||||||
|
NotFoundError,
|
||||||
|
)
|
||||||
|
from open_notebook.repository import (
|
||||||
|
repo_create,
|
||||||
|
repo_delete,
|
||||||
|
repo_query,
|
||||||
|
repo_relate,
|
||||||
|
repo_update,
|
||||||
|
)
|
||||||
|
|
||||||
|
T = TypeVar("T", bound="ObjectModel")
|
||||||
|
|
||||||
|
|
||||||
|
class ObjectModel(BaseModel):
|
||||||
|
id: Optional[str] = None
|
||||||
|
table_name: ClassVar[str] = ""
|
||||||
|
created: Optional[datetime] = None
|
||||||
|
updated: Optional[datetime] = None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_all(cls: Type[T], order_by=None) -> List[T]:
|
||||||
|
try:
|
||||||
|
if order_by:
|
||||||
|
order = f" ORDER BY {order_by}"
|
||||||
|
else:
|
||||||
|
order = ""
|
||||||
|
result = repo_query(f"SELECT * FROM {cls.table_name} {order}")
|
||||||
|
objects = []
|
||||||
|
for obj in result:
|
||||||
|
try:
|
||||||
|
objects.append(cls(**obj))
|
||||||
|
except Exception as e:
|
||||||
|
logger.critical(f"Error creating object: {str(e)}")
|
||||||
|
|
||||||
|
return objects
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error fetching all {cls.table_name}: {str(e)}")
|
||||||
|
logger.exception(e)
|
||||||
|
raise DatabaseOperationError(e)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get(cls: Type[T], id: str) -> Optional[T]:
|
||||||
|
if not id:
|
||||||
|
raise InvalidInputError("ID cannot be empty")
|
||||||
|
try:
|
||||||
|
result = repo_query(f"SELECT * FROM {id}")
|
||||||
|
if result:
|
||||||
|
return cls(**result[0])
|
||||||
|
return None
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error fetching {cls.table_name} with id {id}: {str(e)}")
|
||||||
|
logger.exception(e)
|
||||||
|
raise NotFoundError(f"{cls.table_name} with id {id} not found")
|
||||||
|
|
||||||
|
def needs_embedding(self) -> bool:
|
||||||
|
return False
|
||||||
|
|
||||||
|
def get_embedding_content(self) -> Optional[str]:
|
||||||
|
return None
|
||||||
|
|
||||||
|
def save(self) -> None:
|
||||||
|
from open_notebook.config import EMBEDDING_MODEL
|
||||||
|
|
||||||
|
try:
|
||||||
|
logger.debug(f"Validating {self.__class__.__name__}")
|
||||||
|
self.model_validate(self.model_dump(), strict=True)
|
||||||
|
data = self._prepare_save_data()
|
||||||
|
data["updated"] = datetime.now().isoformat()
|
||||||
|
|
||||||
|
if self.needs_embedding():
|
||||||
|
embedding_content = self.get_embedding_content()
|
||||||
|
if embedding_content:
|
||||||
|
data["embedding"] = EMBEDDING_MODEL.embed(embedding_content)
|
||||||
|
|
||||||
|
if self.id is None:
|
||||||
|
data["created"] = datetime.now().isoformat()
|
||||||
|
logger.debug("Creating new record")
|
||||||
|
repo_result = repo_create(self.__class__.table_name, data)
|
||||||
|
else:
|
||||||
|
logger.debug(f"Updating record with id {self.id}")
|
||||||
|
repo_result = repo_update(self.id, data)
|
||||||
|
|
||||||
|
# Update the current instance with the result
|
||||||
|
for key, value in repo_result[0].items():
|
||||||
|
if hasattr(self, key):
|
||||||
|
if isinstance(getattr(self, key), BaseModel):
|
||||||
|
setattr(self, key, type(getattr(self, key))(**value))
|
||||||
|
else:
|
||||||
|
setattr(self, key, value)
|
||||||
|
|
||||||
|
except ValidationError as e:
|
||||||
|
logger.error(f"Validation failed: {e}")
|
||||||
|
raise
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error saving record: {e}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error saving {self.__class__.table_name}: {str(e)}")
|
||||||
|
logger.exception(e)
|
||||||
|
raise DatabaseOperationError(e)
|
||||||
|
|
||||||
|
def _prepare_save_data(self) -> Dict[str, Any]:
|
||||||
|
data = self.model_dump()
|
||||||
|
# del data["created"]
|
||||||
|
# del data["updated"]
|
||||||
|
return {key: value for key, value in data.items() if value is not None}
|
||||||
|
|
||||||
|
def delete(self) -> bool:
|
||||||
|
if self.id is None:
|
||||||
|
raise InvalidInputError("Cannot delete object without an ID")
|
||||||
|
try:
|
||||||
|
logger.debug(f"Deleting record with id {self.id}")
|
||||||
|
return repo_delete(self.id)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(
|
||||||
|
f"Error deleting {self.__class__.table_name} with id {self.id}: {str(e)}"
|
||||||
|
)
|
||||||
|
raise DatabaseOperationError(
|
||||||
|
f"Failed to delete {self.__class__.table_name}"
|
||||||
|
)
|
||||||
|
|
||||||
|
def relate(self, relationship: str, target_id: str) -> Any:
|
||||||
|
if not relationship or not target_id or not self.id:
|
||||||
|
raise InvalidInputError("Relationship and target ID must be provided")
|
||||||
|
try:
|
||||||
|
return repo_relate(self.id, relationship, target_id)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error creating relationship: {str(e)}")
|
||||||
|
logger.exception(e)
|
||||||
|
raise DatabaseOperationError(e)
|
||||||
|
|
||||||
|
@field_validator("created", "updated", mode="before")
|
||||||
|
@classmethod
|
||||||
|
def parse_datetime(cls, value):
|
||||||
|
if isinstance(value, str):
|
||||||
|
return datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||||
|
return value
|
||||||
|
|
@ -1,153 +1,23 @@
|
||||||
import os
|
import os
|
||||||
from datetime import datetime
|
from typing import Any, ClassVar, Dict, List, Literal, Optional
|
||||||
from typing import Any, ClassVar, Dict, List, Literal, Optional, Type, TypeVar
|
|
||||||
|
|
||||||
from langchain_core.runnables.config import RunnableConfig
|
from langchain_core.runnables.config import RunnableConfig
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
from pydantic import BaseModel, Field, ValidationError, field_validator
|
from pydantic import BaseModel, Field, field_validator
|
||||||
|
|
||||||
|
from open_notebook.config import EMBEDDING_MODEL
|
||||||
|
from open_notebook.domain.base import ObjectModel
|
||||||
from open_notebook.exceptions import (
|
from open_notebook.exceptions import (
|
||||||
DatabaseOperationError,
|
DatabaseOperationError,
|
||||||
InvalidInputError,
|
InvalidInputError,
|
||||||
NotFoundError,
|
|
||||||
)
|
)
|
||||||
from open_notebook.graphs.multipattern import graph as pattern_graph
|
from open_notebook.graphs.multipattern import graph as pattern_graph
|
||||||
from open_notebook.graphs.recursive_toc import graph as toc_graph
|
from open_notebook.graphs.recursive_toc import graph as toc_graph
|
||||||
from open_notebook.repository import (
|
from open_notebook.repository import (
|
||||||
repo_create,
|
repo_create,
|
||||||
repo_delete,
|
|
||||||
repo_query,
|
repo_query,
|
||||||
repo_relate,
|
|
||||||
repo_update,
|
|
||||||
)
|
)
|
||||||
from open_notebook.utils import get_embedding, split_text, surreal_clean
|
from open_notebook.utils import split_text, surreal_clean
|
||||||
|
|
||||||
T = TypeVar("T", bound="ObjectModel")
|
|
||||||
|
|
||||||
|
|
||||||
class ObjectModel(BaseModel):
|
|
||||||
id: Optional[str] = None
|
|
||||||
table_name: ClassVar[str] = ""
|
|
||||||
created: Optional[datetime] = None
|
|
||||||
updated: Optional[datetime] = None
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def get_all(cls: Type[T], order_by=None) -> List[T]:
|
|
||||||
try:
|
|
||||||
if order_by:
|
|
||||||
order = f" ORDER BY {order_by}"
|
|
||||||
else:
|
|
||||||
order = ""
|
|
||||||
result = repo_query(f"SELECT * FROM {cls.table_name} {order}")
|
|
||||||
objects = []
|
|
||||||
for obj in result:
|
|
||||||
try:
|
|
||||||
objects.append(cls(**obj))
|
|
||||||
except Exception as e:
|
|
||||||
logger.critical(f"Error creating object: {str(e)}")
|
|
||||||
|
|
||||||
return objects
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error fetching all {cls.table_name}: {str(e)}")
|
|
||||||
logger.exception(e)
|
|
||||||
raise DatabaseOperationError(e)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def get(cls: Type[T], id: str) -> Optional[T]:
|
|
||||||
if not id:
|
|
||||||
raise InvalidInputError("ID cannot be empty")
|
|
||||||
try:
|
|
||||||
result = repo_query(f"SELECT * FROM {id}")
|
|
||||||
if result:
|
|
||||||
return cls(**result[0])
|
|
||||||
return None
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error fetching {cls.table_name} with id {id}: {str(e)}")
|
|
||||||
logger.exception(e)
|
|
||||||
raise NotFoundError(f"{cls.table_name} with id {id} not found")
|
|
||||||
|
|
||||||
def needs_embedding(self) -> bool:
|
|
||||||
return False
|
|
||||||
|
|
||||||
def get_embedding_content(self) -> Optional[str]:
|
|
||||||
return None
|
|
||||||
|
|
||||||
def save(self) -> None:
|
|
||||||
try:
|
|
||||||
logger.debug(f"Validating {self.__class__.__name__}")
|
|
||||||
self.model_validate(self.model_dump(), strict=True)
|
|
||||||
data = self._prepare_save_data()
|
|
||||||
data["updated"] = datetime.now().isoformat()
|
|
||||||
|
|
||||||
if self.needs_embedding():
|
|
||||||
embedding_content = self.get_embedding_content()
|
|
||||||
if embedding_content:
|
|
||||||
data["embedding"] = get_embedding(embedding_content)
|
|
||||||
|
|
||||||
if self.id is None:
|
|
||||||
data["created"] = datetime.now().isoformat()
|
|
||||||
logger.debug("Creating new record")
|
|
||||||
repo_result = repo_create(self.__class__.table_name, data)
|
|
||||||
else:
|
|
||||||
logger.debug(f"Updating record with id {self.id}")
|
|
||||||
repo_result = repo_update(self.id, data)
|
|
||||||
|
|
||||||
# Update the current instance with the result
|
|
||||||
for key, value in repo_result[0].items():
|
|
||||||
if hasattr(self, key):
|
|
||||||
if isinstance(getattr(self, key), BaseModel):
|
|
||||||
setattr(self, key, type(getattr(self, key))(**value))
|
|
||||||
else:
|
|
||||||
setattr(self, key, value)
|
|
||||||
|
|
||||||
except ValidationError as e:
|
|
||||||
logger.error(f"Validation failed: {e}")
|
|
||||||
raise
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error saving record: {e}")
|
|
||||||
raise
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error saving {self.__class__.table_name}: {str(e)}")
|
|
||||||
logger.exception(e)
|
|
||||||
raise DatabaseOperationError(e)
|
|
||||||
|
|
||||||
def _prepare_save_data(self) -> Dict[str, Any]:
|
|
||||||
data = self.model_dump()
|
|
||||||
# del data["created"]
|
|
||||||
# del data["updated"]
|
|
||||||
return {key: value for key, value in data.items() if value is not None}
|
|
||||||
|
|
||||||
def delete(self) -> bool:
|
|
||||||
if self.id is None:
|
|
||||||
raise InvalidInputError("Cannot delete object without an ID")
|
|
||||||
try:
|
|
||||||
logger.debug(f"Deleting record with id {self.id}")
|
|
||||||
return repo_delete(self.id)
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(
|
|
||||||
f"Error deleting {self.__class__.table_name} with id {self.id}: {str(e)}"
|
|
||||||
)
|
|
||||||
raise DatabaseOperationError(
|
|
||||||
f"Failed to delete {self.__class__.table_name}"
|
|
||||||
)
|
|
||||||
|
|
||||||
def relate(self, relationship: str, target_id: str) -> Any:
|
|
||||||
if not relationship or not target_id or not self.id:
|
|
||||||
raise InvalidInputError("Relationship and target ID must be provided")
|
|
||||||
try:
|
|
||||||
return repo_relate(self.id, relationship, target_id)
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error creating relationship: {str(e)}")
|
|
||||||
logger.exception(e)
|
|
||||||
raise DatabaseOperationError(e)
|
|
||||||
|
|
||||||
@field_validator("created", "updated", mode="before")
|
|
||||||
@classmethod
|
|
||||||
def parse_datetime(cls, value):
|
|
||||||
if isinstance(value, str):
|
|
||||||
return datetime.fromisoformat(value.replace("Z", "+00:00"))
|
|
||||||
return value
|
|
||||||
|
|
||||||
|
|
||||||
class Notebook(ObjectModel):
|
class Notebook(ObjectModel):
|
||||||
|
|
@ -288,7 +158,7 @@ class Source(ObjectModel):
|
||||||
"source": {self.id},
|
"source": {self.id},
|
||||||
"order": {i},
|
"order": {i},
|
||||||
"content": $content,
|
"content": $content,
|
||||||
"embedding": {get_embedding(chunk)},
|
"embedding": {EMBEDDING_MODEL.embed(chunk)},
|
||||||
}};""",
|
}};""",
|
||||||
{"content": surreal_clean(chunk)},
|
{"content": surreal_clean(chunk)},
|
||||||
)
|
)
|
||||||
|
|
@ -322,7 +192,7 @@ class Source(ObjectModel):
|
||||||
if not insight_type or not content:
|
if not insight_type or not content:
|
||||||
raise InvalidInputError("Insight type and content must be provided")
|
raise InvalidInputError("Insight type and content must be provided")
|
||||||
try:
|
try:
|
||||||
embedding = get_embedding(content)
|
embedding = EMBEDDING_MODEL.embed(content)
|
||||||
return repo_query(
|
return repo_query(
|
||||||
f"""
|
f"""
|
||||||
CREATE source_insight CONTENT {{
|
CREATE source_insight CONTENT {{
|
||||||
|
|
@ -396,9 +266,7 @@ class Note(ObjectModel):
|
||||||
return self.content
|
return self.content
|
||||||
|
|
||||||
|
|
||||||
def text_search(
|
def text_search(keyword: str, results: int, source: bool = True, note: bool = True):
|
||||||
keyword: str, results: int, source: bool = True, note: bool = True
|
|
||||||
) -> List[Dict[str, Any]]:
|
|
||||||
if not keyword:
|
if not keyword:
|
||||||
raise InvalidInputError("Search keyword cannot be empty")
|
raise InvalidInputError("Search keyword cannot be empty")
|
||||||
try:
|
try:
|
||||||
|
|
@ -415,9 +283,7 @@ def text_search(
|
||||||
raise DatabaseOperationError("Failed to perform text search")
|
raise DatabaseOperationError("Failed to perform text search")
|
||||||
|
|
||||||
|
|
||||||
def vector_search(
|
def vector_search(keyword: str, results: int, source: bool = True, note: bool = True):
|
||||||
keyword: str, results: int, source: bool = True, note: bool = True
|
|
||||||
) -> List[Dict[str, Any]]:
|
|
||||||
if not keyword:
|
if not keyword:
|
||||||
raise InvalidInputError("Search keyword cannot be empty")
|
raise InvalidInputError("Search keyword cannot be empty")
|
||||||
try:
|
try:
|
||||||
|
|
@ -34,7 +34,7 @@ def repo_query(query_str: str, vars: Optional[Dict[str, Any]] = None):
|
||||||
result = connection.query(query_str, vars)
|
result = connection.query(query_str, vars)
|
||||||
return result
|
return result
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
# logger.debug(f"Query: {query_str}, Variables: {vars}")
|
logger.critical(f"Query: {query_str}, Variables: {vars}")
|
||||||
logger.exception(e)
|
logger.exception(e)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,7 +1,7 @@
|
||||||
import streamlit as st
|
import streamlit as st
|
||||||
from humanize import naturaltime
|
from humanize import naturaltime
|
||||||
|
|
||||||
from open_notebook.domain import Notebook
|
from open_notebook.domain.notebook import Notebook
|
||||||
from stream_app.chat import chat_sidebar
|
from stream_app.chat import chat_sidebar
|
||||||
from stream_app.note import add_note, note_card
|
from stream_app.note import add_note, note_card
|
||||||
from stream_app.source import add_source, source_card
|
from stream_app.source import add_source, source_card
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,6 @@
|
||||||
import streamlit as st
|
import streamlit as st
|
||||||
|
|
||||||
from open_notebook.domain import text_search, vector_search
|
from open_notebook.domain.notebook import text_search, vector_search
|
||||||
from open_notebook.utils import get_embedding
|
from open_notebook.utils import get_embedding
|
||||||
from stream_app.note import note_list_item
|
from stream_app.note import note_list_item
|
||||||
from stream_app.source import source_list_item
|
from stream_app.source import source_list_item
|
||||||
|
|
|
||||||
|
|
@ -1,7 +1,7 @@
|
||||||
import streamlit as st
|
import streamlit as st
|
||||||
from langchain_core.runnables import RunnableConfig
|
from langchain_core.runnables import RunnableConfig
|
||||||
|
|
||||||
from open_notebook.domain import Note, Source
|
from open_notebook.domain.notebook import Note, Source
|
||||||
from open_notebook.graphs.chat import graph as chat_graph
|
from open_notebook.graphs.chat import graph as chat_graph
|
||||||
from open_notebook.plugins.podcasts import PodcastConfig
|
from open_notebook.plugins.podcasts import PodcastConfig
|
||||||
from open_notebook.utils import token_count
|
from open_notebook.utils import token_count
|
||||||
|
|
|
||||||
|
|
@ -3,7 +3,7 @@ from humanize import naturaltime
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
from streamlit_monaco import st_monaco # type: ignore
|
from streamlit_monaco import st_monaco # type: ignore
|
||||||
|
|
||||||
from open_notebook.domain import Note
|
from open_notebook.domain.notebook import Note
|
||||||
from open_notebook.graphs.multipattern import graph as pattern_graph
|
from open_notebook.graphs.multipattern import graph as pattern_graph
|
||||||
from open_notebook.utils import surreal_clean
|
from open_notebook.utils import surreal_clean
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue