1+ from __future__ import annotations
2+ from typing import overload , Union , List , Optional
3+ from dataclasses import dataclass , is_dataclass , field , fields
14from pythonicMySQL .datatypes .mysqltypes import MySQLType , INT
2- from dataclasses import dataclass , field
5+
6+
7+ ID_COLUMN = field (default = None , init = True , metadata = {"name" : "id" , "mysql_type" : INT (11 ), "unsigned" : True ,
8+ "primary" : True , "auto_increment" : True })
39
410
511@dataclass (frozen = True )
@@ -17,16 +23,61 @@ class Column:
1723 @property
1824 def description (self ) -> str :
1925 query_str = f"`{ self .name } ` { self .mysql_type .description } "
20- if self .unsigned :
21- query_str += " unsigned"
22- if not self .null :
23- query_str += " NOT NULL"
24- if self .default is not None :
25- query_str += f" Default { self .default } "
26- if self .auto_increment :
27- query_str += " AUTO_INCREMENT"
26+ if self .unsigned : query_str += " unsigned"
27+ if not self .null : query_str += " NOT NULL"
28+ if self .default is not None : query_str += f" Default { self .default } "
29+ if self .auto_increment : query_str += " AUTO_INCREMENT"
2830 if self .unique :
2931 query_str += " UNIQUE"
3032 if self .primary :
3133 query_str += " PRIMARY KEY"
3234 return query_str
35+
36+
37+ def column (mysql_type : Union [MySQLType , MySQLType .__class__ ], * flags : dict , default = None , unsigned : bool = False ,
38+ null : bool = False , unique : bool = False ):
39+ if isinstance (mysql_type , type ):
40+ mysql_type = mysql_type ()
41+ metadata = {
42+ "mysql_type" : mysql_type ,
43+ "default" : default ,
44+ "unsigned" : unsigned ,
45+ "null" : null ,
46+ "unique" : unique
47+ }
48+ for flag in flags :
49+ metadata = {** metadata , ** flag }
50+ if default is not None :
51+ return field (default = default , metadata = metadata )
52+ else :
53+ return field (metadata = metadata )
54+
55+
56+ @overload
57+ def columns (mysql_object : type , attribute : str ) -> Optional [Column ]: ...
58+
59+
60+ @overload
61+ def columns (mysql_object : type ) -> List [Column ]: ...
62+
63+
64+ def columns (mysql_object : type , attribute : Optional [str ] = None ) -> Union [List [Column ], Column , None ]:
65+ if not is_dataclass (mysql_object ):
66+ raise ValueError ("MySQL objects need to be a dataclass" )
67+ columns_ = []
68+ if attribute is None :
69+ fields_ = fields (mysql_object )
70+ else :
71+ fields_ = [item for item in fields (mysql_object ) if item .name == attribute ]
72+ for field_ in fields_ :
73+ if "mysql_type" in dict (field_ .metadata ).keys ():
74+ metadata = {** {"name" : field_ .name }, ** field_ .metadata }
75+ columns_ .append (Column (** metadata ))
76+ if attribute is not None and len (columns_ ) == 1 :
77+ return columns_ [0 ]
78+ elif attribute is not None and len (columns_ ) == 0 :
79+ return None
80+ elif attribute is None :
81+ return sorted (columns_ , key = lambda i : i .primary , reverse = True )
82+ else :
83+ raise KeyError
0 commit comments