@@ -19,7 +19,9 @@ use crate::semantic_index::{
1919use crate :: types:: bound_super:: BoundSuperError ;
2020use crate :: types:: constraints:: { ConstraintSet , IteratorConstraintsExtension } ;
2121use crate :: types:: context:: InferContext ;
22- use crate :: types:: diagnostic:: { INVALID_TYPE_ALIAS_TYPE , SUPER_CALL_IN_NAMED_TUPLE_METHOD } ;
22+ use crate :: types:: diagnostic:: {
23+ INVALID_TYPE_ALIAS_TYPE , SUPER_CALL_IN_NAMED_TUPLE_METHOD , UNSOUND_DATACLASS_METHOD_OVERRIDE ,
24+ } ;
2325use crate :: types:: enums:: enum_metadata;
2426use crate :: types:: function:: {
2527 DataclassTransformerFlags , DataclassTransformerParams , KnownFunction ,
@@ -1936,6 +1938,69 @@ impl<'db> ClassLiteral<'db> {
19361938 Some ( typed_dict_params_from_class_def ( class_stmt) )
19371939 }
19381940
1941+ /// Returns dataclass params for this class, sourced from both dataclass params and dataclass
1942+ /// transform params
1943+ fn merged_dataclass_params (
1944+ self ,
1945+ db : & ' db dyn Db ,
1946+ field_policy : CodeGeneratorKind < ' db > ,
1947+ ) -> ( Option < DataclassParams < ' db > > , Option < DataclassParams < ' db > > ) {
1948+ let dataclass_params = self . dataclass_params ( db) ;
1949+
1950+ let mut transformer_params =
1951+ if let CodeGeneratorKind :: DataclassLike ( Some ( transformer_params) ) = field_policy {
1952+ Some ( DataclassParams :: from_transformer_params (
1953+ db,
1954+ transformer_params,
1955+ ) )
1956+ } else {
1957+ None
1958+ } ;
1959+
1960+ // Dataclass transformer flags can be overwritten using class arguments.
1961+ if let Some ( transformer_params) = transformer_params. as_mut ( ) {
1962+ if let Some ( class_def) = self . definition ( db) . kind ( db) . as_class ( ) {
1963+ let module = parsed_module ( db, self . file ( db) ) . load ( db) ;
1964+
1965+ if let Some ( arguments) = & class_def. node ( & module) . arguments {
1966+ let mut flags = transformer_params. flags ( db) ;
1967+
1968+ for keyword in & arguments. keywords {
1969+ if let Some ( arg_name) = & keyword. arg {
1970+ if let Some ( is_set) =
1971+ keyword. value . as_boolean_literal_expr ( ) . map ( |b| b. value )
1972+ {
1973+ for ( flag_name, flag) in DATACLASS_FLAGS {
1974+ if arg_name. as_str ( ) == * flag_name {
1975+ flags. set ( * flag, is_set) ;
1976+ }
1977+ }
1978+ }
1979+ }
1980+ }
1981+
1982+ * transformer_params =
1983+ DataclassParams :: new ( db, flags, transformer_params. field_specifiers ( db) ) ;
1984+ }
1985+ }
1986+ }
1987+
1988+ ( dataclass_params, transformer_params)
1989+ }
1990+
1991+ /// Checks if the given dataclass parameter flag is set for this class.
1992+ /// This checks both the `dataclass_params` and `transformer_params`.
1993+ fn has_dataclass_param (
1994+ self ,
1995+ db : & ' db dyn Db ,
1996+ field_policy : CodeGeneratorKind < ' db > ,
1997+ param : DataclassFlags ,
1998+ ) -> bool {
1999+ let ( dataclass_params, transformer_params) = self . merged_dataclass_params ( db, field_policy) ;
2000+ dataclass_params. is_some_and ( |params| params. flags ( db) . contains ( param) )
2001+ || transformer_params. is_some_and ( |params| params. flags ( db) . contains ( param) )
2002+ }
2003+
19392004 /// Return the explicit `metaclass` of this class, if one is defined.
19402005 ///
19412006 /// ## Note
@@ -2345,57 +2410,8 @@ impl<'db> ClassLiteral<'db> {
23452410 inherited_generic_context : Option < GenericContext < ' db > > ,
23462411 name : & str ,
23472412 ) -> Option < Type < ' db > > {
2348- let dataclass_params = self . dataclass_params ( db) ;
2349-
23502413 let field_policy = CodeGeneratorKind :: from_class ( db, self , specialization) ?;
23512414
2352- let mut transformer_params =
2353- if let CodeGeneratorKind :: DataclassLike ( Some ( transformer_params) ) = field_policy {
2354- Some ( DataclassParams :: from_transformer_params (
2355- db,
2356- transformer_params,
2357- ) )
2358- } else {
2359- None
2360- } ;
2361-
2362- // Dataclass transformer flags can be overwritten using class arguments.
2363- // TODO this should be done more generally, not just in `own_synthesized_member`, so that
2364- // `dataclass_params` always reflects the transformer params.
2365- if let Some ( transformer_params) = transformer_params. as_mut ( ) {
2366- if let Some ( class_def) = self . definition ( db) . kind ( db) . as_class ( ) {
2367- let module = parsed_module ( db, self . file ( db) ) . load ( db) ;
2368-
2369- if let Some ( arguments) = & class_def. node ( & module) . arguments {
2370- let mut flags = transformer_params. flags ( db) ;
2371-
2372- for keyword in & arguments. keywords {
2373- if let Some ( arg_name) = & keyword. arg {
2374- if let Some ( is_set) =
2375- keyword. value . as_boolean_literal_expr ( ) . map ( |b| b. value )
2376- {
2377- for ( flag_name, flag) in DATACLASS_FLAGS {
2378- if arg_name. as_str ( ) == * flag_name {
2379- flags. set ( * flag, is_set) ;
2380- }
2381- }
2382- }
2383- }
2384- }
2385-
2386- * transformer_params =
2387- DataclassParams :: new ( db, flags, transformer_params. field_specifiers ( db) ) ;
2388- }
2389- }
2390- }
2391-
2392- let has_dataclass_param = |param| {
2393- dataclass_params. is_some_and ( |params| params. flags ( db) . contains ( param) )
2394- // TODO if we were correctly initializing `dataclass_params` from the
2395- // transformer params, this fallback shouldn't be needed here.
2396- || transformer_params. is_some_and ( |params| params. flags ( db) . contains ( param) )
2397- } ;
2398-
23992415 let instance_ty =
24002416 Type :: instance ( db, self . apply_optional_specialization ( db, specialization) ) ;
24012417
@@ -2514,7 +2530,7 @@ impl<'db> ClassLiteral<'db> {
25142530
25152531 match ( field_policy, name) {
25162532 ( CodeGeneratorKind :: DataclassLike ( _) , "__init__" ) => {
2517- if !has_dataclass_param ( DataclassFlags :: INIT ) {
2533+ if !self . has_dataclass_param ( db , field_policy , DataclassFlags :: INIT ) {
25182534 return None ;
25192535 }
25202536
@@ -2529,7 +2545,7 @@ impl<'db> ClassLiteral<'db> {
25292545 signature_from_fields ( vec ! [ cls_parameter] , Some ( Type :: none ( db) ) )
25302546 }
25312547 ( CodeGeneratorKind :: DataclassLike ( _) , "__lt__" | "__le__" | "__gt__" | "__ge__" ) => {
2532- if !has_dataclass_param ( DataclassFlags :: ORDER ) {
2548+ if !self . has_dataclass_param ( db , field_policy , DataclassFlags :: ORDER ) {
25332549 return None ;
25342550 }
25352551
@@ -2551,9 +2567,10 @@ impl<'db> ClassLiteral<'db> {
25512567 Some ( Type :: function_like_callable ( db, signature) )
25522568 }
25532569 ( CodeGeneratorKind :: DataclassLike ( _) , "__hash__" ) => {
2554- let unsafe_hash = has_dataclass_param ( DataclassFlags :: UNSAFE_HASH ) ;
2555- let frozen = has_dataclass_param ( DataclassFlags :: FROZEN ) ;
2556- let eq = has_dataclass_param ( DataclassFlags :: EQ ) ;
2570+ let unsafe_hash =
2571+ self . has_dataclass_param ( db, field_policy, DataclassFlags :: UNSAFE_HASH ) ;
2572+ let frozen = self . has_dataclass_param ( db, field_policy, DataclassFlags :: FROZEN ) ;
2573+ let eq = self . has_dataclass_param ( db, field_policy, DataclassFlags :: EQ ) ;
25572574
25582575 if unsafe_hash || ( frozen && eq) {
25592576 let signature = Signature :: new (
@@ -2576,11 +2593,12 @@ impl<'db> ClassLiteral<'db> {
25762593 ( CodeGeneratorKind :: DataclassLike ( _) , "__match_args__" )
25772594 if Program :: get ( db) . python_version ( db) >= PythonVersion :: PY310 =>
25782595 {
2579- if !has_dataclass_param ( DataclassFlags :: MATCH_ARGS ) {
2596+ if !self . has_dataclass_param ( db , field_policy , DataclassFlags :: MATCH_ARGS ) {
25802597 return None ;
25812598 }
25822599
2583- let kw_only_default = has_dataclass_param ( DataclassFlags :: KW_ONLY ) ;
2600+ let kw_only_default =
2601+ self . has_dataclass_param ( db, field_policy, DataclassFlags :: KW_ONLY ) ;
25842602
25852603 let fields = self . fields ( db, specialization, field_policy) ;
25862604 let match_args = fields
@@ -2598,8 +2616,8 @@ impl<'db> ClassLiteral<'db> {
25982616 ( CodeGeneratorKind :: DataclassLike ( _) , "__weakref__" )
25992617 if Program :: get ( db) . python_version ( db) >= PythonVersion :: PY311 =>
26002618 {
2601- if !has_dataclass_param ( DataclassFlags :: WEAKREF_SLOT )
2602- || !has_dataclass_param ( DataclassFlags :: SLOTS )
2619+ if !self . has_dataclass_param ( db , field_policy , DataclassFlags :: WEAKREF_SLOT )
2620+ || !self . has_dataclass_param ( db , field_policy , DataclassFlags :: SLOTS )
26032621 {
26042622 return None ;
26052623 }
@@ -2641,7 +2659,7 @@ impl<'db> ClassLiteral<'db> {
26412659 signature_from_fields ( vec ! [ self_parameter] , Some ( instance_ty) )
26422660 }
26432661 ( CodeGeneratorKind :: DataclassLike ( _) , "__setattr__" ) => {
2644- if has_dataclass_param ( DataclassFlags :: FROZEN ) {
2662+ if self . has_dataclass_param ( db , field_policy , DataclassFlags :: FROZEN ) {
26452663 let signature = Signature :: new (
26462664 Parameters :: new (
26472665 db,
@@ -2662,11 +2680,12 @@ impl<'db> ClassLiteral<'db> {
26622680 ( CodeGeneratorKind :: DataclassLike ( _) , "__slots__" )
26632681 if Program :: get ( db) . python_version ( db) >= PythonVersion :: PY310 =>
26642682 {
2665- has_dataclass_param ( DataclassFlags :: SLOTS ) . then ( || {
2666- let fields = self . fields ( db, specialization, field_policy) ;
2667- let slots = fields. keys ( ) . map ( |name| Type :: string_literal ( db, name) ) ;
2668- Type :: heterogeneous_tuple ( db, slots)
2669- } )
2683+ self . has_dataclass_param ( db, field_policy, DataclassFlags :: SLOTS )
2684+ . then ( || {
2685+ let fields = self . fields ( db, specialization, field_policy) ;
2686+ let slots = fields. keys ( ) . map ( |name| Type :: string_literal ( db, name) ) ;
2687+ Type :: heterogeneous_tuple ( db, slots)
2688+ } )
26702689 }
26712690 ( CodeGeneratorKind :: TypedDict , "__setitem__" ) => {
26722691 let fields = self . fields ( db, specialization, field_policy) ;
@@ -3052,6 +3071,42 @@ impl<'db> ClassLiteral<'db> {
30523071 . collect ( )
30533072 }
30543073
3074+ pub ( crate ) fn validate_members ( self , context : & InferContext < ' db , ' _ > ) {
3075+ let db = context. db ( ) ;
3076+ let Some ( field_policy) = CodeGeneratorKind :: from_class ( db, self , None ) else {
3077+ return ;
3078+ } ;
3079+ let class_body_scope = self . body_scope ( db) ;
3080+ let table = place_table ( db, class_body_scope) ;
3081+ let use_def = use_def_map ( db, class_body_scope) ;
3082+ for ( symbol_id, declarations) in use_def. all_end_of_scope_symbol_declarations ( ) {
3083+ let result = place_from_declarations ( db, declarations. clone ( ) ) ;
3084+ let attr = result. ignore_conflicting_declarations ( ) ;
3085+ let symbol = table. symbol ( symbol_id) ;
3086+ let name = symbol. name ( ) ;
3087+ if let Some ( Type :: FunctionLiteral ( literal) ) = attr. place . ignore_possibly_undefined ( )
3088+ && matches ! ( name. as_str( ) , "__setattr__" | "__delattr__" )
3089+ {
3090+ if let Some ( CodeGeneratorKind :: DataclassLike ( _) ) =
3091+ CodeGeneratorKind :: from_class ( db, self , None )
3092+ && self . has_dataclass_param ( db, field_policy, DataclassFlags :: FROZEN )
3093+ {
3094+ if let Some ( builder) = context. report_lint (
3095+ & UNSOUND_DATACLASS_METHOD_OVERRIDE ,
3096+ literal. node ( db, context. file ( ) , context. module ( ) ) ,
3097+ ) {
3098+ let mut diagnostic = builder. into_diagnostic ( format_args ! (
3099+ "Cannot overwrite attribute `{}` in class `{}`" ,
3100+ name,
3101+ self . name( db)
3102+ ) ) ;
3103+ diagnostic. info ( name) ;
3104+ }
3105+ }
3106+ }
3107+ }
3108+ }
3109+
30553110 /// Returns a list of all annotated attributes defined in the body of this class. This is similar
30563111 /// to the `__annotations__` attribute at runtime, but also contains default values.
30573112 ///
0 commit comments