fix: check AssureWritable() return value at all call sites (#26390) Fixes #26389 `AssureWritable()` returns `-1` on failure (OOM in `MutableMessage()`), but 13 call sites were ignoring this return value. When `AssureWritable()` fails, `self->message` still points to the shared read-only default instance. Subsequent mutations (`MergeFrom`, `CopyFrom`, `Clear`, `SetField`, etc.) then corrupt the shared global default instance, affecting all messages of the same type. Same bug class as the recent `FixupMessageAfterMerge` fix (ab14c0f8a). Closes #26390 COPYBARA_INTEGRATE_REVIEW=https://github.com/protocolbuffers/protobuf/pull/26390 from Oblivionsage:fix/check-assure-writable-return-value adbd20477e3a1c49fbe9ce195f65169aac4ce98d PiperOrigin-RevId: 912044414
diff --git a/python/google/protobuf/pyext/message.cc b/python/google/protobuf/pyext/message.cc index 63f4414..b7c0924 100644 --- a/python/google/protobuf/pyext/message.cc +++ b/python/google/protobuf/pyext/message.cc
@@ -1032,7 +1032,7 @@ int InitWKTOrMerge(const Descriptor* descriptor, PyObject* py_message, PyObject* value) { CMessage* cmessage = reinterpret_cast<CMessage*>(py_message); - AssureWritable(cmessage); + if (AssureWritable(cmessage) < 0) return -1; if (PyObject_TypeCheck(value, CMessage_Type)) { ScopedPyObjectPtr merged(MergeFrom(cmessage, value)); if (merged == nullptr) { @@ -1241,7 +1241,7 @@ (descriptor->message_type()->well_known_type() != Descriptor::WELLKNOWNTYPE_STRUCT)) { // Make the message exist even if the dict is empty. - AssureWritable(cmessage); + if (AssureWritable(cmessage) < 0) return -1; if (InitAttributes(cmessage, nullptr, value) < 0) { return -1; } @@ -1655,7 +1655,7 @@ if (InternalReleaseFieldByDescriptor(self, field_descriptor) < 0) { return -1; } - AssureWritable(self); + if (AssureWritable(self) < 0) return -1; Message* message = self->message; message->GetReflection()->ClearField(message, field_descriptor); return 0; @@ -1667,7 +1667,7 @@ if (PyString_AsStringAndSize(arg, &field_name, &field_size) < 0) { return nullptr; } - AssureWritable(self); + if (AssureWritable(self) < 0) return nullptr; bool is_in_oneof; const FieldDescriptor* field_descriptor = FindFieldWithOneofs( self->message, absl::string_view(field_name, field_size), &is_in_oneof); @@ -1689,7 +1689,7 @@ } PyObject* Clear(CMessage* self) { - AssureWritable(self); + if (AssureWritable(self) < 0) return nullptr; // Detach all current fields of this message std::vector<ScopedPyObjectPtr> messages_to_release; std::vector<ScopedPyObjectPtr> containers_to_release; @@ -1912,7 +1912,7 @@ .c_str()); return nullptr; } - AssureWritable(self); + if (AssureWritable(self) < 0) return nullptr; if (MaybeReleaseOneofBeforeMerge(self, *other_message->message) < 0) { return nullptr; @@ -1956,7 +1956,7 @@ return nullptr; } - AssureWritable(self); + if (AssureWritable(self) < 0) return nullptr; // CopyFrom on the message will not clean up self->composite_fields, // which can leave us in an inconsistent state, so clear it out here. @@ -1989,7 +1989,7 @@ return nullptr; } - AssureWritable(self); + if (AssureWritable(self) < 0) return nullptr; PyMessageFactory* factory = GetFactoryForMessage(self); int depth = allow_oversize_protos @@ -2053,7 +2053,7 @@ } static PyObject* SetInParent(CMessage* self, PyObject* args) { - AssureWritable(self); + if (AssureWritable(self) < 0) return nullptr; Py_RETURN_NONE; } @@ -2161,7 +2161,7 @@ } static PyObject* DiscardUnknownFields(CMessage* self) { - AssureWritable(self); + if (AssureWritable(self) < 0) return nullptr; self->message->DiscardUnknownFields(); Py_RETURN_NONE; } @@ -2736,7 +2736,7 @@ Descriptor::WELLKNOWNTYPE_UNSPECIFIED) { ScopedPyObjectPtr sub_message(GetFieldValue(self, field_descriptor)); if (PyObject_HasAttrString(sub_message.get(), "_internal_assign")) { - AssureWritable(self); + if (AssureWritable(self) < 0) return -1; ScopedPyObjectPtr ok(PyObject_CallMethod( sub_message.get(), "_internal_assign", "O", value)); if (ok.get() == nullptr) { @@ -2751,7 +2751,7 @@ std::string(field_descriptor->name()).c_str()); return -1; } else { - AssureWritable(self); + if (AssureWritable(self) < 0) return -1; return InternalSetScalar(self, field_descriptor, value); } } @@ -2937,7 +2937,7 @@ "to a message with extra references"); return nullptr; } - cmessage::AssureWritable(cmsg); + if (cmessage::AssureWritable(cmsg) < 0) return nullptr; return cmsg->message; }